diff --git a/solstone/apps/timeline/tests/conftest.py b/solstone/apps/timeline/tests/conftest.py index d7dc41494..dfee516c7 100644 --- a/solstone/apps/timeline/tests/conftest.py +++ b/solstone/apps/timeline/tests/conftest.py @@ -113,14 +113,17 @@ def mock_agenerate(monkeypatch): async def _fake_agenerate(**kwargs): if not responses: - return json.dumps({"picks": [0], "rationale": "default"}) + return { + "text": json.dumps({"picks": [0], "rationale": "default"}), + "finish_reason": "stop", + } item = responses.pop(0) if isinstance(item, Exception): raise item - return json.dumps(item) + return {"text": json.dumps(item), "finish_reason": "stop"} mock = AsyncMock(side_effect=_fake_agenerate) - monkeypatch.setattr("solstone.think.batch.agenerate", mock) + monkeypatch.setattr("solstone.think.batch.agenerate_with_result", mock) return mock return _install diff --git a/solstone/convey/provider_readiness.py b/solstone/convey/provider_readiness.py index 09589122d..8cffc0694 100644 --- a/solstone/convey/provider_readiness.py +++ b/solstone/convey/provider_readiness.py @@ -274,6 +274,15 @@ _ENTRIES: dict[str, _Entry] = { ), recovery_action=None, ), + "incomplete_text_length": _Entry( + klass="generic", + summary="the answer ran out of room before it finished", + detail=( + "The reply hit its length limit before it could finish. Try again " + "with less at once or choose another provider." + ), + recovery_action=None, + ), "max_turns_exhausted": _Entry( klass="generic", summary="this took too many steps to finish", diff --git a/solstone/convey/static/chat_reasons.js b/solstone/convey/static/chat_reasons.js index 1786a6026..742772bc3 100644 --- a/solstone/convey/static/chat_reasons.js +++ b/solstone/convey/static/chat_reasons.js @@ -122,6 +122,10 @@ "template": "the answer ran out of room before it finished", "action": null }, + "incomplete_text_length": { + "template": "the answer ran out of room before it finished", + "action": null + }, "max_turns_exhausted": { "template": "this took too many steps to finish", "action": null diff --git a/solstone/observe/categories/browsing.md b/solstone/observe/categories/browsing.md index 347625b2c..e214d23f9 100644 --- a/solstone/observe/categories/browsing.md +++ b/solstone/observe/categories/browsing.md @@ -2,7 +2,8 @@ "description": "General web browsing, news, shopping, or reference pages without a dominant social feed or media viewer", "output": "markdown", - "extraction": "Extract when visiting distinctly different websites or search results" + "extraction": "Extract when visiting distinctly different websites or search results", + "max_output_tokens": 2048 } diff --git a/solstone/observe/describe.py b/solstone/observe/describe.py index d4ca9a0fc..9a98dc67d 100644 --- a/solstone/observe/describe.py +++ b/solstone/observe/describe.py @@ -792,7 +792,7 @@ class VideoProcessor: json_output=True, json_schema=_SCHEMA, temperature=0.7, - max_output_tokens=1024, + max_output_tokens=512, thinking_budget=1024, ) diff --git a/solstone/talent/sense.md b/solstone/talent/sense.md index 771b04421..f4f2bae56 100644 --- a/solstone/talent/sense.md +++ b/solstone/talent/sense.md @@ -9,7 +9,7 @@ "tier": 3, "output": "json", "schema": "sense.schema.json", - "max_output_tokens": 12288, + "max_output_tokens": 6144, "timeout_s": 480, "load": {"transcripts": true, "percepts": true, "talents": false} } diff --git a/solstone/talent/sense.schema.json b/solstone/talent/sense.schema.json index f95215503..59bb89af4 100644 --- a/solstone/talent/sense.schema.json +++ b/solstone/talent/sense.schema.json @@ -49,6 +49,7 @@ }, "entities": { "type": "array", + "maxItems": 96, "items": { "type": "object", "additionalProperties": false, @@ -106,6 +107,7 @@ }, "facets": { "type": "array", + "maxItems": 16, "items": { "type": "object", "additionalProperties": false, @@ -146,6 +148,7 @@ }, "speakers": { "type": "array", + "maxItems": 16, "items": { "type": "string" } diff --git a/solstone/think/batch.py b/solstone/think/batch.py index f5379a360..a63240e3a 100644 --- a/solstone/think/batch.py +++ b/solstone/think/batch.py @@ -6,7 +6,7 @@ Async batch processing for LLM API requests. Provides Batch for concurrent execution of multiple LLM API calls with dynamic request queuing and result streaming via async iterator. -Routes requests to providers based on context via the unified agenerate() API. +Routes requests to providers based on context via the unified async generate API. Example: batch = Batch(max_concurrent=5) @@ -26,7 +26,12 @@ import asyncio import time from typing import Any, List, Optional, Union -from solstone.think.models import agenerate, resolve_provider +from solstone.think.models import ( + SchemaValidationError, + agenerate_with_result, + finish_reason_error, + resolve_provider, +) from solstone.think.providers.shared import classify_provider_error @@ -34,7 +39,7 @@ class BatchRequest: """ Mutable request object for a single LLM API call. - Core attributes are passed to agenerate(). Callers can add + Core attributes are passed to agenerate_with_result(). Callers can add arbitrary attributes for tracking (e.g., frame_id, stage, etc). After execution, these attributes are populated: @@ -83,7 +88,7 @@ class Batch: Async batch processor for LLM API requests. Manages concurrent execution with dynamic request queuing and result - streaming via async iterator pattern. Routes to providers via agenerate(). + streaming via async iterator pattern. Routes to providers via async generation. Example: batch = Batch(max_concurrent=5) @@ -265,7 +270,7 @@ class Batch: if request.model is not None: kwargs["model"] = request.model - response = await agenerate( + result = await agenerate_with_result( contents=request.contents, context=request.context, temperature=request.temperature, @@ -277,8 +282,20 @@ class Batch: timeout_s=request.timeout_s, **kwargs, ) + error = finish_reason_error( + result, + json_output=request.json_output, + ) + if error is not None: + raise error + validation = result.get("schema_validation") + if isinstance(validation, dict) and validation.get("valid") is False: + raise SchemaValidationError( + validation.get("errors") or [], + result.get("text", ""), + ) request.duration = time.time() - start_time - request.response = response + request.response = result["text"] request.error = None # Track which model was actually used diff --git a/solstone/think/models.py b/solstone/think/models.py index 48887187e..e88e9c34a 100644 --- a/solstone/think/models.py +++ b/solstone/think/models.py @@ -232,6 +232,26 @@ class IncompleteJSONError(ValueError): super().__init__(f"JSON response incomplete (reason: {reason})") +class IncompleteTextError(ValueError): + """Raised when a non-JSON response is truncated due to token limits.""" + + def __init__(self, reason: str, partial_text: str): + self.reason = reason + self.partial_text = partial_text + self.reason_code = "incomplete_text_length" + super().__init__(f"Text response incomplete (reason: {reason})") + + +class ProviderResponseInvalidError(ValueError): + """Raised when a provider reports a non-success finish for plain text.""" + + reason_code = "provider_response_invalid" + + def __init__(self, reason: str): + self.reason = reason + super().__init__(f"Provider response did not finish cleanly (reason: {reason})") + + class SchemaValidationError(ValueError): """Raised when JSON response text fails local schema validation. @@ -1293,20 +1313,42 @@ def get_usage_cost( # --------------------------------------------------------------------------- +def finish_reason_error( + result: Dict[str, Any], + *, + json_output: bool, +) -> Exception | None: + """Map a finish reason to the error it should raise, or None if acceptable.""" + finish_reason = result.get("finish_reason") + if not finish_reason or finish_reason == "stop": + return None + + if json_output: + return IncompleteJSONError( + reason=finish_reason, + partial_text=result.get("text", ""), + ) + + if str(finish_reason).strip().lower() in _LENGTH_FINISH_REASONS: + return IncompleteTextError( + reason=finish_reason, + partial_text=result.get("text", ""), + ) + return ProviderResponseInvalidError(reason=finish_reason) + + def _validate_json_response(result: Dict[str, Any], json_output: bool) -> None: """Validate response for JSON output mode. - Raises IncompleteJSONError if finish_reason indicates truncation. + Raises IncompleteJSONError if finish_reason is a present non-stop value. """ + # Non-JSON generate() callers (planner, depict, importers, enrich, extract, + # transcribe, detect_*) keep today's leniency; the Batch boundary is strict. if not json_output: return - - finish_reason = result.get("finish_reason") - if finish_reason and finish_reason != "stop": - raise IncompleteJSONError( - reason=finish_reason, - partial_text=result.get("text", ""), - ) + error = finish_reason_error(result, json_output=True) + if error is not None: + raise error def _validate_schema(text: str, schema: dict) -> dict: @@ -1717,6 +1759,75 @@ def generate_with_result( return result +async def agenerate_with_result( + contents: Union[str, List[Any]], + context: str, + temperature: float = 0.3, + max_output_tokens: int = 8192 * 2, + system_instruction: Optional[str] = None, + json_output: bool = False, + *, + json_schema: dict | None = None, + thinking_budget: Optional[int] = None, + timeout_s: Optional[float] = None, + **kwargs: Any, +) -> dict: + """Async generate text and return the full GenerateResult dict.""" + from solstone.think.providers import get_provider_module + + if json_schema is not None: + json_output = True + + model_override = kwargs.pop("model", None) + provider_override = kwargs.pop("provider", None) + + provider, model = resolve_provider(context, "generate") + if provider_override: + provider = provider_override + if not model_override: + model = resolve_model_for_provider(context, provider, "generate") + if model_override: + model = model_override + + _raise_if_no_brain(provider) + _reject_local_cloud_model_override(provider, model_override) + _raise_if_confidential_unverified() + + provider_mod = get_provider_module(provider) + provider_schema = prepare_provider_schema(json_schema, provider) + + timeout_s = DEFAULT_PROVIDER_TIMEOUT_S if timeout_s is None else timeout_s + + result = await provider_mod.run_agenerate( + contents=contents, + model=model, + provider=provider, + temperature=temperature, + max_output_tokens=max_output_tokens, + system_instruction=system_instruction, + json_output=json_output, + json_schema=provider_schema, + thinking_budget=thinking_budget, + timeout_s=timeout_s, + **kwargs, + ) + + if result.get("usage"): + log_token_usage( + model=result.get("model") or model, + usage=result["usage"], + context=context, + type="generate", + ) + + _validate_json_response(result, json_output) + + if json_schema is not None: + result["schema_validation"] = _validate_schema(result["text"], json_schema) + + return result + + async def agenerate( contents: Union[str, List[Any]], context: str, @@ -1855,6 +1966,10 @@ __all__ = [ "generate", "generate_with_result", "agenerate", + "agenerate_with_result", + "finish_reason_error", + "IncompleteTextError", + "ProviderResponseInvalidError", "resolve_provider", "resolve_effective_route", "is_local_provider_needed", diff --git a/solstone/think/providers/local.py b/solstone/think/providers/local.py index 124ed179d..bd5def128 100644 --- a/solstone/think/providers/local.py +++ b/solstone/think/providers/local.py @@ -60,6 +60,13 @@ _QWEN_TOP_P = 0.8 _QWEN_TOP_K = 20 _QWEN_MIN_P = 0.0 _QWEN_PRESENCE_PENALTY = 1.5 +_LOCAL_FINISH_REASON_MAP = { + "stop": "stop", + "length": "max_tokens", + "max_tokens": "max_tokens", + "content_filter": "content_filter", +} +_LOCAL_UNSUPPORTED_FINISH_REASONS = frozenset({"tool_calls", "function_call"}) @dataclass(frozen=True) @@ -280,6 +287,27 @@ def _extract_usage(data: dict[str, Any]) -> dict[str, int] | None: return normalized +def _normalize_finish_reason(raw: Any) -> str: + if not isinstance(raw, str) or not raw.strip(): + raise LocalProviderError( + "provider_response_invalid", + "Local model response did not include a finish reason.", + ) + reason = raw.strip().lower() + normalized = _LOCAL_FINISH_REASON_MAP.get(reason) + if normalized is not None: + return normalized + if reason in _LOCAL_UNSUPPORTED_FINISH_REASONS: + raise LocalProviderError( + "provider_response_invalid", + f"Local model returned unsupported finish reason: {reason}", + ) + raise LocalProviderError( + "provider_response_invalid", + f"Local model returned unknown finish reason: {reason}", + ) + + def _parse_response(data: dict[str, Any]) -> GenerateResult: choices = data.get("choices") if not isinstance(choices, list) or not choices: @@ -298,7 +326,7 @@ def _parse_response(data: dict[str, Any]) -> GenerateResult: text=text, model=LOCAL_MODEL, usage=_extract_usage(data), - finish_reason=choice.get("finish_reason"), + finish_reason=_normalize_finish_reason(choice.get("finish_reason")), thinking=None, ) diff --git a/solstone/think/providers/shared.py b/solstone/think/providers/shared.py index e67a74516..a3a903295 100644 --- a/solstone/think/providers/shared.py +++ b/solstone/think/providers/shared.py @@ -222,6 +222,7 @@ RUNTIME_REASON_CODES = frozenset( "provider_unavailable", "provider_response_invalid", "incomplete_json_length", + "incomplete_text_length", "unknown", } ) diff --git a/solstone/think/talents.py b/solstone/think/talents.py index 4ba7d11e6..fd0ac2a2d 100644 --- a/solstone/think/talents.py +++ b/solstone/think/talents.py @@ -1607,20 +1607,24 @@ async def _execute_generate( exc.reason, ) retries = 1 - gen_result = generate_with_result( - contents=contents, - context=context, - temperature=max(temperature, _LOCAL_LENGTH_RETRY_TEMPERATURE_FLOOR), - max_output_tokens=max_output_tokens, - thinking_budget=thinking_budget, - system_instruction=system_instruction, - json_output=is_json_output, - json_schema=runtime_json_schema, - timeout_s=timeout_s, - provider=config.get("provider"), - model=config.get("model"), - inference_retry_index=1, - ) + try: + gen_result = generate_with_result( + contents=contents, + context=context, + temperature=max(temperature, _LOCAL_LENGTH_RETRY_TEMPERATURE_FLOOR), + max_output_tokens=max_output_tokens, + thinking_budget=thinking_budget, + system_instruction=system_instruction, + json_output=is_json_output, + json_schema=runtime_json_schema, + timeout_s=timeout_s, + provider=config.get("provider"), + model=config.get("model"), + inference_retry_index=1, + ) + except Exception as retry_exc: + retry_exc.retries = retries + raise else: if config.get("fallback_from") or not _should_fallback(exc): raise @@ -2024,6 +2028,9 @@ async def main_async() -> None: event["partial_text_tail"] = e.partial_text[-500:] name = config.get("name", "unknown") if config else "unknown" log_extraction_failure(e, name) + retries = getattr(e, "retries", None) + if retries: + event["retries"] = retries emit_event(event) except Exception as exc: diff --git a/tests/baselines/api/stats/stats.json b/tests/baselines/api/stats/stats.json index 302e3afd5..e5ce75eeb 100644 --- a/tests/baselines/api/stats/stats.json +++ b/tests/baselines/api/stats/stats.json @@ -365,7 +365,7 @@ "talents": false, "transcripts": true }, - "max_output_tokens": 12288, + "max_output_tokens": 6144, "mtime": 0, "output": "json", "path": "/solstone/talent/sense.md", diff --git a/tests/test_bad_media_corpus.py b/tests/test_bad_media_corpus.py index 754dd4d5c..760fe4c35 100644 --- a/tests/test_bad_media_corpus.py +++ b/tests/test_bad_media_corpus.py @@ -54,6 +54,10 @@ SEGMENT = "120000_300" FIXED_NOW = "2026-06-30T12:00:00Z" +def _generate_result(text: str, finish_reason: str = "stop") -> dict[str, Any]: + return {"text": text, "finish_reason": finish_reason} + + @pytest.fixture def observer_env(tmp_path, monkeypatch): """Temp journal + Flask test client factory. @@ -244,11 +248,14 @@ def _drive_describe( output_path: Path, *, agenerate_response: str = "{}", + agenerate_finish_reason: str = "stop", expect_runtime_error: bool = False, ) -> tuple[dict[str, Any], dict[str, Any], AsyncMock]: from solstone.observe import describe, processing_record - agenerate = AsyncMock(return_value=agenerate_response) + agenerate = AsyncMock( + return_value=_generate_result(agenerate_response, agenerate_finish_reason) + ) monkeypatch.setattr( "solstone.think.models.resolve_provider", lambda _context, _interface: ("google", "gemini-test"), @@ -256,7 +263,7 @@ def _drive_describe( monkeypatch.setattr(describe, "callosum_send", lambda *args, **kwargs: None) monkeypatch.setattr(describe, "select_frames_for_extraction", lambda *a, **k: []) monkeypatch.setattr(processing_record, "now_iso_utc", lambda: FIXED_NOW) - monkeypatch.setattr("solstone.think.batch.agenerate", agenerate) + monkeypatch.setattr("solstone.think.batch.agenerate_with_result", agenerate) processor = describe.VideoProcessor(video_path) if expect_runtime_error: @@ -366,7 +373,7 @@ def _run_idle_gate( spawned: list[str] = [] writer_path = journal / "chronicle" / day / "health" / f"idle_{segment}.jsonl" writer = ThinkingJSONLWriter(str(writer_path)) - agenerate = agenerate_spy or AsyncMock(return_value="{}") + agenerate = agenerate_spy or AsyncMock(return_value=_generate_result("{}")) original_callosum = think._callosum original_jsonl = think._jsonl try: @@ -387,7 +394,7 @@ def _run_idle_gate( "wait_for_uses", lambda agent_ids, timeout=600: ({aid: "finish" for aid in agent_ids}, []), ) - monkeypatch.setattr("solstone.think.batch.agenerate", agenerate) + monkeypatch.setattr("solstone.think.batch.agenerate_with_result", agenerate) think._callosum = None think._jsonl = writer result = think.run_segment_sense( @@ -649,6 +656,34 @@ def test_ac5_all_frames_fail_is_analysis_failed_distinct( assert read_segment_data_state(DAY, SEGMENT) == {"screen": DataState.FAILED.value} +def test_truncated_frame_categorization_retries_and_promotes_no_frame_artifact( + segment_journal, + monkeypatch, +): + segment = _segment_dir(segment_journal) + video_path = segment / "screen.mp4" + output_path = segment / "screen.jsonl" + _build_one_frame_mp4(video_path) + + _header, record, agenerate = _drive_describe( + monkeypatch, + video_path, + output_path, + agenerate_response='{"visual_description":"partial"', + agenerate_finish_reason="max_tokens", + expect_runtime_error=True, + ) + + _assert_processing_record( + record, + state=STATE_FAILED, + reason_code=REASON_ANALYSIS_FAILED, + handler=HANDLER_DESCRIBE, + ) + assert agenerate.call_count == 5 + assert len(_read_jsonl(output_path)) == 1 + + def test_ac6_no_model_calls_on_all_empty_segment(segment_journal, monkeypatch): segment = _segment_dir(segment_journal) screen_path = segment / "screen.mp4" diff --git a/tests/test_batch.py b/tests/test_batch.py index 126e46d34..9ab5f4022 100644 --- a/tests/test_batch.py +++ b/tests/test_batch.py @@ -12,6 +12,10 @@ from solstone.think.batch import Batch, BatchRequest from solstone.think.models import GEMINI_FLASH, GEMINI_LITE, SchemaValidationError +def _result(text: str = "Response", finish_reason: str = "stop", **extra): + return {"text": text, "finish_reason": finish_reason, **extra} + + def test_batch_request_creation(): """Test BatchRequest can be created with required and custom params.""" # Required params only @@ -50,10 +54,10 @@ def test_batch_request_custom_attributes(): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_basic(mock_agenerate): """Test basic batch execution with single request.""" - mock_agenerate.return_value = "Response 1" + mock_agenerate.return_value = _result("Response 1") # Create batch and add request batch = Batch(max_concurrent=5) @@ -80,10 +84,10 @@ async def test_batch_basic(mock_agenerate): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_with_model_override(mock_agenerate): """Test batch with explicit model override.""" - mock_agenerate.return_value = "Response" + mock_agenerate.return_value = _result("Response") batch = Batch(max_concurrent=5) req = batch.create(contents="Test", context="test.context", model=GEMINI_FLASH) @@ -102,10 +106,14 @@ async def test_batch_with_model_override(mock_agenerate): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_multiple_requests(mock_agenerate): """Test batch with multiple requests.""" - mock_agenerate.side_effect = ["Response 1", "Response 2", "Response 3"] + mock_agenerate.side_effect = [ + _result("Response 1"), + _result("Response 2"), + _result("Response 3"), + ] batch = Batch(max_concurrent=2) @@ -141,7 +149,7 @@ async def test_batch_multiple_requests(mock_agenerate): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_error_handling(mock_agenerate): """Test that errors are captured in request.error.""" mock_agenerate.side_effect = ValueError("API error") @@ -165,7 +173,7 @@ async def test_batch_error_handling(mock_agenerate): @pytest.mark.asyncio @patch("solstone.think.batch.resolve_provider") -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_preserves_exception_metadata( mock_agenerate, mock_resolve_provider ): @@ -197,7 +205,7 @@ async def test_batch_preserves_exception_metadata( @pytest.mark.asyncio @patch("solstone.think.batch.resolve_provider") -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_classifies_exception_metadata_when_attrs_missing( mock_agenerate, mock_resolve_provider ): @@ -238,10 +246,10 @@ def test_batch_update_clears_error_metadata(): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_dynamic_adding(mock_agenerate): """Test adding requests dynamically during iteration.""" - mock_agenerate.return_value = "Response" + mock_agenerate.return_value = _result("Response") batch = Batch(max_concurrent=5) @@ -271,7 +279,7 @@ async def test_batch_dynamic_adding(mock_agenerate): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_retry_pattern(mock_agenerate): """Test retry pattern - add failed request back with different model.""" # First call fails, second succeeds @@ -282,7 +290,7 @@ async def test_batch_retry_pattern(mock_agenerate): call_count += 1 if call_count == 1: raise ValueError("Transient error") - return "Success on retry" + return _result("Success on retry") mock_agenerate.side_effect = mock_response @@ -314,10 +322,10 @@ async def test_batch_retry_pattern(mock_agenerate): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_factory_method(mock_agenerate): """Test that batch.create() factory method works correctly.""" - mock_agenerate.return_value = "Response" + mock_agenerate.return_value = _result("Response") batch = Batch() @@ -339,10 +347,10 @@ async def test_batch_factory_method(mock_agenerate): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_can_add_after_draining(mock_agenerate): """Test that adding after draining works (reusable batch).""" - mock_agenerate.side_effect = ["Response 1", "Response 2"] + mock_agenerate.side_effect = [_result("Response 1"), _result("Response 2")] batch = Batch() @@ -372,7 +380,7 @@ async def test_batch_can_add_after_draining(mock_agenerate): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_empty_batch(mock_agenerate): """Test that empty batch (no requests) completes immediately.""" batch = Batch() @@ -385,7 +393,7 @@ async def test_batch_empty_batch(mock_agenerate): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_concurrency_limit(mock_agenerate): """Test that semaphore limits concurrent requests.""" # Track concurrent calls @@ -404,7 +412,7 @@ async def test_batch_concurrency_limit(mock_agenerate): async with lock: concurrent_calls -= 1 - return "Response" + return _result("Response") mock_agenerate.side_effect = mock_with_tracking @@ -426,7 +434,7 @@ async def test_batch_concurrency_limit(mock_agenerate): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_update_method(mock_agenerate): """Test batch.update() method for modifying and re-adding requests.""" # Track which model was used in each call @@ -434,7 +442,7 @@ async def test_batch_update_method(mock_agenerate): async def mock_track_model(*args, **kwargs): call_models.append(kwargs.get("model", "unknown")) - return f"Response from {kwargs.get('model', 'unknown')}" + return _result(f"Response from {kwargs.get('model', 'unknown')}") mock_agenerate.side_effect = mock_track_model @@ -494,10 +502,10 @@ def test_batch_request_with_timeout(): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_timeout_passthrough(mock_agenerate): - """Test that timeout_s is passed through to agenerate.""" - mock_agenerate.return_value = "Response" + """Test that timeout_s is passed through to agenerate_with_result.""" + mock_agenerate.return_value = _result("Response") batch = Batch(max_concurrent=5) @@ -511,17 +519,17 @@ async def test_batch_timeout_passthrough(mock_agenerate): assert len(results) == 1 - # Verify timeout_s was passed to agenerate + # Verify timeout_s was passed to agenerate_with_result mock_agenerate.assert_called_once() call_kwargs = mock_agenerate.call_args[1] assert call_kwargs["timeout_s"] == 45 @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_client_passthrough(mock_agenerate): - """Test that client is passed through to agenerate for Google connection reuse.""" - mock_agenerate.return_value = "Response" + """Test that client is passed through for Google connection reuse.""" + mock_agenerate.return_value = _result("Response") # Create a mock client (would be genai.Client for Google) mock_client = object() @@ -542,10 +550,10 @@ async def test_batch_client_passthrough(mock_agenerate): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_passes_json_schema_to_agenerate(mock_agenerate): - """Test that json_schema is passed through to agenerate.""" - mock_agenerate.return_value = "Response" + """Test that json_schema is passed through to agenerate_with_result.""" + mock_agenerate.return_value = _result("Response") batch = Batch(max_concurrent=5) req = batch.create( @@ -566,7 +574,7 @@ async def test_batch_passes_json_schema_to_agenerate(mock_agenerate): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_batch_schema_validation_error_populates_request_error(mock_agenerate): mock_agenerate.side_effect = SchemaValidationError( [{"path": "", "constraint": "json_parse", "message": "empty"}], @@ -588,3 +596,93 @@ async def test_batch_schema_validation_error_populates_request_error(mock_agener assert len(results) == 1 assert results[0].response is None assert "schema validation" in results[0].error + + +@pytest.mark.asyncio +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) +async def test_batch_non_json_length_finish_populates_text_length_error(mock_agenerate): + mock_agenerate.return_value = _result("partial text", finish_reason="max_tokens") + + batch = Batch(max_concurrent=5) + req = batch.create(contents="Test prompt", context="test.context") + batch.add(req) + + results = [] + async for completed_req in batch.drain_batch(): + results.append(completed_req) + + assert len(results) == 1 + assert results[0].response is None + assert results[0].reason_code == "incomplete_text_length" + assert "Text response incomplete" in results[0].error + + +@pytest.mark.asyncio +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) +async def test_batch_non_json_stop_finish_succeeds(mock_agenerate): + mock_agenerate.return_value = _result("complete text", finish_reason="stop") + + batch = Batch(max_concurrent=1) + req = batch.create( + contents="Test", + context="test.context", + json_output=False, + ) + batch.add(req) + + results = [] + async for completed_req in batch.drain_batch(): + results.append(completed_req) + + assert len(results) == 1 + assert req.response == "complete text" + assert req.error is None + assert req.reason_code is None + + +@pytest.mark.asyncio +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) +async def test_batch_non_json_non_length_finish_is_provider_response_invalid( + mock_agenerate, +): + mock_agenerate.return_value = _result("", finish_reason="content_filter") + + batch = Batch(max_concurrent=5) + req = batch.create(contents="Test prompt", context="test.context") + batch.add(req) + + results = [] + async for completed_req in batch.drain_batch(): + results.append(completed_req) + + assert len(results) == 1 + assert results[0].response is None + assert results[0].reason_code == "provider_response_invalid" + + +@pytest.mark.asyncio +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) +async def test_batch_full_result_schema_invalid_stays_hard_failure(mock_agenerate): + mock_agenerate.return_value = _result( + '{"field": "bad"}', + schema_validation={ + "valid": False, + "errors": [{"path": "/field", "constraint": "type", "message": "bad"}], + }, + ) + + batch = Batch(max_concurrent=5) + req = batch.create( + contents="Test prompt", + context="test.context", + json_schema={"type": "object"}, + ) + batch.add(req) + + results = [] + async for completed_req in batch.drain_batch(): + results.append(completed_req) + + assert len(results) == 1 + assert results[0].response is None + assert "schema validation" in results[0].error diff --git a/tests/test_calendar_schema.py b/tests/test_calendar_schema.py index 0df883c0f..27a730804 100644 --- a/tests/test_calendar_schema.py +++ b/tests/test_calendar_schema.py @@ -82,9 +82,12 @@ def test_discover_categories_attaches_calendar_schema(): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_calendar_extract_batch_call_passes_schema(mock_agenerate): - mock_agenerate.return_value = json.dumps(_valid_payload()) + mock_agenerate.return_value = { + "text": json.dumps(_valid_payload()), + "finish_reason": "stop", + } cat_meta = describe_mod.CATEGORIES["calendar"] batch = Batch(max_concurrent=1) diff --git a/tests/test_chat_reasons.py b/tests/test_chat_reasons.py index 0ef182bc6..288f319e9 100644 --- a/tests/test_chat_reasons.py +++ b/tests/test_chat_reasons.py @@ -39,6 +39,7 @@ EXPECTED_CODES = { "chat_timeout", "context_window_exceeded", "incomplete_json_length", + "incomplete_text_length", "max_turns_exhausted", "no_output", "token_budget_exceeded", diff --git a/tests/test_describe_promote.py b/tests/test_describe_promote.py index 37590cfed..a0ddd1c35 100644 --- a/tests/test_describe_promote.py +++ b/tests/test_describe_promote.py @@ -6,6 +6,7 @@ import json import logging from pathlib import Path from types import SimpleNamespace +from unittest.mock import AsyncMock import pytest from PIL import Image @@ -65,6 +66,10 @@ def _jsonl_rows(path: Path) -> list[dict]: ] +def _generate_result(text: str, finish_reason: str = "stop") -> dict: + return {"text": text, "finish_reason": finish_reason} + + def _canned_detection() -> dict: return { "image": {"width": 8, "height": 8}, @@ -282,6 +287,101 @@ async def test_success_with_mixed_results_promotes_byte_identical_jsonl( _assert_no_describe_temp(output_path.parent) +@pytest.mark.asyncio +async def test_browsing_truncation_does_not_promote_category_content( + tmp_path, monkeypatch +): + from solstone.think import batch as batch_module + from solstone.think import models + + video_path = _video_path(tmp_path) + output_path = video_path.with_suffix(".jsonl") + frame_bytes = _png_bytes() + processor = _processor( + video_path, + [_frame(1, 0.0, frame_bytes)], + monkeypatch, + ) + completed = [] + real_batch = batch_module.Batch + + class SpyBatch(real_batch): + async def drain_batch(self): + async for request in super().drain_batch(): + completed.append( + { + "request_type": getattr(request, "request_type", None), + "reason_code": getattr(request, "reason_code", None), + "retry_count": getattr(request, "retry_count", None), + "extraction_category": getattr( + request, "extraction_category", None + ), + } + ) + yield request + + categorize = _generate_result( + json.dumps( + { + "visual_description": "A browser page is open.", + "primary": "browsing", + "secondary": "none", + "overlap": True, + } + ) + ) + truncated = _generate_result("partial browsing notes", "max_tokens") + agenerate = AsyncMock(return_value=truncated) + agenerate.side_effect = [categorize, *[truncated for _ in range(5)]] + + monkeypatch.setattr(batch_module, "Batch", SpyBatch) + monkeypatch.setattr(batch_module, "agenerate_with_result", agenerate) + monkeypatch.setattr( + batch_module, + "resolve_provider", + lambda _context, _interface: ("google", "gemini-test"), + ) + monkeypatch.setattr( + models, + "resolve_provider", + lambda _context, _interface: ("google", "gemini-test"), + ) + monkeypatch.setattr( + processing_record_module, "now_iso_utc", lambda: "2026-06-30T12:00:00Z" + ) + monkeypatch.setattr(describe_module, "callosum_send", lambda *a, **k: None) + monkeypatch.setattr(describe_module, "get_config", lambda: {"describe": {}}) + monkeypatch.setattr( + describe_module, + "select_frames_for_extraction", + lambda *_args, **_kwargs: [1], + ) + + await processor.process_with_vision( + max_concurrent=1, + output_path=output_path, + work_key="20250101/143022_300/screen", + ) + + rows = _jsonl_rows(output_path) + result = rows[1] + category_requests = [ + request + for request in completed + if request["request_type"] == describe_module.RequestType.CATEGORY + ] + + assert agenerate.call_count == 6 + assert len(category_requests) == 5 + assert category_requests[-1]["reason_code"] == "incomplete_text_length" + assert result["enhanced"] is True + assert result["content"] == {} + assert "browsing" not in result["content"] + assert result["error"] == "Text response incomplete (reason: max_tokens)" + assert result["requests"][-1]["category"] == "browsing" + assert result["requests"][-1]["retries"] == 4 + + @pytest.mark.asyncio async def test_detection_blocks_attach_to_media_and_social_frames( tmp_path, monkeypatch diff --git a/tests/test_generate_talents.py b/tests/test_generate_talents.py index c3fe0ffc6..c024c038c 100644 --- a/tests/test_generate_talents.py +++ b/tests/test_generate_talents.py @@ -48,20 +48,28 @@ def test_json_extraction_talents_pin_output_cap_and_timeout(): largest_observed_legitimate_completion = 3560 - for name in ("sense", "participation"): - config = get_talent(name) - params = _generation_params(config) - max_output_tokens = params["max_output_tokens"] - # Mirror talents.py's resolution: frontmatter timeout_s short-circuits - # the derivation. - resolved = config.get("timeout_s") or min( - 480, - max(120, (max_output_tokens + params["thinking_budget"]) // 100), - ) + sense_config = get_talent("sense") + sense_params = _generation_params(sense_config) + assert sense_params["max_output_tokens"] == 6144 + assert sense_config.get("timeout_s") == 480 + assert "temperature" not in sense_config - assert max_output_tokens >= 2 * largest_observed_legitimate_completion - assert max_output_tokens < 8192 * 6 - assert config.get("timeout_s") == 480 - assert resolved == config["timeout_s"] - assert resolved >= 480 - assert "temperature" not in config + participation_config = get_talent("participation") + participation_params = _generation_params(participation_config) + participation_tokens = participation_params["max_output_tokens"] + # Participation did not change and still keeps 2x headroom over the + # largest legitimate completion observed when this guard was added. + assert participation_tokens == 12288 + assert participation_tokens >= 2 * largest_observed_legitimate_completion + assert participation_tokens < 8192 * 6 + assert participation_config.get("timeout_s") == 480 + resolved = participation_config.get("timeout_s") or min( + 480, + max( + 120, + (participation_tokens + participation_params["thinking_budget"]) // 100, + ), + ) + assert resolved == participation_config["timeout_s"] + assert resolved >= 480 + assert "temperature" not in participation_config diff --git a/tests/test_local.py b/tests/test_local.py index b5cedc1a3..c18fa2365 100644 --- a/tests/test_local.py +++ b/tests/test_local.py @@ -65,6 +65,17 @@ def _schema_keyword_paths(schema, keywords): return found +def _local_response(finish_reason): + return { + "choices": [ + { + "message": {"content": "ok"}, + "finish_reason": finish_reason, + } + ] + } + + def test_local_model_prefix_maps_to_provider(): assert get_model_provider(LOCAL_MODEL) == "local" @@ -114,6 +125,33 @@ def test_context_budget_exceeded_classifies_by_reason_code(): ) +@pytest.mark.parametrize( + ("raw", "expected"), + [ + ("stop", "stop"), + ("length", "max_tokens"), + ("max_tokens", "max_tokens"), + ("content_filter", "content_filter"), + ], +) +def test_parse_response_normalizes_known_finish_reasons(raw, expected): + provider = _provider() + + result = provider._parse_response(_local_response(raw)) + + assert result["finish_reason"] == expected + + +@pytest.mark.parametrize("raw", [None, "", "weird", "tool_calls", "function_call"]) +def test_parse_response_fails_closed_on_bad_finish_reasons(raw): + provider = _provider() + + with pytest.raises(provider.LocalProviderError) as exc_info: + provider._parse_response(_local_response(raw)) + + assert exc_info.value.reason_code == "provider_response_invalid" + + def test_cloud_generate_providers_do_not_reference_local_budget(): root = Path(__file__).resolve().parents[1] @@ -839,7 +877,9 @@ def test_run_generate_byo_omits_auth_header_without_credential(monkeypatch): return None def json(self): - return {"choices": [{"message": {"content": "ok"}}]} + return { + "choices": [{"message": {"content": "ok"}, "finish_reason": "stop"}] + } def fake_post(url, **kwargs): captured.update({"url": url, **kwargs}) @@ -854,7 +894,20 @@ def test_run_generate_byo_omits_auth_header_without_credential(monkeypatch): assert "headers" not in captured -def test_generate_schema_files_do_not_declare_bounds(): +def _load_schema(path: str) -> dict: + return json.loads(Path(path).read_text(encoding="utf-8")) + + +def _sense_collection_bounds(schema: dict) -> dict[str, int]: + properties = schema["properties"] + return { + "entities": properties["entities"]["maxItems"], + "facets": properties["facets"]["maxItems"], + "speakers": properties["speakers"]["maxItems"], + } + + +def test_generate_schema_files_declare_only_safe_sense_collection_bounds(): bounded_keys = { "minItems", "maxItems", @@ -863,29 +916,60 @@ def test_generate_schema_files_do_not_declare_bounds(): "minimum", "maximum", } - paths = [ - Path("solstone/talent/sense.schema.json"), - Path("solstone/talent/participation.schema.json"), - Path("solstone/talent/participation_entry.schema.json"), + sense = _load_schema("solstone/talent/sense.schema.json") + + assert _sense_collection_bounds(sense) == { + "entities": 96, + "facets": 16, + "speakers": 16, + } + assert _schema_keyword_paths(sense, {"maxItems"}) == [ + "$/properties/entities/maxItems", + "$/properties/facets/maxItems", + "$/properties/speakers/maxItems", ] - found = {} + assert _schema_keyword_paths(sense, {"pattern", "minLength", "maxLength"}) == [] - def walk(node, keys): - if isinstance(node, dict): - keys.update(bounded_keys & node.keys()) - for value in node.values(): - walk(value, keys) - elif isinstance(node, list): - for item in node: - walk(item, keys) + for path in ( + "solstone/talent/participation.schema.json", + "solstone/talent/participation_entry.schema.json", + ): + assert _schema_keyword_paths(_load_schema(path), bounded_keys) == [] + + +def test_sense_collection_bounds_survive_runtime_and_local_schema_prep(): + from solstone.think.talent import hydrate_runtime_enums + + provider = _provider() + sense = _load_schema("solstone/talent/sense.schema.json") + + hydrated = hydrate_runtime_enums(sense) + prepared = provider._prepare_local_schema(hydrated) + + assert _sense_collection_bounds(hydrated) == { + "entities": 96, + "facets": 16, + "speakers": 16, + } + assert _sense_collection_bounds(prepared) == { + "entities": 96, + "facets": 16, + "speakers": 16, + } + + +def test_local_input_budget_reserve_for_changed_caps(): + from solstone.think.providers.local_budget import compute_input_budget - for path in paths: - keys = set() - walk(json.loads(path.read_text(encoding="utf-8")), keys) - if keys: - found[str(path)] = sorted(keys) + floor = 16384 + capable = 32768 - assert found == {} + assert compute_input_budget(512, floor) - compute_input_budget(1024, floor) == 512 + assert compute_input_budget(2048, floor) - compute_input_budget(4096, floor) == 2048 + assert compute_input_budget(12288, floor) == compute_input_budget(6144, floor) + assert compute_input_budget(12288, floor) == floor - 4096 - 256 + assert compute_input_budget(12288, capable) == capable - 8192 - 256 + assert compute_input_budget(6144, capable) == capable - 6144 - 256 def test_run_generate_byo_network_error_maps_to_unreachable(monkeypatch): diff --git a/tests/test_meeting_schema.py b/tests/test_meeting_schema.py index 393729650..33b8a377d 100644 --- a/tests/test_meeting_schema.py +++ b/tests/test_meeting_schema.py @@ -127,12 +127,15 @@ def test_discover_categories_attaches_meeting_schema(): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_meeting_extract_batch_call_passes_schema(mock_agenerate): - mock_agenerate.return_value = ( - '{"platform":"zoom","participants":[{"name":"Alice","status":"active",' - '"video":true}],"screen_share":null}' - ) + mock_agenerate.return_value = { + "text": ( + '{"platform":"zoom","participants":[{"name":"Alice","status":"active",' + '"video":true}],"screen_share":null}' + ), + "finish_reason": "stop", + } cat_meta = describe_mod.CATEGORIES["meeting"] batch = Batch(max_concurrent=1) diff --git a/tests/test_messaging_schema.py b/tests/test_messaging_schema.py index 157bd8ab7..8c652917c 100644 --- a/tests/test_messaging_schema.py +++ b/tests/test_messaging_schema.py @@ -74,9 +74,12 @@ def test_discover_categories_attaches_messaging_schema(): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_messaging_extract_batch_call_passes_schema(mock_agenerate): - mock_agenerate.return_value = json.dumps(_valid_payload()) + mock_agenerate.return_value = { + "text": json.dumps(_valid_payload()), + "finish_reason": "stop", + } cat_meta = describe_mod.CATEGORIES["messaging"] batch = Batch(max_concurrent=1) diff --git a/tests/test_models.py b/tests/test_models.py index bfeef6cbc..d8da8160b 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -35,7 +35,9 @@ from solstone.think.models import ( TIER_PRO, TYPE_DEFAULTS, IncompleteJSONError, + IncompleteTextError, NoBrainConfiguredError, + ProviderResponseInvalidError, SchemaValidationError, _Family, _find_pricing_fallback, @@ -44,8 +46,10 @@ from solstone.think.models import ( _parse_family_openai, _validate_schema, agenerate, + agenerate_with_result, calc_agent_cost, calc_token_cost, + finish_reason_error, generate, generate_with_result, get_context_registry, @@ -1573,6 +1577,38 @@ class TestValidateSchema: class TestGenerateJsonSchemaPlumbing: + def test_finish_reason_predicate_json_rejects_any_non_stop(self): + error = finish_reason_error( + {"text": "{}", "finish_reason": "content_filter"}, + json_output=True, + ) + + assert isinstance(error, IncompleteJSONError) + + def test_finish_reason_predicate_plain_text_rejects_length(self): + error = finish_reason_error( + {"text": "partial", "finish_reason": "max_tokens"}, + json_output=False, + ) + + assert isinstance(error, IncompleteTextError) + assert error.reason_code == "incomplete_text_length" + + def test_finish_reason_predicate_plain_text_rejects_non_length(self): + error = finish_reason_error( + {"text": "", "finish_reason": "content_filter"}, + json_output=False, + ) + + assert isinstance(error, ProviderResponseInvalidError) + assert error.reason_code == "provider_response_invalid" + + def test_validate_json_response_plain_text_leniency_is_unchanged(self): + models_module._validate_json_response( + {"text": "partial", "finish_reason": "max_tokens"}, + json_output=False, + ) + def test_generate_forces_json_output_with_schema(self): schema = {"type": "object"} provider_module = SimpleNamespace( @@ -1643,6 +1679,34 @@ class TestGenerateJsonSchemaPlumbing: assert result["schema_validation"] == validation + def test_agenerate_with_result_adds_schema_validation(self): + provider_module = SimpleNamespace( + run_agenerate=AsyncMock( + return_value={"text": "{}", "finish_reason": "stop"} + ) + ) + validation = {"valid": True, "errors": []} + + with ( + patch( + "solstone.think.models.resolve_provider", return_value=("fake", "model") + ), + patch( + "solstone.think.providers.get_provider_module", + return_value=provider_module, + ), + patch("solstone.think.models._validate_schema", return_value=validation), + ): + result = asyncio.run( + agenerate_with_result( + "hello", + "test.context", + json_schema={"type": "object"}, + ) + ) + + assert result["schema_validation"] == validation + def test_generate_with_result_returns_failed_schema_validation_without_raising( self, ): diff --git a/tests/test_observe_describe_schema.py b/tests/test_observe_describe_schema.py index dc1d494c9..9020d7480 100644 --- a/tests/test_observe_describe_schema.py +++ b/tests/test_observe_describe_schema.py @@ -138,12 +138,15 @@ def test_describe_schema_accepts_and_rejects_expected_values(): @pytest.mark.asyncio -@patch("solstone.think.batch.agenerate", new_callable=AsyncMock) +@patch("solstone.think.batch.agenerate_with_result", new_callable=AsyncMock) async def test_describe_batch_call_passes_schema(mock_agenerate): - mock_agenerate.return_value = ( - '{"visual_description":"A code editor is visible.","primary":"code",' - '"secondary":"none","overlap":false}' - ) + mock_agenerate.return_value = { + "text": ( + '{"visual_description":"A code editor is visible.","primary":"code",' + '"secondary":"none","overlap":false}' + ), + "finish_reason": "stop", + } batch = Batch(max_concurrent=1) req = batch.create( diff --git a/tests/test_provider_readiness_presenter.py b/tests/test_provider_readiness_presenter.py index e156413e3..5216a0e4b 100644 --- a/tests/test_provider_readiness_presenter.py +++ b/tests/test_provider_readiness_presenter.py @@ -148,6 +148,7 @@ def test_blocking_reason_classification(): "chat_timeout", "network_unreachable", "provider_response_invalid", + "incomplete_text_length", "no_output", "unknown", "ready", diff --git a/tests/test_provider_state.py b/tests/test_provider_state.py index 0a0788904..afdd7db45 100644 --- a/tests/test_provider_state.py +++ b/tests/test_provider_state.py @@ -130,6 +130,7 @@ def test_runtime_reason_codes_are_state_reason_codes(): "provider_response_invalid", "context_window_exceeded", "incomplete_json_length", + "incomplete_text_length", "max_turns_exhausted", "unknown", } diff --git a/tests/test_sense_schema.py b/tests/test_sense_schema.py index 05102f57e..ee0fa9490 100644 --- a/tests/test_sense_schema.py +++ b/tests/test_sense_schema.py @@ -6,6 +6,7 @@ import json from pathlib import Path import frontmatter +import pytest from jsonschema import Draft202012Validator from solstone.think.activities import DEFAULT_ACTIVITIES @@ -39,6 +40,49 @@ def _facet_naming_template() -> str: return FACET_NAMING_PATH.read_text(encoding="utf-8").strip() +def _valid_sense_payload() -> dict: + return { + "density": "active", + "content_type": "coding", + "activity_summary": "Writing tests.", + "entities": [], + "facets": [ + { + "facet": RUNTIME_FACETS_SENTINEL, + "activity": "Writing tests.", + "level": "low", + } + ], + "speculative_facet": "test-planning", + "meeting_detected": False, + "speakers": [], + "recommend": { + "screen_record": False, + "speaker_attribution": False, + }, + "emotional_register": "focused", + } + + +def _sense_entity(index: int) -> dict: + return { + "type": "Person", + "name": f"Person {index}", + "role": "mentioned", + "source": "screen", + "context": "visible in the workspace", + "level": "low", + } + + +def _sense_facet(index: int) -> dict: + return { + "facet": RUNTIME_FACETS_SENTINEL, + "activity": f"Activity {index}", + "level": "low", + } + + def _write_prompt_journal(tmp_path: Path, *, with_facet: bool) -> None: config_dir = tmp_path / "config" config_dir.mkdir(parents=True) @@ -119,27 +163,7 @@ def test_sense_schema_facet_uses_runtime_sentinel_constant(): def test_sense_schema_speculative_facet_nullable_and_required(): schema = json.loads(SENSE_SCHEMA_PATH.read_text(encoding="utf-8")) validator = Draft202012Validator(schema) - valid = { - "density": "active", - "content_type": "coding", - "activity_summary": "Writing tests.", - "entities": [], - "facets": [ - { - "facet": RUNTIME_FACETS_SENTINEL, - "activity": "Writing tests.", - "level": "low", - } - ], - "speculative_facet": "test-planning", - "meeting_detected": False, - "speakers": [], - "recommend": { - "screen_record": False, - "speaker_attribution": False, - }, - "emotional_register": "focused", - } + valid = _valid_sense_payload() assert schema["properties"]["speculative_facet"] == {"type": ["string", "null"]} assert "speculative_facet" in schema["required"] @@ -154,6 +178,36 @@ def test_sense_schema_speculative_facet_nullable_and_required(): assert list(validator.iter_errors(missing)) +@pytest.mark.parametrize( + ("collection", "ok_count", "bad_count", "factory"), + [ + ("entities", 96, 97, _sense_entity), + ("facets", 16, 17, _sense_facet), + ("speakers", 16, 17, lambda index: f"Speaker {index}"), + ], +) +def test_sense_schema_collection_bounds_independently( + collection, + ok_count, + bad_count, + factory, +): + schema = json.loads(SENSE_SCHEMA_PATH.read_text(encoding="utf-8")) + validator = Draft202012Validator(schema) + payload = _valid_sense_payload() + + payload[collection] = [factory(index) for index in range(ok_count)] + assert list(validator.iter_errors(payload)) == [] + + payload[collection] = [factory(index) for index in range(bad_count)] + errors = list(validator.iter_errors(payload)) + + assert any( + error.validator == "maxItems" and list(error.path) == [collection] + for error in errors + ) + + def test_sense_prompt_renders_speculative_facet_instruction_in_steady_state( tmp_path, monkeypatch ): diff --git a/tests/test_talent_fallback.py b/tests/test_talent_fallback.py index 3038492cc..705e8daed 100644 --- a/tests/test_talent_fallback.py +++ b/tests/test_talent_fallback.py @@ -772,6 +772,70 @@ def test_execute_generate_local_length_retry_succeeds( assert events[-1]["retries"] == 1 +def test_execute_generate_local_length_retry_success_writes_once_and_runs_hook_once( + tmp_path, monkeypatch +): + from solstone.think import models, talents + from solstone.think.talents import _execute_generate + + output_path = tmp_path / "out.json" + events = [] + generate_calls = [] + hook_calls = [] + + def mock_generate_with_result(**kwargs): + generate_calls.append(kwargs) + if len(generate_calls) == 1: + raise IncompleteJSONError("length", '{"partial":') + return { + "text": '{"summary": "ok"}', + "usage": {"input_tokens": 1, "output_tokens": 2}, + } + + def post_hook(result, config): + hook_calls.append((result, config["name"])) + return result + + def fail_backup(_agent_type): + raise AssertionError("local retry must not consult cloud backup") + + monkeypatch.setattr( + "solstone.think.talent.key_to_context", lambda _name: "talent.system.default" + ) + monkeypatch.setattr(models, "generate_with_result", mock_generate_with_result) + monkeypatch.setattr(models, "get_backup_provider", fail_backup) + monkeypatch.setattr(talents, "load_post_hook", lambda _config: post_hook) + + asyncio.run( + _execute_generate( + { + "name": "chat", + "provider": "local", + "model": LOCAL_MODEL, + "prompt": "hello", + "health_stale": False, + "output": "json", + "output_path": str(output_path), + "hook": {"post": "test"}, + "json_schema": { + "type": "object", + "additionalProperties": False, + "required": ["summary"], + "properties": {"summary": {"type": "string"}}, + }, + }, + events.append, + ) + ) + + assert len(generate_calls) == 2 + assert hook_calls == [('{"summary": "ok"}', "chat")] + assert output_path.read_text(encoding="utf-8") == '{"summary": "ok"}' + assert [path.name for path in tmp_path.iterdir()] == ["out.json"] + assert events[-1]["event"] == "finish" + assert events[-1]["retries"] == 1 + + def test_main_async_local_length_retry_second_failure_emits_one_error( monkeypatch, capsys ): @@ -830,8 +894,8 @@ def test_main_async_local_length_retry_second_failure_emits_one_error( assert calls["count"] == 2 assert len(error_events) == 1 assert error_events[0]["reason_code"] == "incomplete_json_length" + assert error_events[0]["retries"] == 1 assert finish_events == [] - assert all("retries" not in event for event in events) @pytest.mark.parametrize(