Skip to content

Commit e492c74

Browse files
JerrettDavisCopilot
andcommitted
test: extend traffic learner branches
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
1 parent fb7a0ec commit e492c74

1 file changed

Lines changed: 123 additions & 0 deletions

File tree

tests/test_memory/test_traffic_learner.py

Lines changed: 123 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -255,6 +255,64 @@ async def test_no_pattern_from_success_only(self, learner: TrafficLearner):
255255
# Only environment patterns possible, no error_recovery
256256
assert stats["requests_processed"] == 1
257257

258+
@pytest.mark.asyncio
259+
async def test_on_messages_ignores_non_user_and_empty_content(self, learner: TrafficLearner):
260+
await learner.on_messages(
261+
[
262+
{"role": "assistant", "content": "do not store"},
263+
{"role": "user", "content": [{"type": "image", "source": "ignored"}]},
264+
]
265+
)
266+
assert learner.get_stats()["patterns_extracted"] == 0
267+
268+
def test_extract_tool_results_handles_non_list_and_string_results(self, learner: TrafficLearner):
269+
results = learner.extract_tool_results_from_messages(
270+
[
271+
{"role": "assistant", "content": "ignored"},
272+
{
273+
"role": "user",
274+
"content": [
275+
{"type": "tool_result", "tool_use_id": "missing", "content": "Permission denied"},
276+
],
277+
},
278+
]
279+
)
280+
assert results == [
281+
{
282+
"tool_name": "unknown",
283+
"input": {},
284+
"output": "Permission denied",
285+
"is_error": True,
286+
}
287+
]
288+
289+
def test_build_recovery_for_grep_and_command_variants(self, learner: TrafficLearner):
290+
grep_pattern = learner._build_recovery_pattern(
291+
{"tool_name": "Grep", "input": {"pattern": "foo"}, "error_category": "runtime_error"},
292+
{"tool_name": "Grep", "input": {"pattern": "bar"}},
293+
)
294+
assert grep_pattern is not None
295+
assert "Use `bar` instead" in grep_pattern.content
296+
297+
assert (
298+
learner._build_command_recovery(
299+
{"input": {"command": "same"}, "output": "", "error_category": "unknown"},
300+
{"input": {"command": "same"}},
301+
)
302+
is None
303+
)
304+
module_pattern = learner._build_command_recovery(
305+
{
306+
"input": {"command": "python test.py"},
307+
"output": "ModuleNotFoundError: No module named 'rich'",
308+
"error_category": "module_not_found",
309+
},
310+
{"input": {"command": "python -m pytest"}},
311+
)
312+
assert module_pattern is not None
313+
assert module_pattern.importance == 0.8
314+
assert module_pattern.entity_refs == ["rich"]
315+
258316

259317
# =============================================================================
260318
# Pattern Model Tests
@@ -798,6 +856,44 @@ def mk() -> ExtractedPattern:
798856
assert rows[0][0] == "seed-id"
799857
assert rows[0][2]["evidence_count"] == 4
800858

859+
@pytest.mark.asyncio
860+
async def test_stop_drains_queue_and_set_backend(self, tmp_path):
861+
db = tmp_path / "memory.db"
862+
_init_db(db)
863+
backend = _FakeBackend(db)
864+
learner = TrafficLearner(backend=None, min_evidence=1)
865+
learner.set_backend(backend)
866+
await learner._save_queue.put(
867+
ExtractedPattern(
868+
category=PatternCategory.ENVIRONMENT,
869+
content="Use tool",
870+
importance=0.6,
871+
evidence_count=2,
872+
)
873+
)
874+
await learner.stop()
875+
rows = _read_traffic_rows(db)
876+
assert len(rows) == 1
877+
assert learner.get_stats()["patterns_saved"] == 1
878+
879+
@pytest.mark.asyncio
880+
async def test_stop_breaks_when_backend_save_raises(self, tmp_path):
881+
class BrokenBackend:
882+
async def save_memory(self, **kwargs):
883+
raise RuntimeError("boom")
884+
885+
learner = TrafficLearner(backend=BrokenBackend(), min_evidence=1)
886+
await learner._save_queue.put(
887+
ExtractedPattern(
888+
category=PatternCategory.ENVIRONMENT,
889+
content="bad",
890+
importance=0.5,
891+
evidence_count=2,
892+
)
893+
)
894+
await learner.stop()
895+
assert learner.get_stats()["patterns_saved"] == 0
896+
801897

802898
# =============================================================================
803899
# flush_to_file end-to-end + early-return paths
@@ -1140,6 +1236,20 @@ async def test_missing_db_file(self, tmp_path):
11401236
assert learner._saved_hashes == set()
11411237
assert learner._persisted_ids == {}
11421238

1239+
@pytest.mark.asyncio
1240+
async def test_hydrate_to_thread_failure_is_swallowed(self, tmp_path, monkeypatch):
1241+
db = tmp_path / "memory.db"
1242+
_init_db(db)
1243+
backend = _FakeBackend(db)
1244+
learner = TrafficLearner(backend=backend, min_evidence=1)
1245+
1246+
async def boom(_func):
1247+
raise RuntimeError("boom")
1248+
1249+
monkeypatch.setattr("headroom.memory.traffic_learner.asyncio.to_thread", boom)
1250+
await learner._hydrate_persisted_state()
1251+
assert learner._saved_hashes == set()
1252+
11431253

11441254
class TestBumpEdgeCases:
11451255
@pytest.mark.asyncio
@@ -1164,6 +1274,19 @@ async def test_bump_unknown_id_is_noop(self, tmp_path):
11641274
await learner._bump_persisted_evidence("no-such-id")
11651275
assert _read_traffic_rows(db) == []
11661276

1277+
@pytest.mark.asyncio
1278+
async def test_bump_to_thread_failure_is_swallowed(self, tmp_path, monkeypatch):
1279+
db = tmp_path / "memory.db"
1280+
_init_db(db)
1281+
backend = _FakeBackend(db)
1282+
learner = TrafficLearner(backend=backend, min_evidence=1)
1283+
1284+
async def boom(_func):
1285+
raise RuntimeError("boom")
1286+
1287+
monkeypatch.setattr("headroom.memory.traffic_learner.asyncio.to_thread", boom)
1288+
await learner._bump_persisted_evidence("x")
1289+
11671290

11681291
# =============================================================================
11691292
# stop() cancels the flush task

0 commit comments

Comments
 (0)