diff --git a/tests/test_generate_full.py b/tests/test_generate_full.py index e8fad0d61..99eab5e43 100644 --- a/tests/test_generate_full.py +++ b/tests/test_generate_full.py @@ -14,6 +14,7 @@ import io import json import os from pathlib import Path +from unittest.mock import MagicMock from tests.conftest import copytree_tracked from think.utils import day_path @@ -63,6 +64,22 @@ def run_generator_with_config(mod, config: dict, monkeypatch) -> list[dict]: return events +def _write_generator_file( + tmp_path: Path, + name: str, + metadata: dict, + body: str = "Test prompt", +) -> None: + (tmp_path / f"{name}.md").write_text( + f"{json.dumps(metadata, indent=2)}\n\n{body}\n", + encoding="utf-8", + ) + + +def _write_schema_file(tmp_path: Path, name: str, schema: dict) -> None: + (tmp_path / name).write_text(json.dumps(schema, indent=2), encoding="utf-8") + + def test_generate_output_ndjson(tmp_path, monkeypatch): """Test basic output generation via NDJSON protocol.""" mod = importlib.import_module("think.talents") @@ -109,6 +126,200 @@ def test_generate_output_ndjson(tmp_path, monkeypatch): assert finish_events[0]["result"] == MOCK_RESULT["text"] +def test_dispatcher_passes_json_schema(tmp_path, monkeypatch): + """Test that generator execution forwards json_schema to the model layer.""" + mod = importlib.import_module("think.talents") + copy_day(tmp_path) + + import think.models + import think.talent + + monkeypatch.setattr(think.talent, "TALENT_DIR", tmp_path) + schema = {"type": "object", "properties": {"summary": {"type": "string"}}} + _write_schema_file(tmp_path, "schema.json", schema) + _write_generator_file( + tmp_path, + "schema_gen", + { + "type": "generate", + "schedule": "daily", + "priority": 10, + "output": "json", + "schema": "schema.json", + "load": {"transcripts": True, "percepts": True}, + }, + ) + + mock_generate = MagicMock( + return_value={ + "text": '{"summary":"ok"}', + "usage": {"input_tokens": 10, "output_tokens": 5}, + } + ) + monkeypatch.setattr(think.models, "generate_with_result", mock_generate) + monkeypatch.setenv("GOOGLE_API_KEY", "x") + monkeypatch.setenv("_SOLSTONE_JOURNAL_OVERRIDE", str(tmp_path)) + + events = run_generator_with_config( + mod, + { + "name": "schema_gen", + "day": "20240101", + "output": "json", + "provider": "google", + "model": "gemini-2.0-flash", + }, + monkeypatch, + ) + + assert mock_generate.call_args.kwargs["json_schema"] == schema + finish_events = [e for e in events if e["event"] == "finish"] + assert len(finish_events) == 1 + + +def test_dispatcher_omits_json_schema_when_absent(tmp_path, monkeypatch): + """Test that generator execution passes json_schema=None when absent.""" + mod = importlib.import_module("think.talents") + copy_day(tmp_path) + + import think.models + import think.talent + + monkeypatch.setattr(think.talent, "TALENT_DIR", tmp_path) + _write_generator_file( + tmp_path, + "plain_gen", + { + "type": "generate", + "schedule": "daily", + "priority": 10, + "output": "md", + "load": {"transcripts": True, "percepts": True}, + }, + ) + + mock_generate = MagicMock(return_value=MOCK_RESULT) + monkeypatch.setattr(think.models, "generate_with_result", mock_generate) + monkeypatch.setenv("GOOGLE_API_KEY", "x") + monkeypatch.setenv("_SOLSTONE_JOURNAL_OVERRIDE", str(tmp_path)) + + run_generator_with_config( + mod, + { + "name": "plain_gen", + "day": "20240101", + "output": "md", + "provider": "google", + "model": "gemini-2.0-flash", + }, + monkeypatch, + ) + + assert mock_generate.call_args.kwargs["json_schema"] is None + + +def test_finish_event_includes_schema_validation(tmp_path, monkeypatch): + """Test that finish events surface schema_validation when returned.""" + mod = importlib.import_module("think.talents") + copy_day(tmp_path) + + import think.models + import think.talent + + monkeypatch.setattr(think.talent, "TALENT_DIR", tmp_path) + schema = {"type": "object", "properties": {"summary": {"type": "string"}}} + validation = {"valid": True, "errors": []} + _write_schema_file(tmp_path, "schema.json", schema) + _write_generator_file( + tmp_path, + "schema_validation_gen", + { + "type": "generate", + "schedule": "daily", + "priority": 10, + "output": "json", + "schema": "schema.json", + "load": {"transcripts": True, "percepts": True}, + }, + ) + + monkeypatch.setattr( + think.models, + "generate_with_result", + MagicMock( + return_value={ + "text": '{"summary":"ok"}', + "usage": {"input_tokens": 10, "output_tokens": 5}, + "schema_validation": validation, + } + ), + ) + monkeypatch.setenv("GOOGLE_API_KEY", "x") + monkeypatch.setenv("_SOLSTONE_JOURNAL_OVERRIDE", str(tmp_path)) + + events = run_generator_with_config( + mod, + { + "name": "schema_validation_gen", + "day": "20240101", + "output": "json", + "provider": "google", + "model": "gemini-2.0-flash", + }, + monkeypatch, + ) + + finish_events = [e for e in events if e["event"] == "finish"] + assert len(finish_events) == 1 + assert finish_events[0]["schema_validation"] == validation + + +def test_finish_event_omits_schema_validation_when_absent(tmp_path, monkeypatch): + """Test that finish events omit schema_validation when not returned.""" + mod = importlib.import_module("think.talents") + copy_day(tmp_path) + + import think.models + import think.talent + + monkeypatch.setattr(think.talent, "TALENT_DIR", tmp_path) + _write_generator_file( + tmp_path, + "no_schema_validation_gen", + { + "type": "generate", + "schedule": "daily", + "priority": 10, + "output": "md", + "load": {"transcripts": True, "percepts": True}, + }, + ) + + monkeypatch.setattr( + think.models, + "generate_with_result", + MagicMock(return_value=MOCK_RESULT), + ) + monkeypatch.setenv("GOOGLE_API_KEY", "x") + monkeypatch.setenv("_SOLSTONE_JOURNAL_OVERRIDE", str(tmp_path)) + + events = run_generator_with_config( + mod, + { + "name": "no_schema_validation_gen", + "day": "20240101", + "output": "md", + "provider": "google", + "model": "gemini-2.0-flash", + }, + monkeypatch, + ) + + finish_events = [e for e in events if e["event"] == "finish"] + assert len(finish_events) == 1 + assert "schema_validation" not in finish_events[0] + + def test_generate_hook_invoked_with_context(tmp_path, monkeypatch): """Test that hooks receive correct context including span flag.""" mod = importlib.import_module("think.talents") diff --git a/tests/test_talent.py b/tests/test_talent.py index 0be352919..a39ce0b62 100644 --- a/tests/test_talent.py +++ b/tests/test_talent.py @@ -3,8 +3,12 @@ """Tests for think.talent module.""" +import json +from pathlib import Path + import pytest +from think import talent as talent_module from think.talent import ( _validate_cwd, get_talent, @@ -105,3 +109,162 @@ def test_get_agent_normalizes_cwd_for_cogitate(): def test_get_agent_preserves_repo_cwd_for_coder(): config = get_talent("coder") assert config["cwd"] == "repo" + + +def _write_talent_file(tmp_path: Path, name: str, metadata: dict) -> Path: + md_path = tmp_path / f"{name}.md" + md_path.write_text( + f"{json.dumps(metadata, indent=2)}\n\nTest prompt\n", + encoding="utf-8", + ) + return md_path + + +def _write_schema_file(path: Path, schema: dict) -> None: + path.write_text(json.dumps(schema, indent=2), encoding="utf-8") + + +def test_schema_absent_no_json_schema_key(tmp_path, monkeypatch): + monkeypatch.setattr(talent_module, "TALENT_DIR", tmp_path) + _write_talent_file( + tmp_path, "schema_absent", {"type": "generate", "output": "json"} + ) + + config = get_talent("schema_absent") + + assert "json_schema" not in config + assert "schema" not in config + + +def test_schema_loads_valid_file(tmp_path, monkeypatch): + monkeypatch.setattr(talent_module, "TALENT_DIR", tmp_path) + schema = {"type": "object", "properties": {"value": {"type": "string"}}} + _write_schema_file(tmp_path / "schema.json", schema) + _write_talent_file( + tmp_path, + "schema_valid", + {"type": "generate", "output": "json", "schema": "schema.json"}, + ) + + config = get_talent("schema_valid") + + assert config["json_schema"] == schema + assert "schema" not in config + + +def test_schema_absolute_path_rejected(tmp_path, monkeypatch): + monkeypatch.setattr(talent_module, "TALENT_DIR", tmp_path) + _write_talent_file( + tmp_path, + "schema_absolute", + {"type": "generate", "output": "json", "schema": "/etc/passwd"}, + ) + + with pytest.raises( + ValueError, + match=r"talent schema_absolute: schema path must be relative: /etc/passwd", + ): + get_talent("schema_absolute") + + +def test_schema_parent_traversal_rejected(tmp_path, monkeypatch): + monkeypatch.setattr(talent_module, "TALENT_DIR", tmp_path) + _write_talent_file( + tmp_path, + "schema_parent", + {"type": "generate", "output": "json", "schema": "../escape.json"}, + ) + + with pytest.raises( + ValueError, + match=r"talent schema_parent: schema path must not contain '\.\.': \.\./escape\.json", + ): + get_talent("schema_parent") + + +def test_schema_symlink_escape_rejected(tmp_path, monkeypatch): + monkeypatch.setattr(talent_module, "TALENT_DIR", tmp_path) + outside_schema = tmp_path.parent / "outside_schema.json" + _write_schema_file(outside_schema, {"type": "object"}) + try: + (tmp_path / "schema.json").symlink_to(outside_schema) + except (NotImplementedError, OSError) as exc: + pytest.skip(f"symlinks unavailable on this filesystem: {exc}") + + _write_talent_file( + tmp_path, + "schema_symlink", + {"type": "generate", "output": "json", "schema": "schema.json"}, + ) + + with pytest.raises( + ValueError, + match=r"talent schema_symlink: schema path escapes talent directory:", + ): + get_talent("schema_symlink") + + +def test_schema_missing_file(tmp_path, monkeypatch): + monkeypatch.setattr(talent_module, "TALENT_DIR", tmp_path) + _write_talent_file( + tmp_path, + "schema_missing", + {"type": "generate", "output": "json", "schema": "missing.json"}, + ) + + with pytest.raises( + FileNotFoundError, + match=r"talent schema_missing: schema file not found: .*missing\.json", + ): + get_talent("schema_missing") + + +def test_schema_malformed_json(tmp_path, monkeypatch): + monkeypatch.setattr(talent_module, "TALENT_DIR", tmp_path) + schema_path = tmp_path / "broken.json" + schema_path.write_text("{\n", encoding="utf-8") + _write_talent_file( + tmp_path, + "schema_malformed", + {"type": "generate", "output": "json", "schema": "broken.json"}, + ) + + with pytest.raises( + ValueError, + match=r"talent schema_malformed: schema file is not valid JSON: .*broken\.json", + ): + get_talent("schema_malformed") + + +def test_schema_invalid_schema_draft(tmp_path, monkeypatch): + monkeypatch.setattr(talent_module, "TALENT_DIR", tmp_path) + _write_schema_file(tmp_path / "invalid_schema.json", {"type": 3}) + _write_talent_file( + tmp_path, + "schema_invalid", + {"type": "generate", "output": "json", "schema": "invalid_schema.json"}, + ) + + with pytest.raises( + ValueError, + match=( + r"talent schema_invalid: schema file is not a valid JSON Schema: " + r".*invalid_schema\.json" + ), + ): + get_talent("schema_invalid") + + +def test_schema_not_string(tmp_path, monkeypatch): + monkeypatch.setattr(talent_module, "TALENT_DIR", tmp_path) + _write_talent_file( + tmp_path, + "schema_not_string", + {"type": "generate", "output": "json", "schema": 42}, + ) + + with pytest.raises( + ValueError, + match=r"talent schema_not_string: schema must be a string, got int: 42", + ): + get_talent("schema_not_string") diff --git a/think/talent.py b/think/talent.py index d6001e40f..958cd4922 100644 --- a/think/talent.py +++ b/think/talent.py @@ -18,11 +18,13 @@ use think.prompts.load_prompt() directly. from __future__ import annotations import importlib.util +import json import os from pathlib import Path from typing import Any, Callable import frontmatter +from jsonschema import Draft202012Validator, SchemaError # Import core prompt utilities from think.prompts from think.prompts import _load_prompt_metadata, load_prompt @@ -447,6 +449,54 @@ def get_talent_filter(value: bool | str | dict) -> dict[str, bool | str] | None: # --------------------------------------------------------------------------- +def _load_talent_schema( + *, + name: str, + md_path: Path, + raw_schema: Any, +) -> dict[str, Any]: + """Load and validate a talent JSON Schema from a relative file path.""" + if not isinstance(raw_schema, str): + raise ValueError( + f"talent {name}: schema must be a string, got {type(raw_schema).__name__}: " + f"{raw_schema!r}" + ) + + raw_path = Path(raw_schema) + if raw_path.is_absolute(): + raise ValueError(f"talent {name}: schema path must be relative: {raw_schema}") + if ".." in raw_path.parts: + raise ValueError( + f"talent {name}: schema path must not contain '..': {raw_schema}" + ) + + talent_dir = md_path.parent.resolve() + schema_path = (md_path.parent / raw_schema).resolve() + if not schema_path.is_relative_to(talent_dir): + raise ValueError( + f"talent {name}: schema path escapes talent directory: {schema_path}" + ) + if not schema_path.exists(): + raise FileNotFoundError(f"talent {name}: schema file not found: {schema_path}") + + try: + with open(schema_path, encoding="utf-8") as f: + parsed = json.load(f) + except json.JSONDecodeError as exc: + raise ValueError( + f"talent {name}: schema file is not valid JSON: {schema_path}" + ) from exc + + try: + Draft202012Validator.check_schema(parsed) + except SchemaError as exc: + raise ValueError( + f"talent {name}: schema file is not a valid JSON Schema: {schema_path}" + ) from exc + + return parsed + + def get_talent( name: str = "unified", facet: str | None = None, @@ -501,6 +551,14 @@ def get_talent( # Store path for later use config["path"] = str(md_path) + if "schema" in config: + config["json_schema"] = _load_talent_schema( + name=name, + md_path=md_path, + raw_schema=config["schema"], + ) + del config["schema"] + # Extract source config from 'load' key (replaces instructions.sources) config["sources"] = config.pop("load", _DEFAULT_LOAD.copy()) diff --git a/think/talents.py b/think/talents.py index 10dbed7b5..bbf5c9652 100644 --- a/think/talents.py +++ b/think/talents.py @@ -970,6 +970,7 @@ async def _execute_generate( thinking_budget=thinking_budget, system_instruction=system_instruction, json_output=is_json_output, + json_schema=config.get("json_schema"), timeout_s=timeout_s, ) except Exception as exc: @@ -1015,6 +1016,7 @@ async def _execute_generate( thinking_budget=thinking_budget, system_instruction=system_instruction, json_output=is_json_output, + json_schema=config.get("json_schema"), timeout_s=timeout_s, provider=backup, model=backup_model, @@ -1046,6 +1048,8 @@ async def _execute_generate( } if usage_data: finish_event["usage"] = usage_data + if "schema_validation" in gen_result: + finish_event["schema_validation"] = gen_result["schema_validation"] emit_event(finish_event)