diff --git a/examples/structured_output_model.py b/examples/structured_output_model.py index 10328df..d847001 100644 --- a/examples/structured_output_model.py +++ b/examples/structured_output_model.py @@ -7,7 +7,7 @@ from pydantic import BaseModel -from droid_sdk import run +from droid_sdk import RunSuccess, SessionConfig, run class Finding(BaseModel): @@ -22,11 +22,29 @@ class Review(BaseModel): async def main() -> None: result = await run( - "Return a short repository summary and zero or more findings.", + ( + 'Return summary exactly "Structured output works." and one finding ' + 'with severity "low" and message exactly "Example complete."' + ), output=Review, - timeout=180, + timeout=60, + config=SessionConfig( + disable_builtin_skills=True, + restrict_tools=(), + ), + ) + assert isinstance(result, RunSuccess), ( + result.error.message if result.error else result.subtype + ) + assert result.output == Review( + summary="Structured output works.", + findings=[ + Finding( + severity="low", + message="Example complete.", + ) + ], ) - assert result.output is not None, result.output_validation_error print(result.output.summary) diff --git a/src/droid_sdk/_high_level/output.py b/src/droid_sdk/_high_level/output.py index fa5d883..abb3c84 100644 --- a/src/droid_sdk/_high_level/output.py +++ b/src/droid_sdk/_high_level/output.py @@ -47,11 +47,17 @@ def adapt(self, raw: object | None) -> OutputAdaptation[T_co]: return OutputAdaptation(None, None, None) if self._model is not None: + structured_output = _raw_object(raw) + validation_input = ( + thaw_json(structured_output) + if structured_output is not None + else raw + ) try: - value = self._model.model_validate(raw) + value = self._model.model_validate(validation_input) except ValidationError as exc: - return OutputAdaptation(None, _raw_object(raw), exc) - return OutputAdaptation(cast("T_co", value), _raw_object(raw), None) + return OutputAdaptation(None, structured_output, exc) + return OutputAdaptation(cast("T_co", value), structured_output, None) raw_value = _validate_json_object(raw) return OutputAdaptation( diff --git a/tests/test_v5_streaming_core.py b/tests/test_v5_streaming_core.py index 35cd57f..d3730a7 100644 --- a/tests/test_v5_streaming_core.py +++ b/tests/test_v5_streaming_core.py @@ -6,7 +6,7 @@ from typing import Any, cast import pytest -from pydantic import BaseModel +from pydantic import BaseModel, ConfigDict from droid_sdk import ( AssistantMessage, @@ -780,6 +780,20 @@ class NumericOutput(BaseModel): value: float +class StrictFinding(BaseModel): + model_config = ConfigDict(strict=True, extra="forbid") + + severity: str + message: str + + +class StrictReview(BaseModel): + model_config = ConfigDict(strict=True, extra="forbid") + + summary: str + findings: list[StrictFinding] + + @pytest.mark.asyncio async def test_structured_output_raw_adapted_fallback_and_validation_error() -> None: valid = RunStream[Review]( @@ -845,6 +859,50 @@ async def test_structured_output_raw_adapted_fallback_and_validation_error() -> assert raw.result.output == {"summary": "fallback"} +@pytest.mark.asyncio +async def test_strict_nested_output_validates_frozen_stream_data() -> None: + stream = RunStream[StrictReview]( + expected_turn_id="turn", + session_id="session", + output_adapter=prepare_output_adapter(StrictReview), + ) + raw = { + "summary": "ok", + "findings": [ + { + "severity": "high", + "message": "Validate ordinary JSON containers", + } + ], + } + stream.feed_notification( + { + "type": "structured_output", + "messageId": "assistant", + "structuredOutput": raw, + } + ) + stream.feed_notification(_complete()) + + await _collect(stream) + + assert isinstance(stream.result, RunSuccess) + assert stream.result.output == StrictReview.model_validate(raw) + structured_output = stream.result.structured_output + assert structured_output is not None + assert structured_output == { + "summary": "ok", + "findings": ( + { + "severity": "high", + "message": "Validate ordinary JSON containers", + }, + ), + } + assert isinstance(structured_output["findings"], tuple) + assert stream.result.output_validation_error is None + + @pytest.mark.asyncio async def test_requested_output_that_never_arrives_is_a_failure() -> None: missing = RunStream[Mapping[str, object]](