diff --git a/observe/extract.py b/observe/extract.py index ffd27a2d6..0c80cb9d4 100644 --- a/observe/extract.py +++ b/observe/extract.py @@ -22,6 +22,10 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) +_SCHEMA = json.loads( + (Path(__file__).parent / "extract.schema.json").read_text(encoding="utf-8") +) + # Default maximum frames to extract content from DEFAULT_MAX_EXTRACTIONS = 20 @@ -244,6 +248,7 @@ def _ai_select_frames( context="observe.extract.selection", system_instruction=prompt_content.text, json_output=True, + json_schema=_SCHEMA, thinking_budget=4096, max_output_tokens=1024, temperature=0.3, diff --git a/observe/extract.schema.json b/observe/extract.schema.json new file mode 100644 index 000000000..e8fca4e56 --- /dev/null +++ b/observe/extract.schema.json @@ -0,0 +1,5 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "type": "array", + "items": {"type": "integer", "minimum": 0} +} diff --git a/tests/test_extract_schema.py b/tests/test_extract_schema.py new file mode 100644 index 000000000..a68c0015c --- /dev/null +++ b/tests/test_extract_schema.py @@ -0,0 +1,60 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +import importlib +import json +from pathlib import Path + +from jsonschema import Draft202012Validator + +import think.models as models + +extract_mod = importlib.import_module("observe.extract") + +_SCHEMA = json.loads( + (Path(__file__).resolve().parents[1] / "observe" / "extract.schema.json").read_text( + encoding="utf-8" + ) +) + + +def test_extract_schema_file_is_valid_draft_2020_12(): + Draft202012Validator.check_schema(_SCHEMA) + + +def test_extract_schema_accepts_and_rejects_expected_values(): + validator = Draft202012Validator(_SCHEMA) + + assert validator.is_valid([]) + assert validator.is_valid([1, 15, 42, 89]) + assert validator.is_valid([1, 0]) + assert not validator.is_valid(["1"]) + assert not validator.is_valid([-1]) + assert not validator.is_valid([1.5]) + assert not validator.is_valid(42) + assert not validator.is_valid({"ids": [1]}) + assert not validator.is_valid([[1, 2]]) + + +def test_ai_select_frames_passes_schema_to_generate(monkeypatch): + captured = {} + + def fake_generate(**kwargs): + captured.update(kwargs) + return "[1]" + + monkeypatch.setattr(models, "generate", fake_generate) + + frames = [ + {"frame_id": 1, "timestamp": 1.0, "analysis": {"primary": "code"}}, + ] + categories = {"code": {"description": "Code editors"}} + + result = extract_mod._ai_select_frames( + frames, + max_extractions=5, + categories=categories, + ) + + assert captured["json_schema"] is extract_mod._SCHEMA + assert result == [1]