diff --git a/pyproject.toml b/pyproject.toml index fbfe3ad24..deffd1c0c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -51,6 +51,7 @@ dependencies = [ "openai-agents>=0.1.0", "anthropic", "httpx", + "jsonschema>=4.26,<5", "genai-prices", "pypdf", "pdf2image", diff --git a/tests/test_anthropic.py b/tests/test_anthropic.py index e7d519b82..7249313c0 100644 --- a/tests/test_anthropic.py +++ b/tests/test_anthropic.py @@ -7,6 +7,7 @@ import json import sys import types from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock from think.models import CLAUDE_SONNET_4 @@ -93,8 +94,12 @@ def _setup_anthropic_stub( else: self.messages = DummyMessages() + class DummyBadRequestError(Exception): + pass + anthropic_stub.Anthropic = DummyClient anthropic_stub.AsyncAnthropic = DummyClient # Add async version + anthropic_stub.BadRequestError = DummyBadRequestError # Add types to the types module anthropic_types_stub.MessageParam = dict @@ -401,3 +406,122 @@ def test_claude_outfile_error(monkeypatch, tmp_path, capsys): events = [json.loads(line) for line in out_lines if line] if events: assert any(e["event"] == "error" for e in events) + + +class TestRunGenerateJsonSchema: + def test_no_schema_keeps_prompt_append(self, monkeypatch): + provider = importlib.reload( + importlib.import_module("think.providers.anthropic") + ) + mock_client = MagicMock() + mock_response = MagicMock() + mock_response.content = [SimpleNamespace(type="text", text="{}")] + mock_response.usage = None + mock_response.stop_reason = "end_turn" + mock_client.messages.create.return_value = mock_response + monkeypatch.setattr(provider, "_get_anthropic_client", lambda: mock_client) + + provider.run_generate( + "hello", + json_output=True, + system_instruction="base", + ) + + call_kwargs = mock_client.messages.create.call_args.kwargs + assert call_kwargs["system"].endswith( + "Respond with valid JSON only. No explanation or markdown." + ) + + def test_with_schema_uses_output_config(self, monkeypatch): + provider = importlib.reload( + importlib.import_module("think.providers.anthropic") + ) + mock_client = MagicMock() + mock_response = MagicMock() + mock_response.content = [SimpleNamespace(type="text", text="{}")] + mock_response.usage = None + mock_response.stop_reason = "end_turn" + mock_client.messages.create.return_value = mock_response + monkeypatch.setattr(provider, "_get_anthropic_client", lambda: mock_client) + schema = {"type": "object"} + + provider.run_generate( + "hello", + system_instruction="base", + json_schema=schema, + ) + + call_kwargs = mock_client.messages.create.call_args.kwargs + assert call_kwargs["output_config"] == { + "format": {"type": "json_schema", "schema": schema} + } + assert call_kwargs["system"] == "base" + + def test_fallback_on_bad_request(self, monkeypatch): + provider = importlib.reload( + importlib.import_module("think.providers.anthropic") + ) + mock_client = MagicMock() + + class DummyBadRequestError(Exception): + pass + + fallback_response = MagicMock() + fallback_response.content = [ + SimpleNamespace(type="tool_use", input={"key": "value"}), + ] + fallback_response.usage = None + fallback_response.stop_reason = "end_turn" + mock_client.messages.create.side_effect = [ + DummyBadRequestError("bad schema"), + fallback_response, + ] + + monkeypatch.setattr(provider, "BadRequestError", DummyBadRequestError) + monkeypatch.setattr(provider, "_get_anthropic_client", lambda: mock_client) + schema = {"type": "object"} + + result = provider.run_generate("hello", json_schema=schema) + + assert mock_client.messages.create.call_count == 2 + retry_kwargs = mock_client.messages.create.call_args_list[1].kwargs + assert retry_kwargs["tools"] == [ + { + "name": "response", + "description": "Generate the requested JSON response.", + "input_schema": schema, + } + ] + assert retry_kwargs["tool_choice"] == {"type": "tool", "name": "response"} + assert "output_config" not in retry_kwargs + assert result["text"] == json.dumps({"key": "value"}) + + def test_async_with_schema_uses_output_config(self, monkeypatch): + provider = importlib.reload( + importlib.import_module("think.providers.anthropic") + ) + mock_client = MagicMock() + mock_client.messages.create = AsyncMock() + mock_response = MagicMock() + mock_response.content = [SimpleNamespace(type="text", text="{}")] + mock_response.usage = None + mock_response.stop_reason = "end_turn" + mock_client.messages.create.return_value = mock_response + monkeypatch.setattr( + provider, "_get_async_anthropic_client", lambda: mock_client + ) + schema = {"type": "object"} + + asyncio.run( + provider.run_agenerate( + "hello", + system_instruction="base", + json_schema=schema, + ) + ) + + call_kwargs = mock_client.messages.create.call_args.kwargs + assert call_kwargs["output_config"] == { + "format": {"type": "json_schema", "schema": schema} + } + assert call_kwargs["system"] == "base" diff --git a/tests/test_google.py b/tests/test_google.py index 6e7c081e1..62ebc8ffd 100644 --- a/tests/test_google.py +++ b/tests/test_google.py @@ -6,7 +6,7 @@ import importlib import json import sys from types import SimpleNamespace -from unittest.mock import AsyncMock +from unittest.mock import AsyncMock, MagicMock from tests.conftest import setup_google_genai_stub from think.models import GEMINI_FLASH @@ -257,3 +257,106 @@ def test_format_completion_message_none(): """Test message when finish_reason is None.""" msg = _format_completion_message(None, had_tool_calls=False) assert msg == "Completed (unknown)." + + +class TestRunGenerateJsonSchema: + def test_no_schema_kwargs_unchanged(self, monkeypatch): + setup_google_genai_stub(monkeypatch, with_thinking=False) + sys.modules.pop("think.providers.google", None) + provider = importlib.reload(importlib.import_module("think.providers.google")) + + mock_client = MagicMock() + mock_client.models.generate_content.return_value = SimpleNamespace( + text="[]", + candidates=[], + usage_metadata=None, + ) + monkeypatch.setattr( + provider, "get_or_create_client", lambda _client=None: mock_client + ) + + provider.run_generate("hello", model=GEMINI_FLASH, json_output=True) + + config = mock_client.models.generate_content.call_args.kwargs["config"] + assert config.response_mime_type == "application/json" + assert getattr(config, "response_json_schema", None) is None + + def test_with_schema_adds_json_schema(self, monkeypatch): + setup_google_genai_stub(monkeypatch, with_thinking=False) + sys.modules.pop("think.providers.google", None) + provider = importlib.reload(importlib.import_module("think.providers.google")) + + schema = {"type": "object"} + mock_client = MagicMock() + mock_client.models.generate_content.return_value = SimpleNamespace( + text="[]", + candidates=[], + usage_metadata=None, + ) + monkeypatch.setattr( + provider, "get_or_create_client", lambda _client=None: mock_client + ) + + provider.run_generate( + "hello", model=GEMINI_FLASH, json_output=True, json_schema=schema + ) + + config = mock_client.models.generate_content.call_args.kwargs["config"] + assert config.response_mime_type == "application/json" + assert config.response_json_schema == schema + + def test_async_no_schema_kwargs_unchanged(self, monkeypatch): + setup_google_genai_stub(monkeypatch, with_thinking=False) + sys.modules.pop("think.providers.google", None) + provider = importlib.reload(importlib.import_module("think.providers.google")) + + mock_client = MagicMock() + mock_client.aio.models.generate_content = AsyncMock( + return_value=SimpleNamespace( + text="[]", + candidates=[], + usage_metadata=None, + ) + ) + monkeypatch.setattr( + provider, "get_or_create_client", lambda _client=None: mock_client + ) + + asyncio.run( + provider.run_agenerate("hello", model=GEMINI_FLASH, json_output=True) + ) + + config = mock_client.aio.models.generate_content.call_args.kwargs["config"] + assert config.response_mime_type == "application/json" + assert getattr(config, "response_json_schema", None) is None + + def test_async_with_schema_adds_json_schema(self, monkeypatch): + setup_google_genai_stub(monkeypatch, with_thinking=False) + sys.modules.pop("think.providers.google", None) + provider = importlib.reload(importlib.import_module("think.providers.google")) + + schema = {"type": "object"} + mock_client = MagicMock() + mock_client.aio.models.generate_content = AsyncMock( + return_value=SimpleNamespace( + text="[]", + candidates=[], + usage_metadata=None, + ) + ) + monkeypatch.setattr( + provider, "get_or_create_client", lambda _client=None: mock_client + ) + + asyncio.run( + provider.run_agenerate( + "hello", + model=GEMINI_FLASH, + json_output=True, + json_schema=schema, + ) + ) + + config = mock_client.aio.models.generate_content.call_args.kwargs["config"] + assert config.response_mime_type == "application/json" + assert config.response_json_schema == schema diff --git a/tests/test_models.py b/tests/test_models.py index 88015b6d0..54ee4a225 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -3,6 +3,11 @@ """Tests for think.models module.""" +import asyncio +import logging +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch + import pytest from think.models import ( @@ -21,7 +26,12 @@ from think.models import ( TIER_LITE, TIER_PRO, TYPE_DEFAULTS, + IncompleteJSONError, + _validate_schema, + agenerate, calc_token_cost, + generate, + generate_with_result, get_context_registry, get_usage_cost, iter_token_log, @@ -738,3 +748,244 @@ def test_log_token_usage_passes_through_cache_creation_tokens(tmp_path, monkeypa entry = json.loads(log_file.read_text().strip()) assert entry["usage"]["cache_creation_tokens"] == 2000 assert entry["usage"]["cached_tokens"] == 3000 + + +class TestValidateSchema: + def test_valid_instance(self): + schema = { + "type": "object", + "properties": {"field": {"type": "string"}}, + "required": ["field"], + } + + result = _validate_schema('{"field": "ok"}', schema) + + assert result == {"valid": True, "errors": []} + + def test_schema_violation_type(self): + schema = { + "type": "object", + "properties": {"field": {"type": "integer"}}, + "required": ["field"], + } + + result = _validate_schema('{"field": "ok"}', schema) + + assert result["valid"] is False + assert len(result["errors"]) == 1 + assert result["errors"][0]["path"] == "/field" + assert result["errors"][0]["constraint"] == "type" + assert result["errors"][0]["message"] + + def test_schema_violation_required(self): + schema = {"type": "object", "required": ["field"]} + + result = _validate_schema("{}", schema) + + assert result["valid"] is False + assert len(result["errors"]) == 1 + assert result["errors"][0]["path"] == "" + assert result["errors"][0]["constraint"] == "required" + + def test_multiple_violations(self): + schema = { + "type": "object", + "properties": { + "field": {"type": "integer"}, + "other": {"type": "string"}, + }, + "required": ["field", "other"], + } + + result = _validate_schema('{"field": "bad"}', schema) + + assert result["valid"] is False + assert len(result["errors"]) == 2 + assert {error["constraint"] for error in result["errors"]} == { + "required", + "type", + } + + def test_parse_failure(self): + schema = {"type": "object"} + + result = _validate_schema("{", schema) + + assert result["valid"] is False + assert result["errors"] == [ + { + "path": "", + "constraint": "json_parse", + "message": result["errors"][0]["message"], + } + ] + + def test_json_pointer_escape(self): + schema = { + "type": "object", + "properties": { + "a/b": { + "type": "object", + "properties": {"c~d": {"type": "integer"}}, + } + }, + } + + result = _validate_schema('{"a/b": {"c~d": "bad"}}', schema) + + assert result["valid"] is False + assert result["errors"][0]["path"] == "/a~1b/c~0d" + + def test_warning_logged_for_violations(self, caplog): + schema = { + "type": "object", + "properties": {"field": {"type": "integer"}}, + } + + with caplog.at_level(logging.WARNING): + _validate_schema('{"field": "bad"}', schema) + + assert any( + record.levelno == logging.WARNING + and "schema_validation:" in record.getMessage() + for record in caplog.records + ) + + def test_invalid_schema_does_not_raise(self): + result = _validate_schema('{"field": "ok"}', {"type": "not-a-real-type"}) + + assert result["valid"] is False + assert len(result["errors"]) == 1 + assert result["errors"][0]["constraint"] == "schema_validation" + + +class TestGenerateJsonSchemaPlumbing: + def test_generate_forces_json_output_with_schema(self): + schema = {"type": "object"} + provider_module = SimpleNamespace( + run_generate=MagicMock(return_value={"text": "{}", "finish_reason": "stop"}) + ) + + with ( + patch("think.models.resolve_provider", return_value=("fake", "model")), + patch("think.providers.get_provider_module", return_value=provider_module), + ): + result = generate("hello", "test.context", json_schema=schema) + + assert result == "{}" + call_kwargs = provider_module.run_generate.call_args.kwargs + assert call_kwargs["json_output"] is True + assert call_kwargs["json_schema"] == schema + + def test_agenerate_forces_json_output_with_schema(self): + schema = {"type": "object"} + provider_module = SimpleNamespace( + run_agenerate=AsyncMock( + return_value={"text": "{}", "finish_reason": "stop"} + ) + ) + + with ( + patch("think.models.resolve_provider", return_value=("fake", "model")), + patch("think.providers.get_provider_module", return_value=provider_module), + ): + result = asyncio.run(agenerate("hello", "test.context", json_schema=schema)) + + assert result == "{}" + call_kwargs = provider_module.run_agenerate.call_args.kwargs + assert call_kwargs["json_output"] is True + assert call_kwargs["json_schema"] == schema + + def test_generate_with_result_adds_schema_validation(self): + provider_module = SimpleNamespace( + run_generate=MagicMock(return_value={"text": "{}", "finish_reason": "stop"}) + ) + validation = {"valid": True, "errors": []} + + with ( + patch("think.models.resolve_provider", return_value=("fake", "model")), + patch("think.providers.get_provider_module", return_value=provider_module), + patch("think.models._validate_schema", return_value=validation), + ): + result = generate_with_result( + "hello", + "test.context", + json_schema={"type": "object"}, + ) + + assert result["schema_validation"] == validation + + def test_generate_with_result_omits_schema_validation_without_schema(self): + provider_module = SimpleNamespace( + run_generate=MagicMock(return_value={"text": "{}", "finish_reason": "stop"}) + ) + + with ( + patch("think.models.resolve_provider", return_value=("fake", "model")), + patch("think.providers.get_provider_module", return_value=provider_module), + patch("think.models._validate_schema") as mock_validate_schema, + ): + result = generate_with_result("hello", "test.context") + + assert "schema_validation" not in result + mock_validate_schema.assert_not_called() + + def test_generate_and_agenerate_do_not_surface_schema_validation(self): + sync_provider = SimpleNamespace( + run_generate=MagicMock(return_value={"text": "{}", "finish_reason": "stop"}) + ) + async_provider = SimpleNamespace( + run_agenerate=AsyncMock( + return_value={"text": "{}", "finish_reason": "stop"} + ) + ) + + with ( + patch("think.models.resolve_provider", return_value=("fake", "model")), + patch("think.providers.get_provider_module", return_value=sync_provider), + patch( + "think.models._validate_schema", + return_value={"valid": True, "errors": []}, + ) as mock_validate_schema, + ): + sync_result = generate( + "hello", "test.context", json_schema={"type": "object"} + ) + + with ( + patch("think.models.resolve_provider", return_value=("fake", "model")), + patch("think.providers.get_provider_module", return_value=async_provider), + patch( + "think.models._validate_schema", + return_value={"valid": True, "errors": []}, + ) as mock_async_validate, + ): + async_result = asyncio.run( + agenerate("hello", "test.context", json_schema={"type": "object"}) + ) + + assert sync_result == "{}" + assert async_result == "{}" + mock_validate_schema.assert_called_once() + mock_async_validate.assert_called_once() + + def test_truncation_raises_before_schema_validation(self): + provider_module = SimpleNamespace( + run_generate=MagicMock(return_value={"text": "{}", "finish_reason": "stop"}) + ) + + with ( + patch("think.models.resolve_provider", return_value=("fake", "model")), + patch("think.providers.get_provider_module", return_value=provider_module), + patch( + "think.models._validate_json_response", + side_effect=IncompleteJSONError("max_tokens", "{}"), + ), + patch("think.models._validate_schema") as mock_validate_schema, + ): + with pytest.raises(IncompleteJSONError): + generate_with_result( + "hello", "test.context", json_schema={"type": "object"} + ) + + mock_validate_schema.assert_not_called() diff --git a/tests/test_ollama.py b/tests/test_ollama.py index 196873989..d5e13f4b9 100644 --- a/tests/test_ollama.py +++ b/tests/test_ollama.py @@ -141,6 +141,20 @@ class TestBuildRequestBody: ) assert body["format"] == "json" + def test_json_schema_dict(self): + provider = _ollama_provider() + schema = {"type": "object"} + body = provider._build_request_body( + "m", + [{"role": "user", "content": "hi"}], + 0.3, + 1024, + True, + None, + schema, + ) + assert body["format"] == schema + def test_no_json_output(self): provider = _ollama_provider() body = provider._build_request_body( @@ -421,6 +435,26 @@ class TestRunGenerate: body = call_kwargs.kwargs["json"] assert body["format"] == "json" + def test_json_schema_dict(self): + provider = _ollama_provider() + mock_response = MagicMock() + mock_response.json.return_value = _make_ollama_response( + content='{"key": "value"}' + ) + mock_response.raise_for_status = MagicMock() + schema = {"type": "object"} + + with patch.object(provider, "_get_client") as mock_get: + mock_client = MagicMock() + mock_client.post.return_value = mock_response + mock_get.return_value = mock_client + + provider.run_generate("hello", model=OLLAMA_FLASH, json_schema=schema) + + call_kwargs = mock_client.post.call_args + body = call_kwargs.kwargs["json"] + assert body["format"] == schema + def test_system_instruction(self): provider = _ollama_provider() mock_response = MagicMock() @@ -481,6 +515,28 @@ class TestRunAgenerate: assert result["text"] == "Hello!" assert result["finish_reason"] == "stop" + def test_async_json_schema_dict(self): + provider = _ollama_provider() + mock_response = MagicMock() + mock_response.json.return_value = _make_ollama_response( + content='{"key": "value"}' + ) + mock_response.raise_for_status = MagicMock() + schema = {"type": "object"} + + with patch.object(provider, "_get_async_client") as mock_get: + mock_client = MagicMock() + mock_client.post = AsyncMock(return_value=mock_response) + mock_get.return_value = mock_client + + asyncio.run( + provider.run_agenerate("hello", model=OLLAMA_FLASH, json_schema=schema) + ) + + call_kwargs = mock_client.post.call_args + body = call_kwargs.kwargs["json"] + assert body["format"] == schema + # --------------------------------------------------------------------------- # _translate_opencode diff --git a/tests/test_openai.py b/tests/test_openai.py index 274474860..0fe998cb2 100644 --- a/tests/test_openai.py +++ b/tests/test_openai.py @@ -726,6 +726,107 @@ class TestRunGenerate: called_kwargs = mock_client.responses.create.call_args.kwargs assert called_kwargs["text"] == {"format": {"type": "json_object"}} + def test_no_schema_format_unchanged(self): + provider = _openai_provider() + mock_client = MagicMock() + mock_client.responses.create = MagicMock() + mock_response = MagicMock() + mock_response.output_text = "Hello" + mock_response.status = "completed" + mock_response.incomplete_details = None + mock_response.usage = None + mock_response.output = [] + mock_client.responses.create.return_value = mock_response + + with patch( + "think.providers.openai._get_openai_client", return_value=mock_client + ): + provider.run_generate( + "hello", + model="gpt-5.2", + json_output=True, + json_schema=None, + ) + + called_kwargs = mock_client.responses.create.call_args.kwargs + assert called_kwargs["text"] == {"format": {"type": "json_object"}} + + def test_with_schema_format_shape(self): + provider = _openai_provider() + mock_client = MagicMock() + mock_client.responses.create = MagicMock() + mock_response = MagicMock() + mock_response.output_text = "{}" + mock_response.status = "completed" + mock_response.incomplete_details = None + mock_response.usage = None + mock_response.output = [] + mock_client.responses.create.return_value = mock_response + schema = {"type": "object"} + + with patch( + "think.providers.openai._get_openai_client", return_value=mock_client + ): + provider.run_generate("hello", model="gpt-5.2", json_schema=schema) + + called_kwargs = mock_client.responses.create.call_args.kwargs + assert called_kwargs["text"] == { + "format": { + "type": "json_schema", + "name": "response", + "schema": schema, + "strict": True, + } + } + + def test_schema_title_becomes_name(self): + provider = _openai_provider() + mock_client = MagicMock() + mock_client.responses.create = MagicMock() + mock_response = MagicMock() + mock_response.output_text = "{}" + mock_response.status = "completed" + mock_response.incomplete_details = None + mock_response.usage = None + mock_response.output = [] + mock_client.responses.create.return_value = mock_response + + with patch( + "think.providers.openai._get_openai_client", return_value=mock_client + ): + provider.run_generate( + "hello", + model="gpt-5.2", + json_schema={"title": "MyThing", "type": "object"}, + ) + + called_kwargs = mock_client.responses.create.call_args.kwargs + assert called_kwargs["text"]["format"]["name"] == "MyThing" + + def test_schema_bad_title_falls_back(self): + provider = _openai_provider() + mock_client = MagicMock() + mock_client.responses.create = MagicMock() + mock_response = MagicMock() + mock_response.output_text = "{}" + mock_response.status = "completed" + mock_response.incomplete_details = None + mock_response.usage = None + mock_response.output = [] + mock_client.responses.create.return_value = mock_response + + with patch( + "think.providers.openai._get_openai_client", return_value=mock_client + ): + provider.run_generate( + "hello", + model="gpt-5.2", + json_schema={"title": "bad name with spaces", "type": "object"}, + ) + + called_kwargs = mock_client.responses.create.call_args.kwargs + assert called_kwargs["text"]["format"]["name"] == "response" + def test_with_system_instruction(self): provider = _openai_provider() mock_client = MagicMock() @@ -821,3 +922,112 @@ class TestRunAgenerate: result = asyncio.run(provider.run_agenerate("hello", model="gpt-5.2")) assert result["thinking"] == [{"summary": "Let me think..."}] + + def test_no_schema_format_unchanged(self): + provider = _openai_provider() + mock_client = MagicMock() + mock_client.responses.create = AsyncMock() + mock_response = MagicMock() + mock_response.output_text = "Hello" + mock_response.status = "completed" + mock_response.incomplete_details = None + mock_response.usage = None + mock_response.output = [] + mock_client.responses.create.return_value = mock_response + + with patch( + "think.providers.openai._get_async_openai_client", return_value=mock_client + ): + asyncio.run( + provider.run_agenerate( + "hello", + model="gpt-5.2", + json_output=True, + json_schema=None, + ) + ) + + called_kwargs = mock_client.responses.create.call_args.kwargs + assert called_kwargs["text"] == {"format": {"type": "json_object"}} + + def test_with_schema_format_shape(self): + provider = _openai_provider() + mock_client = MagicMock() + mock_client.responses.create = AsyncMock() + mock_response = MagicMock() + mock_response.output_text = "{}" + mock_response.status = "completed" + mock_response.incomplete_details = None + mock_response.usage = None + mock_response.output = [] + mock_client.responses.create.return_value = mock_response + schema = {"type": "object"} + + with patch( + "think.providers.openai._get_async_openai_client", return_value=mock_client + ): + asyncio.run( + provider.run_agenerate("hello", model="gpt-5.2", json_schema=schema) + ) + + called_kwargs = mock_client.responses.create.call_args.kwargs + assert called_kwargs["text"] == { + "format": { + "type": "json_schema", + "name": "response", + "schema": schema, + "strict": True, + } + } + + def test_schema_title_becomes_name(self): + provider = _openai_provider() + mock_client = MagicMock() + mock_client.responses.create = AsyncMock() + mock_response = MagicMock() + mock_response.output_text = "{}" + mock_response.status = "completed" + mock_response.incomplete_details = None + mock_response.usage = None + mock_response.output = [] + mock_client.responses.create.return_value = mock_response + + with patch( + "think.providers.openai._get_async_openai_client", return_value=mock_client + ): + asyncio.run( + provider.run_agenerate( + "hello", + model="gpt-5.2", + json_schema={"title": "MyThing", "type": "object"}, + ) + ) + + called_kwargs = mock_client.responses.create.call_args.kwargs + assert called_kwargs["text"]["format"]["name"] == "MyThing" + + def test_schema_bad_title_falls_back(self): + provider = _openai_provider() + mock_client = MagicMock() + mock_client.responses.create = AsyncMock() + mock_response = MagicMock() + mock_response.output_text = "{}" + mock_response.status = "completed" + mock_response.incomplete_details = None + mock_response.usage = None + mock_response.output = [] + mock_client.responses.create.return_value = mock_response + + with patch( + "think.providers.openai._get_async_openai_client", return_value=mock_client + ): + asyncio.run( + provider.run_agenerate( + "hello", + model="gpt-5.2", + json_schema={"title": "bad name with spaces", "type": "object"}, + ) + ) + + called_kwargs = mock_client.responses.create.call_args.kwargs + assert called_kwargs["text"]["format"]["name"] == "response" diff --git a/think/models.py b/think/models.py index e770e511e..0b764cdd9 100644 --- a/think/models.py +++ b/think/models.py @@ -13,9 +13,12 @@ from pathlib import Path from typing import Any, Dict, List, Optional, Union import frontmatter +from jsonschema import Draft202012Validator from think.utils import get_config, get_journal +logger = logging.getLogger(__name__) + # --------------------------------------------------------------------------- # Tier constants # --------------------------------------------------------------------------- @@ -908,6 +911,83 @@ def _validate_json_response(result: Dict[str, Any], json_output: bool) -> None: ) +def _validate_schema(text: str, schema: dict) -> dict: + """Validate JSON text against a JSON Schema and log any violations.""" + + def truncate_repr(value: Any) -> str: + value_repr = repr(value) + if len(value_repr) <= 80: + return value_repr + return value_repr[:77] + "..." + + def build_pointer(path: Any) -> str: + segments = list(path) + if not segments: + return "" + escaped_segments = [] + for segment in segments: + escaped = str(segment).replace("~", "~0").replace("/", "~1") + escaped_segments.append(escaped) + return "/" + "/".join(escaped_segments) + + try: + parsed = json.loads(text) + except ValueError as exc: + error = { + "path": "", + "constraint": "json_parse", + "message": str(exc), + } + logger.warning( + "schema_validation: %s: %s: %s (value=%s)", + "", + "json_parse", + str(exc), + truncate_repr(text), + ) + return {"valid": False, "errors": [error]} + + errors = [] + try: + validator = Draft202012Validator(schema) + validation_errors = list(validator.iter_errors(parsed)) + except Exception as exc: + error = { + "path": "", + "constraint": "schema_validation", + "message": str(exc), + } + logger.warning( + "schema_validation: %s: %s: %s (value=%s)", + "", + "schema_validation", + str(exc), + truncate_repr(parsed), + ) + return {"valid": False, "errors": [error]} + + for error in validation_errors: + path = build_pointer(error.absolute_path) + constraint = str(error.validator) + message = error.message + errors.append( + { + "path": path, + "constraint": constraint, + "message": message, + } + ) + logger.warning( + "schema_validation: %s: %s: %s (value=%s)", + path, + constraint, + message, + truncate_repr(error.instance), + ) + + return {"valid": len(errors) == 0, "errors": errors} + + def generate( contents: Union[str, List[Any]], context: str, @@ -915,6 +995,8 @@ def generate( 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, @@ -939,6 +1021,10 @@ def generate( System instruction for the model. json_output : bool Whether to request JSON response format. + json_schema : dict, optional + JSON Schema to request structured output from the provider. When supplied, + this forces json_output=True and runs advisory local validation on the + returned text after truncation checks. thinking_budget : int, optional Token budget for model thinking (ignored by providers that don't support it). timeout_s : float, optional @@ -960,6 +1046,9 @@ def generate( """ from think.providers import get_provider_module + if json_schema is not None: + json_output = True + # Allow model override via kwargs (used by callers with explicit model selection) model_override = kwargs.pop("model", None) @@ -978,6 +1067,7 @@ def generate( max_output_tokens=max_output_tokens, system_instruction=system_instruction, json_output=json_output, + json_schema=json_schema, thinking_budget=thinking_budget, timeout_s=timeout_s, **kwargs, @@ -996,6 +1086,9 @@ def generate( # Validate JSON output if requested _validate_json_response(result, json_output) + if json_schema is not None: + _validate_schema(result["text"], json_schema) + return result["text"] @@ -1099,6 +1192,8 @@ def generate_with_result( 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, @@ -1109,13 +1204,43 @@ def generate_with_result( just the text. Used by cortex-managed generators that need usage data for event emission. + Parameters + ---------- + contents : str or List + The content to send to the model. + context : str + Context string for routing and token logging. + temperature : float + Temperature for generation (default: 0.3). + max_output_tokens : int + Maximum tokens for the model's response output. + system_instruction : str, optional + System instruction for the model. + json_output : bool + Whether to request JSON response format. + json_schema : dict, optional + JSON Schema to request structured output from the provider. When supplied, + this forces json_output=True and runs advisory local validation on the + returned text after truncation checks. + thinking_budget : int, optional + Token budget for model thinking (ignored by providers that don't support it). + timeout_s : float, optional + Request timeout in seconds. + **kwargs + Additional provider-specific options passed through to the backend. + Returns ------- dict - GenerateResult with: text, usage, finish_reason, thinking. + GenerateResult with: text, usage, finish_reason, thinking, and + schema_validation when json_schema is supplied. Validation is advisory + and runs after truncation checks succeed. """ from 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) @@ -1136,6 +1261,7 @@ def generate_with_result( max_output_tokens=max_output_tokens, system_instruction=system_instruction, json_output=json_output, + json_schema=json_schema, thinking_budget=thinking_budget, timeout_s=timeout_s, **kwargs, @@ -1154,6 +1280,9 @@ def generate_with_result( # Validate JSON output if requested _validate_json_response(result, json_output) + if json_schema is not None: + result["schema_validation"] = _validate_schema(result["text"], json_schema) + return result @@ -1164,6 +1293,8 @@ async def agenerate( 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, @@ -1188,6 +1319,10 @@ async def agenerate( System instruction for the model. json_output : bool Whether to request JSON response format. + json_schema : dict, optional + JSON Schema to request structured output from the provider. When supplied, + this forces json_output=True and runs advisory local validation on the + returned text after truncation checks. thinking_budget : int, optional Token budget for model thinking (ignored by providers that don't support it). timeout_s : float, optional @@ -1209,6 +1344,9 @@ async def agenerate( """ from think.providers import get_provider_module + if json_schema is not None: + json_output = True + # Allow model override via kwargs (used by Batch for explicit model selection) model_override = kwargs.pop("model", None) @@ -1227,6 +1365,7 @@ async def agenerate( max_output_tokens=max_output_tokens, system_instruction=system_instruction, json_output=json_output, + json_schema=json_schema, thinking_budget=thinking_budget, timeout_s=timeout_s, **kwargs, @@ -1245,6 +1384,9 @@ async def agenerate( # Validate JSON output if requested _validate_json_response(result, json_output) + if json_schema is not None: + _validate_schema(result["text"], json_schema) + return result["text"] diff --git a/think/providers/anthropic.py b/think/providers/anthropic.py index 25f7f4bfe..83dee9f48 100644 --- a/think/providers/anthropic.py +++ b/think/providers/anthropic.py @@ -31,13 +31,15 @@ timeout_s : float, optional from __future__ import annotations +import json import logging import os +import re import traceback from pathlib import Path from typing import Any, Callable -from anthropic import AsyncAnthropic +from anthropic import AsyncAnthropic, BadRequestError from anthropic.types import ( MessageParam, RedactedThinkingBlock, @@ -64,6 +66,7 @@ from .shared import ( _DEFAULT_MODEL = CLAUDE_SONNET_4 logger = logging.getLogger(__name__) +_TOOL_NAME_RE = re.compile(r"^[a-zA-Z0-9_-]{1,64}$") _DEFAULT_MAX_TOKENS = 8096 * 2 _MIN_THINKING_BUDGET = 1024 # Anthropic minimum @@ -393,6 +396,23 @@ def _extract_text_and_thinking(response: Any) -> tuple[str, list | None]: return text, thinking_blocks if thinking_blocks else None +def _derive_tool_name(schema: dict | None) -> str: + """Return a valid tool name for Anthropic schema output requests.""" + if isinstance(schema, dict): + title = schema.get("title") + if isinstance(title, str) and title and _TOOL_NAME_RE.fullmatch(title): + return title + return "response" + + +def _extract_first_tool_use_json(response: Any) -> str: + """Serialize the first tool_use block input from an Anthropic response.""" + for block in getattr(response, "content", []): + if getattr(block, "type", None) == "tool_use": + return json.dumps(getattr(block, "input", None)) + raise ValueError("Anthropic schema fallback response missing tool_use block") + + # Cache for Anthropic clients _anthropic_client = None _async_anthropic_client = None @@ -447,6 +467,7 @@ def run_generate( system_instruction: str | None = None, json_output: bool = False, thinking_budget: int | None = None, + json_schema: dict | None = None, timeout_s: float | None = None, **kwargs: Any, ) -> GenerateResult: @@ -460,7 +481,7 @@ def run_generate( # Handle JSON output by adding to system instruction system = system_instruction or "" - if json_output: + if json_schema is None and json_output: json_instruction = "Respond with valid JSON only. No explanation or markdown." system = f"{system}\n\n{json_instruction}" if system else json_instruction @@ -487,9 +508,32 @@ def run_generate( if timeout_s: request_kwargs["timeout"] = timeout_s - response = client.messages.create(**request_kwargs) + if json_schema is not None: + tool_name = _derive_tool_name(json_schema) + request_kwargs["output_config"] = { + "format": {"type": "json_schema", "schema": json_schema} + } + try: + response = client.messages.create(**request_kwargs) + text, thinking = _extract_text_and_thinking(response) + except BadRequestError: + retry_kwargs = dict(request_kwargs) + retry_kwargs.pop("output_config", None) + retry_kwargs["tools"] = [ + { + "name": tool_name, + "description": "Generate the requested JSON response.", + "input_schema": json_schema, + } + ] + retry_kwargs["tool_choice"] = {"type": "tool", "name": tool_name} + response = client.messages.create(**retry_kwargs) + text = _extract_first_tool_use_json(response) + _, thinking = _extract_text_and_thinking(response) + else: + response = client.messages.create(**request_kwargs) + text, thinking = _extract_text_and_thinking(response) - text, thinking = _extract_text_and_thinking(response) return GenerateResult( text=text, usage=_extract_usage_dict(response), @@ -506,6 +550,7 @@ async def run_agenerate( system_instruction: str | None = None, json_output: bool = False, thinking_budget: int | None = None, + json_schema: dict | None = None, timeout_s: float | None = None, **kwargs: Any, ) -> GenerateResult: @@ -519,7 +564,7 @@ async def run_agenerate( # Handle JSON output by adding to system instruction system = system_instruction or "" - if json_output: + if json_schema is None and json_output: json_instruction = "Respond with valid JSON only. No explanation or markdown." system = f"{system}\n\n{json_instruction}" if system else json_instruction @@ -546,9 +591,32 @@ async def run_agenerate( if timeout_s: request_kwargs["timeout"] = timeout_s - response = await client.messages.create(**request_kwargs) + if json_schema is not None: + tool_name = _derive_tool_name(json_schema) + request_kwargs["output_config"] = { + "format": {"type": "json_schema", "schema": json_schema} + } + try: + response = await client.messages.create(**request_kwargs) + text, thinking = _extract_text_and_thinking(response) + except BadRequestError: + retry_kwargs = dict(request_kwargs) + retry_kwargs.pop("output_config", None) + retry_kwargs["tools"] = [ + { + "name": tool_name, + "description": "Generate the requested JSON response.", + "input_schema": json_schema, + } + ] + retry_kwargs["tool_choice"] = {"type": "tool", "name": tool_name} + response = await client.messages.create(**retry_kwargs) + text = _extract_first_tool_use_json(response) + _, thinking = _extract_text_and_thinking(response) + else: + response = await client.messages.create(**request_kwargs) + text, thinking = _extract_text_and_thinking(response) - text, thinking = _extract_text_and_thinking(response) return GenerateResult( text=text, usage=_extract_usage_dict(response), diff --git a/think/providers/google.py b/think/providers/google.py index 5e30ffcf1..8db17d52e 100644 --- a/think/providers/google.py +++ b/think/providers/google.py @@ -212,6 +212,7 @@ def _build_generate_config( system_instruction: str | None, json_output: bool, thinking_budget: int | None, + json_schema: dict | None = None, timeout_s: float | None = None, ) -> types.GenerateContentConfig: """Build the GenerateContentConfig. @@ -232,6 +233,8 @@ def _build_generate_config( if json_output: config_args["response_mime_type"] = "application/json" + if json_schema is not None: + config_args["response_json_schema"] = json_schema # Set thinking config when caller explicitly specified a budget. # thinking_budget=0 must explicitly disable thinking (not omit config), @@ -456,6 +459,7 @@ def run_generate( system_instruction: str | None = None, json_output: bool = False, thinking_budget: int | None = None, + json_schema: dict | None = None, timeout_s: float | None = None, **kwargs: Any, ) -> GenerateResult: @@ -475,6 +479,7 @@ def run_generate( system_instruction=system_instruction, json_output=json_output, thinking_budget=thinking_budget, + json_schema=json_schema, timeout_s=timeout_s, ) @@ -500,6 +505,7 @@ async def run_agenerate( system_instruction: str | None = None, json_output: bool = False, thinking_budget: int | None = None, + json_schema: dict | None = None, timeout_s: float | None = None, **kwargs: Any, ) -> GenerateResult: @@ -519,6 +525,7 @@ async def run_agenerate( system_instruction=system_instruction, json_output=json_output, thinking_budget=thinking_budget, + json_schema=json_schema, timeout_s=timeout_s, ) diff --git a/think/providers/ollama.py b/think/providers/ollama.py index 63a201361..43a02777a 100644 --- a/think/providers/ollama.py +++ b/think/providers/ollama.py @@ -163,6 +163,7 @@ def _build_request_body( max_output_tokens: int, json_output: bool, thinking_budget: int | None, + json_schema: dict | None = None, ) -> dict[str, Any]: """Build the native Ollama /api/chat request body. @@ -203,7 +204,9 @@ def _build_request_body( else: body["think"] = False - if json_output: + if json_schema is not None: + body["format"] = json_schema + elif json_output: body["format"] = "json" return body @@ -281,6 +284,7 @@ def run_generate( system_instruction: str | None = None, json_output: bool = False, thinking_budget: int | None = None, + json_schema: dict | None = None, timeout_s: float | None = None, **kwargs: Any, ) -> GenerateResult: @@ -299,6 +303,7 @@ def run_generate( max_output_tokens, json_output, thinking_budget, + json_schema, ) response = client.post( @@ -318,6 +323,7 @@ async def run_agenerate( system_instruction: str | None = None, json_output: bool = False, thinking_budget: int | None = None, + json_schema: dict | None = None, timeout_s: float | None = None, **kwargs: Any, ) -> GenerateResult: @@ -336,6 +342,7 @@ async def run_agenerate( max_output_tokens, json_output, thinking_budget, + json_schema, ) response = await client.post( diff --git a/think/providers/openai.py b/think/providers/openai.py index 32e029cd8..e407f207e 100644 --- a/think/providers/openai.py +++ b/think/providers/openai.py @@ -35,6 +35,7 @@ from __future__ import annotations import functools import logging import os +import re import traceback from pathlib import Path from typing import Any, Callable @@ -57,6 +58,7 @@ from .shared import ( # Agent configuration is now loaded via get_talent() in cortex.py LOG = logging.getLogger("think.providers.openai") +_SCHEMA_NAME_RE = re.compile(r"^[a-zA-Z0-9_-]{1,64}$") def _parse_model_effort(model: str) -> tuple[str, str | None]: @@ -296,6 +298,15 @@ def _build_input( return str(contents), system_instruction +def _derive_schema_name(schema: dict | None) -> str: + """Return a valid schema name for OpenAI structured outputs.""" + if isinstance(schema, dict): + title = schema.get("title") + if isinstance(title, str) and title and _SCHEMA_NAME_RE.fullmatch(title): + return title + return "response" + + def _normalize_finish_reason(response: Any) -> str | None: """Normalize OpenAI finish_reason to standard values. @@ -370,6 +381,7 @@ def run_generate( max_output_tokens: int = 8192 * 2, system_instruction: str | None = None, json_output: bool = False, + json_schema: dict | None = None, timeout_s: float | None = None, **kwargs: Any, ) -> GenerateResult: @@ -395,7 +407,16 @@ def run_generate( if effort is not None: request_kwargs["reasoning"] = {"effort": effort} - if json_output: + if json_schema is not None: + request_kwargs["text"] = { + "format": { + "type": "json_schema", + "name": _derive_schema_name(json_schema), + "schema": json_schema, + "strict": True, + } + } + elif json_output: request_kwargs["text"] = {"format": {"type": "json_object"}} if timeout_s: @@ -416,6 +437,7 @@ async def run_agenerate( max_output_tokens: int = 8192 * 2, system_instruction: str | None = None, json_output: bool = False, + json_schema: dict | None = None, timeout_s: float | None = None, **kwargs: Any, ) -> GenerateResult: @@ -441,7 +463,16 @@ async def run_agenerate( if effort is not None: request_kwargs["reasoning"] = {"effort": effort} - if json_output: + if json_schema is not None: + request_kwargs["text"] = { + "format": { + "type": "json_schema", + "name": _derive_schema_name(json_schema), + "schema": json_schema, + "strict": True, + } + } + elif json_output: request_kwargs["text"] = {"format": {"type": "json_object"}} if timeout_s: diff --git a/think/providers/shared.py b/think/providers/shared.py index a50676a9e..5e6069db1 100644 --- a/think/providers/shared.py +++ b/think/providers/shared.py @@ -169,6 +169,7 @@ class GenerateResult(TypedDict, total=False): usage: Optional[dict] # Normalized usage dict (input_tokens, output_tokens, etc.) finish_reason: Optional[str] # Normalized: "stop", "max_tokens", "safety", etc. thinking: Optional[list] # List of thinking block dicts + schema_validation: Optional[dict] # Validation result when json_schema is supplied # --------------------------------------------------------------------------- diff --git a/uv.lock b/uv.lock index 475444e42..6c4fa6663 100644 --- a/uv.lock +++ b/uv.lock @@ -3527,6 +3527,7 @@ dependencies = [ { name = "google-genai" }, { name = "httpx" }, { name = "icalendar" }, + { name = "jsonschema" }, { name = "markdown" }, { name = "mistune" }, { name = "mypy" }, @@ -3574,6 +3575,7 @@ requires-dist = [ { name = "google-genai" }, { name = "httpx" }, { name = "icalendar" }, + { name = "jsonschema", specifier = ">=4.26,<5" }, { name = "markdown" }, { name = "mistune" }, { name = "mypy" },