From 3fb36812e8d0f5c06b2228146e71599561f95992 Mon Sep 17 00:00:00 2001 From: superG Date: Mon, 28 Sep 2026 15:52:19 +0800 Subject: [PATCH 1/2] fix(ai): avoid quadratic coverage counting during analysis --- src/wechat_decrypt_tool/ai/agent_workspace.py | 6 +- tests/test_ai_workspace_coverage.py | 59 +++++++++++++++++++ 2 files changed, 63 insertions(+), 2 deletions(-) create mode 100644 tests/test_ai_workspace_coverage.py diff --git a/src/wechat_decrypt_tool/ai/agent_workspace.py b/src/wechat_decrypt_tool/ai/agent_workspace.py index 40c4f47f..a391e0e8 100644 --- a/src/wechat_decrypt_tool/ai/agent_workspace.py +++ b/src/wechat_decrypt_tool/ai/agent_workspace.py @@ -80,8 +80,10 @@ def saved_note_coverage(db, run_id, version): SELECT source,max(hi) covered_end FROM ordered GROUP BY source HAVING min(lo)=0 AND sum(CASE WHEN lo>coalesce(previous_end,0) THEN 1 ELSE 0 END)=0 ) - SELECT m.username,count(*) FROM whole w JOIN agent_material m ON m.source=w.source - WHERE m.run_id=? AND w.covered_end>=length(coalesce(json_extract(m.body,'$.text'),'')) + -- 先遍历已覆盖来源,再按 (run_id,source) 主键找原文,避免反向逐行扫描 whole。 + SELECT m.username,count(*) FROM whole w CROSS JOIN agent_material m + WHERE m.run_id=? AND m.source=w.source + AND w.covered_end>=length(coalesce(json_extract(m.body,'$.text'),'')) GROUP BY m.username ''', (run_id, version, run_id)).fetchall() count = db.execute("SELECT count(*) FROM agent_piece WHERE run_id=? AND version=? AND kind='stage_note'", diff --git a/tests/test_ai_workspace_coverage.py b/tests/test_ai_workspace_coverage.py new file mode 100644 index 00000000..db9b4ffb --- /dev/null +++ b/tests/test_ai_workspace_coverage.py @@ -0,0 +1,59 @@ +"""已分析覆盖统计的计数与查询规模回归。""" +import json + +from wechat_decrypt_tool.ai.agent_workspace import Workspace +from wechat_decrypt_tool.ai.storage import AIStore + + +def test_saved_note_coverage_scales_without_changing_counts(tmp_path): + def measure(size): + store = AIStore(tmp_path / str(size)) + workspace = Workspace(store) + run_id = 'run' + sources = [f'{index:024x}' for index in range(size)] + materials = [ + (run_id, source, 'group' if index == size - 1 else 'friend', index, + source, json.dumps({'text': 'same'})) + for index, source in enumerate(sources) + ] + notes = [] + for index, source in enumerate(sources): + # 第 0 条有缺口;第 1 条重复覆盖。相同正文仍是不同消息。 + intervals = [(0, 2), (3, 4)] if index == 0 else ( + [(0, 4), (0, 4)] if index == 1 else [(0, 4)]) + for part, (start, end) in enumerate(intervals): + notes.append((run_id, 1, f'note:{index:08d}:{part}', 'stage_note', + json.dumps({'covered': [{'source': source, 'start': start, 'end': end}]}))) + with store.connection() as db: + db.executemany('INSERT INTO agent_material VALUES(?,?,?,?,?,?)', materials) + db.executemany('INSERT INTO agent_piece VALUES(?,?,?,?,?)', notes) + # 同 source 的其他任务与旧版本笔记不能算进本轮。 + db.execute('INSERT INTO agent_material VALUES(?,?,?,?,?,?)', + ('other', sources[0], 'foreign', 0, sources[0], json.dumps({'text': 'same'}))) + db.execute('INSERT INTO agent_piece VALUES(?,?,?,?,?)', + ('other', 1, 'foreign-note', 'stage_note', + json.dumps({'covered': [{'source': sources[0], 'start': 0, 'end': 4}]}))) + db.execute('INSERT INTO agent_piece VALUES(?,?,?,?,?)', + (run_id, 2, 'old-version', 'stage_note', + json.dumps({'covered': [{'source': sources[0], 'start': 0, 'end': 4}]}))) + + steps = 0 + + def progress(): + nonlocal steps + steps += 1 + return 0 + + db.set_progress_handler(progress, 1000) + try: + counts, segments = workspace.saved_note_coverage(db, run_id, 1) + finally: + db.set_progress_handler(None, 0) + + assert counts == {'friend': size - 2, 'group': 1} + assert segments == size + 2 + return steps + + small, large = measure(100), measure(1000) + # 数据量放大十倍,不应重复交叉扫描整份原文和覆盖清单。 + assert large < small * 25, (small, large) From bfc9eb96934b8946ade4f25901aa20d137f713a7 Mon Sep 17 00:00:00 2001 From: superG Date: Mon, 28 Sep 2026 17:10:08 +0800 Subject: [PATCH 2/2] test(ai): keep SSE notification check ahead of heartbeat race Windows CI intermittently failed test_stream_waits_for_notification_ instead_of_polling...: the next frame after the event write was the ': heartbeat' comment, so parsing 'data: ' raised IndexError. A slow synchronous SQLite write can let the 30ms test heartbeat win the race; a heartbeat comment before an event frame is protocol-valid SSE, so keep production behavior and instead raise the heartbeat above the 0.5s assertion deadline for the notification check. A broken notification now fails with TimeoutError instead of passing on a heartbeat frame. --- tests/test_ai_agent_sse.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/tests/test_ai_agent_sse.py b/tests/test_ai_agent_sse.py index 270bcae7..a9354c02 100644 --- a/tests/test_ai_agent_sse.py +++ b/tests/test_ai_agent_sse.py @@ -71,10 +71,12 @@ def counted(*args, **kwargs): # 空闲期间只查询一次后等待通知,不再每 100ms 查询 SQLite。 assert await asyncio.wait_for(anext(response.body_iterator), 0.5) == ': heartbeat\n\n' assert calls == 1 - pending = asyncio.create_task(anext(response.body_iterator)) - await asyncio.sleep(0) - service.store.event('account', 'agent', {'type': 'live'}) - event = await asyncio.wait_for(pending, 0.5) - assert json.loads(event.split('data: ', 1)[1]) == {'type': 'live'} + # 通知验证的心跳要晚于断言截止时间;通知失效时应超时,而非靠心跳重查通过。 + with patch.object(ai_agent, 'SSE_HEARTBEAT_SECONDS', 2.0): + pending = asyncio.create_task(anext(response.body_iterator)) + await asyncio.sleep(0) + service.store.event('account', 'agent', {'type': 'live'}) + event = await asyncio.wait_for(pending, 0.5) + assert json.loads(event.split('data: ', 1)[1]) == {'type': 'live'} await response.body_iterator.aclose() asyncio.run(run())