From d81bb318fc9b5f2d800d4ff97450d86b6cdbb2d2 Mon Sep 17 00:00:00 2001 From: Jer Miller Date: Sat, 11 Jul 2026 07:32:17 -0600 Subject: [PATCH] fix(think): contain local output tails Cap the Qwen-heavy describe and Sense output paths so truncated local runs have less tail room to leak partial output, while leaving the category loader default unchanged. Add explicit Sense collection bounds that survive runtime enum hydration and local GBNF schema prep instead of falling back to the global array cap. Normalize local finish reasons and make the Batch boundary reject any present non-stop finish while preserving hard schema-invalid JSON failures. This keeps truncated, malformed, timed-out, or otherwise bad local responses from landing as normal-looking journal artifacts at frame categorization, browsing extraction, and Segment Sense boundaries. Carry local JSON length retry counts into the terminal error when the retry itself fails. Co-Authored-By: OpenAI Codex --- solstone/apps/timeline/tests/conftest.py | 9 +- solstone/convey/provider_readiness.py | 9 ++ solstone/convey/static/chat_reasons.js | 4 + solstone/observe/categories/browsing.md | 3 +- solstone/observe/describe.py | 2 +- solstone/talent/sense.md | 2 +- solstone/talent/sense.schema.json | 3 + solstone/think/batch.py | 29 +++- solstone/think/models.py | 131 +++++++++++++++- solstone/think/providers/local.py | 30 +++- solstone/think/providers/shared.py | 1 + solstone/think/talents.py | 35 +++-- tests/baselines/api/stats/stats.json | 2 +- tests/test_bad_media_corpus.py | 43 +++++- tests/test_batch.py | 164 ++++++++++++++++----- tests/test_calendar_schema.py | 7 +- tests/test_chat_reasons.py | 1 + tests/test_describe_promote.py | 100 +++++++++++++ tests/test_generate_talents.py | 40 +++-- tests/test_local.py | 126 +++++++++++++--- tests/test_meeting_schema.py | 13 +- tests/test_messaging_schema.py | 7 +- tests/test_models.py | 64 ++++++++ tests/test_observe_describe_schema.py | 13 +- tests/test_provider_readiness_presenter.py | 1 + tests/test_provider_state.py | 1 + tests/test_sense_schema.py | 96 +++++++++--- tests/test_talent_fallback.py | 66 ++++++++- 28 files changed, 856 insertions(+), 146 deletions(-) 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( -- 2.51.2