personal memory agent
1# SPDX-License-Identifier: AGPL-3.0-only
2# Copyright (c) 2026 sol pbc
3
4import importlib
5import json
6from pathlib import Path
7
8from jsonschema import Draft202012Validator
9
10import solstone.think.models as models
11
12mod = importlib.import_module("solstone.think.detect_transcript")
13
14DETECT_TRANSCRIPT_SEGMENT_SCHEMA_PATH = (
15 Path(__file__).resolve().parents[1]
16 / "solstone"
17 / "think"
18 / "detect_transcript_segment.schema.json"
19)
20DETECT_TRANSCRIPT_JSON_SCHEMA_PATH = (
21 Path(__file__).resolve().parents[1]
22 / "solstone"
23 / "think"
24 / "detect_transcript_json.schema.json"
25)
26
27
28def _load_detect_transcript_segment_schema() -> dict:
29 return json.loads(DETECT_TRANSCRIPT_SEGMENT_SCHEMA_PATH.read_text(encoding="utf-8"))
30
31
32def _load_detect_transcript_json_schema() -> dict:
33 return json.loads(DETECT_TRANSCRIPT_JSON_SCHEMA_PATH.read_text(encoding="utf-8"))
34
35
36def test_detect_transcript_segment_schema_file_is_valid_draft_2020_12():
37 Draft202012Validator.check_schema(_load_detect_transcript_segment_schema())
38
39
40def test_detect_transcript_json_schema_file_is_valid_draft_2020_12():
41 Draft202012Validator.check_schema(_load_detect_transcript_json_schema())
42
43
44def test_detect_transcript_segment_schema_accepts_and_rejects_expected_values():
45 schema = _load_detect_transcript_segment_schema()
46 validator = Draft202012Validator(schema)
47 valid = {"segments": [{"start_at": "12:34:56", "line": 1}]}
48
49 assert validator.is_valid(valid)
50 assert not validator.is_valid([{"start_at": "12:34:56", "line": 1}])
51 assert not validator.is_valid({"segments": [{"start_at": "12:34:56"}]})
52 assert not validator.is_valid({"segments": [{"start_at": "12:34", "line": 1}]})
53 assert not validator.is_valid({"segments": [{"start_at": "12:34:56", "line": "1"}]})
54 assert not validator.is_valid(
55 {"segments": [{"start_at": "12:34:56", "line": 1, "extra": "x"}]}
56 )
57
58
59def test_detect_transcript_json_schema_accepts_and_rejects_expected_values():
60 schema = _load_detect_transcript_json_schema()
61 validator = Draft202012Validator(schema)
62 valid = {
63 "entries": [{"start": "12:34:56", "speaker": "Alice", "text": "Hello"}],
64 "topics": "planning, budget",
65 "setting": "workplace",
66 }
67
68 assert validator.is_valid(valid)
69 assert validator.is_valid({**valid, "topics": "", "setting": ""})
70 assert not validator.is_valid(
71 {"topics": "planning, budget", "setting": "workplace"}
72 )
73 assert not validator.is_valid(
74 {
75 **valid,
76 "entries": [{"start": "12:34", "speaker": "Alice", "text": "Hello"}],
77 }
78 )
79 assert not validator.is_valid(
80 {
81 **valid,
82 "entries": [{"start": "12:34:56", "speaker": 1, "text": "Hello"}],
83 }
84 )
85 assert not validator.is_valid(
86 {
87 **valid,
88 "entries": [{"start": "12:34:56", "speaker": "Alice", "text": 7}],
89 }
90 )
91 assert not validator.is_valid({**valid, "extra": "x"})
92 assert not validator.is_valid(
93 {
94 **valid,
95 "entries": [
96 {
97 "start": "12:34:56",
98 "speaker": "Alice",
99 "text": "Hello",
100 "extra": "x",
101 }
102 ],
103 }
104 )
105
106
107def test_detect_transcript_segment_passes_schema_to_generate(monkeypatch):
108 captured = {}
109
110 def fake_generate(**kwargs):
111 captured.update(kwargs)
112 return '{"segments": [{"start_at": "12:00:00", "line": 1}]}'
113
114 monkeypatch.setattr(models, "generate", fake_generate)
115
116 result = mod.detect_transcript_segment("01\n02\n", "12:00:00")
117
118 assert captured["json_schema"] is mod._SEGMENT_SCHEMA
119 assert result
120 assert all(isinstance(item, tuple) and len(item) == 2 for item in result)
121
122
123def test_detect_transcript_segment_schema_validation_error_returns_empty(monkeypatch):
124 def fake_generate(**kwargs):
125 raise models.SchemaValidationError(
126 [{"path": "", "constraint": "json_parse", "message": "empty"}],
127 "",
128 )
129
130 monkeypatch.setattr(models, "generate", fake_generate)
131
132 assert mod.detect_transcript_segment("01\n02\n", "12:00:00") == []
133
134
135def test_detect_transcript_json_passes_schema_to_generate(monkeypatch):
136 captured = {}
137
138 def fake_generate(**kwargs):
139 captured.update(kwargs)
140 return (
141 '{"entries": [{"start": "12:00:00", "speaker": "Alice", "text": "Hello"}], '
142 '"topics": "planning", "setting": "workplace"}'
143 )
144
145 monkeypatch.setattr(models, "generate", fake_generate)
146
147 result = mod.detect_transcript_json("some text", "12:00:00")
148
149 assert captured["json_schema"] is mod._JSON_SCHEMA
150 assert result == {
151 "entries": [{"start": "12:00:00", "speaker": "Alice", "text": "Hello"}],
152 "topics": "planning",
153 "setting": "workplace",
154 }
155
156
157def test_detect_transcript_json_schema_validation_error_returns_none(monkeypatch):
158 def fake_generate(**kwargs):
159 raise models.SchemaValidationError(
160 [{"path": "", "constraint": "json_parse", "message": "empty"}],
161 "",
162 )
163
164 monkeypatch.setattr(models, "generate", fake_generate)
165
166 assert mod.detect_transcript_json("some text", "12:00:00") is None