From 2e7bb53ad25f9833215c5ec00b6dc6c1462aa9e0 Mon Sep 17 00:00:00 2001 From: Jer Miller Date: Sun, 19 Apr 2026 20:08:35 -0600 Subject: [PATCH] observe/transcribe: schema-constrain + drop shape tolerance Add a Draft 2020-12 JSON schema for Gemini's transcribe response and pass it via json_schema on generate(). Tighten _extract_segments to accept only the documented {"segments": [...]} wrapper; bare-list, array-wrapped-dict, and {"transcript": ...} fallbacks now raise. Co-Authored-By: Claude Opus 4.7 (1M context) --- observe/transcribe/gemini.py | 67 ++++---------- observe/transcribe/gemini.schema.json | 30 +++++++ tests/test_transcribe_gemini.py | 58 ++++++------ tests/test_transcribe_gemini_schema.py | 118 +++++++++++++++++++++++++ 4 files changed, 198 insertions(+), 75 deletions(-) create mode 100644 observe/transcribe/gemini.schema.json create mode 100644 tests/test_transcribe_gemini_schema.py diff --git a/observe/transcribe/gemini.py b/observe/transcribe/gemini.py index abfa64d9e..bc331fbfc 100644 --- a/observe/transcribe/gemini.py +++ b/observe/transcribe/gemini.py @@ -36,6 +36,10 @@ from think.prompts import load_prompt logger = logging.getLogger(__name__) +_SCHEMA = json.loads( + (Path(__file__).parent / "gemini.schema.json").read_text(encoding="utf-8") +) + # Regex for parsing speaker strings like "Speaker 1", "Speaker 2" SPEAKER_PATTERN = re.compile(r"(?:speaker\s*)?(\d+)", re.IGNORECASE) @@ -123,54 +127,20 @@ def _parse_speaker(speaker: str | int | None) -> int | None: return None -def _extract_segments(result: list | dict) -> list: - """Extract a segments list from Gemini's JSON response. - - Gemini may return any of these shapes: - - ``{"segments": [...]}`` — expected format per prompt - - ``[{"start": ..., "text": ...}, ...]`` — bare list of segment dicts - - ``[{"segments": [...]}]`` — array-wrapped dict - - ``{"transcript": [...]}`` or other single-key wrapper - - Args: - result: Parsed JSON (list or dict). +def _extract_segments(result: dict) -> list: + """Extract the segments list from Gemini's schema-constrained response. - Returns: - List of segment dicts, empty list on unexpected input. + Raises RuntimeError if the result does not match the documented + {"segments": [...]} wrapper shape. """ - # Unwrap: [{"segments": [...]}] — array-wrapped dict with a segments key. - # Only unwrap if the inner dict has "segments", not if it's a segment itself. - if ( - isinstance(result, list) - and len(result) == 1 - and isinstance(result[0], dict) - and "segments" in result[0] - ): - result = result[0] - - # Bare list of segment dicts - if isinstance(result, list): - return result - - if isinstance(result, dict): - # Preferred key - if "segments" in result: - val = result["segments"] - if isinstance(val, list): - return val - - # Fallback: single-key dict whose value is a list (e.g. {"transcript": [...]}) - if len(result) == 1: - val = next(iter(result.values())) - if isinstance(val, list): - logger.warning( - f"Gemini used unexpected key {next(iter(result.keys()))!r} " - f"instead of 'segments'" - ) - return val - - logger.warning(f"Gemini returned unexpected result shape: {type(result)}") - return [] + if isinstance(result, dict) and isinstance(result.get("segments"), list): + return result["segments"] + logger.warning( + "Gemini returned unexpected shape: type=%s keys=%s", + type(result).__name__, + list(result.keys()) if isinstance(result, dict) else None, + ) + raise RuntimeError(f"Gemini returned unexpected shape: {type(result).__name__}") def _build_chunk_contents( @@ -380,6 +350,7 @@ def transcribe( max_output_tokens=16384, json_output=True, thinking_budget=0, + json_schema=_SCHEMA, ) transcribe_time = time.perf_counter() - t0 @@ -395,10 +366,6 @@ def transcribe( logger.debug(f"Response text: {response_text[:500]}") raise RuntimeError(f"Gemini returned invalid JSON: {e}") from e - # Extract segments — Gemini may return different shapes: - # [...] bare list of segments - # {"segments": [...]} expected wrapper - # {"transcript": [...]} alternate key name segments = _extract_segments(result) # Normalize to standard statement format diff --git a/observe/transcribe/gemini.schema.json b/observe/transcribe/gemini.schema.json new file mode 100644 index 000000000..b769e8bd7 --- /dev/null +++ b/observe/transcribe/gemini.schema.json @@ -0,0 +1,30 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "type": "object", + "additionalProperties": false, + "required": ["segments"], + "properties": { + "segments": { + "type": "array", + "items": { + "type": "object", + "additionalProperties": false, + "required": ["start", "speaker", "text"], + "properties": { + "start": { + "type": "string", + "pattern": "^\\d{2}:\\d{2}$" + }, + "speaker": { + "type": "string", + "minLength": 1 + }, + "text": { + "type": "string", + "minLength": 1 + } + } + } + } + } +} diff --git a/tests/test_transcribe_gemini.py b/tests/test_transcribe_gemini.py index 46955b584..eab4d0806 100644 --- a/tests/test_transcribe_gemini.py +++ b/tests/test_transcribe_gemini.py @@ -4,6 +4,7 @@ """Tests for the Gemini STT backend.""" import numpy as np +import pytest from observe.transcribe.gemini import ( _build_chunk_contents, @@ -259,47 +260,54 @@ class TestBuildChunkContents: class TestExtractSegments: - """Tests for _extract_segments — robust response parsing.""" + """Tests for _extract_segments strict wrapper parsing.""" def test_expected_dict_wrapper(self): """Standard {"segments": [...]} response.""" segs = [{"start": "00:00", "speaker": "Speaker 1", "text": "Hi"}] assert _extract_segments({"segments": segs}) == segs - def test_bare_list(self): - """Gemini returns bare list of segment dicts.""" + def test_bare_list_raises(self): + """Bare list is rejected.""" segs = [{"start": "00:00", "speaker": "Speaker 1", "text": "Hi"}] - assert _extract_segments(segs) == segs + with pytest.raises(RuntimeError): + _extract_segments(segs) - def test_alternate_key(self): - """Single-key dict with alternate key name.""" + def test_alternate_key_raises(self): + """Alternate wrapper key is rejected.""" segs = [{"start": "00:00", "text": "Hi"}] - assert _extract_segments({"transcript": segs}) == segs + with pytest.raises(RuntimeError): + _extract_segments({"transcript": segs}) - def test_array_wrapped_dict(self): - """Gemini wraps response in array: [{"segments": [...]}].""" + def test_array_wrapped_dict_raises(self): + """Array-wrapped dict is rejected.""" segs = [{"start": "00:00", "speaker": "Speaker 1", "text": "Hi"}] - assert _extract_segments([{"segments": segs}]) == segs + with pytest.raises(RuntimeError): + _extract_segments([{"segments": segs}]) def test_empty_segments(self): """Empty segments list in dict.""" assert _extract_segments({"segments": []}) == [] - def test_empty_bare_list(self): - """Empty bare list.""" - assert _extract_segments([]) == [] - - def test_non_list_segments_value(self): - """segments key has non-list value.""" - assert _extract_segments({"segments": "not a list"}) == [] - - def test_unexpected_type(self): - """Completely unexpected type returns empty.""" - assert _extract_segments("unexpected") == [] - - def test_dict_with_no_segments_key(self): - """Dict without segments or single-list key returns empty.""" - assert _extract_segments({"other": "value", "more": "stuff"}) == [] + def test_empty_bare_list_raises(self): + """Empty bare list is rejected.""" + with pytest.raises(RuntimeError): + _extract_segments([]) + + def test_non_list_segments_value_raises(self): + """Non-list segments value is rejected.""" + with pytest.raises(RuntimeError): + _extract_segments({"segments": "not a list"}) + + def test_unexpected_type_raises(self): + """Unexpected type is rejected.""" + with pytest.raises(RuntimeError): + _extract_segments("unexpected") + + def test_dict_with_no_segments_key_raises(self): + """Dict without segments key is rejected.""" + with pytest.raises(RuntimeError): + _extract_segments({"other": 1}) class TestGetModelInfo: diff --git a/tests/test_transcribe_gemini_schema.py b/tests/test_transcribe_gemini_schema.py new file mode 100644 index 000000000..2d2b8ed56 --- /dev/null +++ b/tests/test_transcribe_gemini_schema.py @@ -0,0 +1,118 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +import json +from pathlib import Path +from types import SimpleNamespace + +import numpy as np +from jsonschema import Draft202012Validator + +import observe.transcribe.gemini as gemini_mod + + +def _load_schema() -> dict: + with ( + Path(__file__).resolve().parents[1] + / "observe" + / "transcribe" + / "gemini.schema.json" + ).open(encoding="utf-8") as f: + return json.load(f) + + +def test_gemini_schema_file_is_valid_draft_2020_12(): + Draft202012Validator.check_schema(_load_schema()) + + +def test_gemini_schema_accepts_and_rejects_expected_values(): + validator = Draft202012Validator(_load_schema()) + + assert validator.is_valid({"segments": []}) + assert validator.is_valid( + {"segments": [{"start": "01:23", "speaker": "Speaker 1", "text": "hi"}]} + ) + assert validator.is_valid( + { + "segments": [ + {"start": "00:00", "speaker": "Speaker 1", "text": "hello"}, + {"start": "00:05", "speaker": "Speaker 2", "text": "hi back"}, + ] + } + ) + assert not validator.is_valid( + [{"start": "01:23", "speaker": "Speaker 1", "text": "hi"}] + ) + assert not validator.is_valid( + {"transcript": [{"start": "01:23", "speaker": "Speaker 1", "text": "hi"}]} + ) + assert not validator.is_valid({"segments": [], "extra": 1}) + assert not validator.is_valid( + {"segments": [{"speaker": "Speaker 1", "text": "hi"}]} + ) + assert not validator.is_valid({"segments": [{"start": "01:23", "text": "hi"}]}) + assert not validator.is_valid( + {"segments": [{"start": "01:23", "speaker": "Speaker 1"}]} + ) + assert not validator.is_valid( + { + "segments": [ + { + "start": "01:23", + "speaker": "s", + "text": "t", + "confidence": 0.9, + } + ] + } + ) + assert not validator.is_valid( + {"segments": [{"start": "01:23", "speaker": "Speaker 1", "text": ""}]} + ) + assert not validator.is_valid( + {"segments": [{"start": "01:23", "speaker": "", "text": "hi"}]} + ) + assert not validator.is_valid( + {"segments": [{"start": "1:23", "speaker": "Speaker 1", "text": "hi"}]} + ) + assert not validator.is_valid( + {"segments": [{"start": "01:23:45", "speaker": "Speaker 1", "text": "hi"}]} + ) + assert not validator.is_valid( + {"segments": [{"start": "01-23", "speaker": "Speaker 1", "text": "hi"}]} + ) + assert not validator.is_valid( + {"segments": [{"start": 83, "speaker": "Speaker 1", "text": "hi"}]} + ) + + +def test_transcribe_passes_schema_to_generate(monkeypatch): + captured = {} + + def fake_generate(**kwargs): + captured.update(kwargs) + return json.dumps( + {"segments": [{"start": "00:00", "speaker": "Speaker 1", "text": "hello"}]} + ) + + monkeypatch.setattr(gemini_mod, "generate", fake_generate) + monkeypatch.setattr(gemini_mod, "audio_to_flac_bytes", lambda *_args: b"flac") + monkeypatch.setattr( + gemini_mod.types.Part, + "from_bytes", + staticmethod(lambda data, mime_type: {"data": data, "mime_type": mime_type}), + ) + monkeypatch.setattr( + gemini_mod, + "load_prompt", + lambda *_args, **_kwargs: SimpleNamespace(text="Prompt"), + ) + + gemini_mod.transcribe( + np.zeros(16000, dtype=np.float32), + 16000, + {}, + [(0.0, 1.0)], + ) + + assert captured["json_schema"] is gemini_mod._SCHEMA -- 2.51.2