Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 4 additions & 2 deletions src/wechat_decrypt_tool/ai/agent_workspace.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'",
Expand Down
12 changes: 7 additions & 5 deletions tests/test_ai_agent_sse.py
Original file line number Diff line number Diff line change
Expand Up @@ -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())
59 changes: 59 additions & 0 deletions tests/test_ai_workspace_coverage.py
Original file line number Diff line number Diff line change
@@ -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)
Loading