Skip to content
Closed
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
15 changes: 14 additions & 1 deletion src/wechat_decrypt_tool/local_search/frozen.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,8 @@ def __init__(self, path, job):
CREATE TABLE IF NOT EXISTS segments(position INTEGER PRIMARY KEY, body TEXT);
CREATE TABLE IF NOT EXISTS messages(segment INTEGER, ordinal INTEGER, source TEXT, body TEXT,
PRIMARY KEY(segment,ordinal), UNIQUE(segment,source));
CREATE INDEX IF NOT EXISTS messages_source ON messages(source);
CREATE TABLE IF NOT EXISTS reconcile(source TEXT PRIMARY KEY);
''')
if not self._get(db, 'plan'):
self._put(db, 'plan', {'job_id': job['id'], 'base': job['processed'],
Expand Down Expand Up @@ -65,9 +67,10 @@ def append(self, position, result, checkpoint):
(position, state['ordinal'] + 1, message['source'], json.dumps(message, ensure_ascii=False)))
if inserted.rowcount:
state['ordinal'] += 1
warning = ';'.join(dict.fromkeys(filter(None, [state.get('warning', ''), result.get('warning', '')])))
state.update(offset=state['offset'] + len(result['messages']), cursor=result.get('cursor'),
complete=not result.get('has_more', False), name=result.get('name', ''),
source=result.get('source', 'snapshot'), warning=result.get('warning', ''))
source=result.get('source', 'snapshot'), warning=warning)
checkpoint()
db.execute('INSERT OR REPLACE INTO segments VALUES(?,?)',
(position, json.dumps(state, ensure_ascii=False)))
Expand Down Expand Up @@ -104,6 +107,16 @@ def page(self, position, offset, size, checkpoint):
return {'messages': messages, 'has_more': more,
'name': state['name'], 'source': state['source'], 'warning': state['warning']}

def reconciliation_scopes(self, targets):
"""只返回读取完整且无来源警告、可安全核对删除的快照范围。"""
with self.connection() as db:
states = {row[0]: json.loads(row[1])
for row in db.execute('SELECT position,body FROM segments')}
return [{'segment': position, 'username': target['username'],
'start': target['start'], 'end': target['end']}
for position, target in enumerate(targets)
if (state := states.get(position)) and state.get('complete') and not state.get('warning')]

def discard(self):
# 只删除由任务 ID 推导的临时快照文件;索引与聊天数据库不受影响。
self.path.unlink(missing_ok=True)
44 changes: 43 additions & 1 deletion src/wechat_decrypt_tool/local_search/index.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,13 +58,28 @@ def stats(self, generation):
return {'messages': messages, 'chunks': chunks}

@observed('index.commit')
def commit(self, generation, messages, chunks, vectors, job, checkpoint=None):
def commit(self, generation, messages, chunks, vectors, job, checkpoint=None, reconcile_snapshot=None):
import sqlite_vec
if len(chunks) != len(vectors):
raise ValueError('向量数量与片段数量不一致,未提交当前批次')
with self.connection() as db:
if reconcile_snapshot:
db.execute('ATTACH DATABASE ? AS snapshot', (str(reconcile_snapshot),))
db.execute('BEGIN IMMEDIATE')
if checkpoint: checkpoint()
if reconcile_snapshot:
# 共享片段的向量包含全部成员文本;删除任一成员时必须移除整片,
# 随后由 messages/chunks 参数原子地重建仍存在的相邻消息。
db.execute('CREATE TEMP TABLE reconcile_chunks(id TEXT PRIMARY KEY)')
db.execute('''INSERT OR IGNORE INTO reconcile_chunks
SELECT DISTINCT c.id FROM chunks c
JOIN members m ON m.chunk=c.id
JOIN snapshot.reconcile r ON r.source=m.source
WHERE c.generation=?''', (generation,))
db.execute('DELETE FROM members WHERE chunk IN (SELECT id FROM reconcile_chunks)')
db.execute('DELETE FROM chunks WHERE id IN (SELECT id FROM reconcile_chunks)')
db.execute('DELETE FROM messages WHERE generation=? AND source IN '
'(SELECT source FROM snapshot.reconcile)', (generation,))
for position, message in enumerate(messages):
if checkpoint and position % 100 == 0: checkpoint()
source = message['source']
Expand All @@ -87,6 +102,33 @@ def commit(self, generation, messages, chunks, vectors, job, checkpoint=None):
diagnostic_event('index.checkpoint.committed', task_id=job['id'], generation=generation, processed=job.get('processed'),
offset=job.get('offset'), chat_index=job.get('chat_index'), chunks=len(chunks), count=len(messages), committed=True)

@observed('index.reconcile')
def reconciliation(self, generation, snapshot, scopes):
"""记录快照中缺失的来源,并返回需要重建共享片段的现存相邻消息。"""
if not scopes:
return {'missing': 0, 'messages': []}
with self.connection() as db:
db.execute('ATTACH DATABASE ? AS snapshot', (str(snapshot),))
db.execute('CREATE TEMP TABLE reconcile_scopes(segment INTEGER PRIMARY KEY, username TEXT, start INTEGER, end INTEGER)')
db.executemany('INSERT INTO reconcile_scopes VALUES(?,?,?,?)',
[(s['segment'], s['username'], s['start'], s['end']) for s in scopes])
db.execute('DELETE FROM snapshot.reconcile')
db.execute('''INSERT OR IGNORE INTO snapshot.reconcile
SELECT m.source FROM messages m JOIN reconcile_scopes s
ON s.username=m.username AND m.created>=s.start AND m.created<=s.end
WHERE m.generation=? AND NOT EXISTS
(SELECT 1 FROM snapshot.messages frozen WHERE frozen.source=m.source)''', (generation,))
missing = db.execute('SELECT count(*) FROM snapshot.reconcile').fetchone()[0]
rows = db.execute('''SELECT DISTINCT current.body,current.username,current.created,current.source
FROM snapshot.reconcile removed
JOIN members old_member ON old_member.source=removed.source
JOIN chunks old_chunk ON old_chunk.id=old_member.chunk AND old_chunk.generation=?
JOIN members neighbor ON neighbor.chunk=old_chunk.id
JOIN messages current ON current.generation=old_chunk.generation AND current.source=neighbor.source
WHERE EXISTS (SELECT 1 FROM snapshot.messages frozen WHERE frozen.source=current.source)
ORDER BY current.username,current.created,current.source''', (generation,)).fetchall()
return {'missing': missing, 'messages': [json.loads(row['body']) for row in rows]}

def existing(self, generation, messages):
with self.connection() as db:
return {m['source'] for m in messages if (row := db.execute('SELECT body FROM messages WHERE generation=? AND source=?', (generation, m['source'])).fetchone()) and json.loads(row[0]) == m}
Expand Down
71 changes: 49 additions & 22 deletions src/wechat_decrypt_tool/local_search/service.py
Original file line number Diff line number Diff line change
Expand Up @@ -283,6 +283,39 @@ def check():
self.update(job, status='running', read_count=job['processed'], embedded_count=job['embedded'])
plan = await self.count_message_total(job, check)
segments = job.get('segments')
targets = segments or [{'username': username,
'start': job.get('read_starts', {}).get(username, job.get('read_start', job['start'])),
'end': job['end']} for username in cfg['usernames']]

async def encode_chunks(chunks, embedded_base):
vectors = []
last_embedding_update = time.monotonic()
if chunks:
self.update(job, stage='embedding')
for batch_start in range(0, len(chunks), 8):
check()
await self.yield_to_queries(check)
batch = chunks[batch_start:batch_start + 8]
strategy = cfg['device']
# 语音任务占用显卡时,本地索引主动让出。
try:
from ..voice_transcription import _VOICE_MODEL_ACTIVITY
if any(_VOICE_MODEL_ACTIVITY.values()):
strategy = 'cpu'
except ImportError:
pass
diagnostic_event('index.batch.started', index=batch_start, batch_size=len(batch), strategy=strategy,
reason_code='voice_priority' if strategy != cfg['device'] else 'configured')
values = await asyncio.to_thread(self.engine.encode, root, spec, [c['text'] for c in batch],
strategy, cfg['device_id'], False, cancelled)
diagnostic_event('index.batch.finished', index=batch_start, count=len(values),
actual_device=self.engine.status.get('actual_device'))
vectors.extend(values)
if len(vectors) == len(chunks) or time.monotonic() - last_embedding_update >= 0.25:
self.update(job, embedded_count=embedded_base + len(vectors))
last_embedding_update = time.monotonic()
return vectors

for position in range(job['chat_index'], len(segments) if segments is not None else len(cfg['usernames'])):
segment = segments[position] if segments is not None else None
username = segment['username'] if segment else cfg['usernames'][position]
Expand Down Expand Up @@ -327,28 +360,7 @@ def batch_size_changed(size):
changed = await asyncio.to_thread(index.affected_messages, job['generation'], [m for m in messages if m['source'] not in unchanged])
chunks = await asyncio.to_thread(make_chunks, changed, tokenizer)
diagnostic_event('index.page.organized', count=len(changed), unchanged=len(unchanged), chunks=len(chunks))
vectors = []
last_embedding_update = time.monotonic()
if chunks: self.update(job, stage='embedding')
for batch_start in range(0, len(chunks), 8):
check()
await self.yield_to_queries(check)
batch = chunks[batch_start:batch_start + 8]
strategy = cfg['device']
# 语音任务占用显卡时,本地索引主动让出。
try:
from ..voice_transcription import _VOICE_MODEL_ACTIVITY
if any(_VOICE_MODEL_ACTIVITY.values()): strategy = 'cpu'
except ImportError:
pass
diagnostic_event('index.batch.started', index=batch_start, batch_size=len(batch), strategy=strategy,
reason_code='voice_priority' if strategy!=cfg['device'] else 'configured')
values = await asyncio.to_thread(self.engine.encode, root, spec, [c['text'] for c in batch], strategy, cfg['device_id'], False, cancelled)
diagnostic_event('index.batch.finished', index=batch_start, count=len(values), actual_device=self.engine.status.get('actual_device'))
vectors.extend(values)
if len(vectors) == len(chunks) or time.monotonic() - last_embedding_update >= 0.25:
self.update(job, embedded_count=job['embedded'] + len(vectors))
last_embedding_update = time.monotonic()
vectors = await encode_chunks(chunks, job['embedded'])
check()
more = result.get('has_more', False)
next_job = {**job, 'chat_index': position if more else position + 1,
Expand All @@ -375,6 +387,21 @@ def batch_size_changed(size):
raise InferenceFailure('本轮消息清单与已保存数量不一致,已保留进度,请重试。', 'count_mismatch')
current = self.config(account)
if current.get('revision') != cfg.get('revision'): raise InferenceFailure('配置已更新', 'cancelled')
if not job.get('incremental'):
scopes = await asyncio.to_thread(plan.reconciliation_scopes, targets)
reconciliation = await asyncio.to_thread(index.reconciliation, job['generation'], plan.path, scopes)
if reconciliation['missing']:
changed = reconciliation['messages']
chunks = await asyncio.to_thread(make_chunks, changed, tokenizer)
vectors = await encode_chunks(chunks, job['embedded'])
next_job = {**job, 'embedded': job['embedded'] + len(chunks),
'removed': job.get('removed', 0) + reconciliation['missing']}
self.update(job, stage='saving')
await asyncio.to_thread(index.commit, job['generation'], changed, chunks, vectors,
next_job, check, plan.path)
job.update(next_job)
diagnostic_event('index.reconcile.finished', generation=job['generation'],
dropped_count=reconciliation['missing'], count=len(changed), chunks=len(chunks))
# 完成清理、范围约束和统计后才发布成功状态。
await asyncio.to_thread(index.prune, job['generation'], cfg['usernames'], job['start'], job['end'])
stats = await asyncio.to_thread(index.stats, job['generation'])
Expand Down
45 changes: 45 additions & 0 deletions tests/test_local_search.py
Original file line number Diff line number Diff line change
Expand Up @@ -466,6 +466,51 @@ def reader(account, username, start, end, offset):
asyncio.run(run())


def test_full_reconciliation_removes_missing_messages_and_rebuilds_neighbors(tmp_path, monkeypatch):
from tokenizers import Tokenizer, models
from wechat_decrypt_tool.local_search.catalog import model_dir

async def run():
rows = [message('kept', '继续保留'), message('deleted', '已经删除的秘密', timestamp=101)]
warning = ['']

def reader(account, username, start, end, offset):
return {'messages': [dict(m) for m in rows if start <= m['time'] <= end],
'name': username, 'has_more': False, 'warning': warning[0]}
service = LocalSearch(tmp_path/'state', tmp_path/'models', reader=reader, engine=FakeEngine())
root = model_dir(service.downloads.root, 'bge-small-zh')
root.mkdir(parents=True)
Tokenizer(models.WordLevel({'[UNK]': 0}, unk_token='[UNK]')).save(str(root/'tokenizer.json'))
monkeypatch.setattr(service.downloads, 'available', lambda _: True)
monkeypatch.setattr(service, 'enrichment_version', lambda _: [])
await service.configure('a', {'enabled': True, 'model': 'bge-small-zh', 'days': 0,
'start': 0, 'end': 1000, 'usernames': ['allowed']})
first = await service.build('a')
await service.jobs[first['id']]
assert first['status'] == 'done'
assert service.index('a').keyword(first['generation'], '秘密', ['allowed'])

rows.pop()
warning[0] = '数据源暂时不完整'
incomplete = await service.build('a', incremental=False)
await service.jobs[incomplete['id']]
assert incomplete['status'] == 'done'
assert service.index('a').keyword(first['generation'], '秘密', ['allowed'])

warning[0] = ''
reconciled = await service.build('a', incremental=False)
await service.jobs[reconciled['id']]
index = service.index('a')
assert reconciled['status'] == 'done' and reconciled['removed'] == 1
assert index.keyword(first['generation'], '秘密', ['allowed']) == []
assert index.stats(first['generation'])['messages'] == 1
assert {r['message']['source'] for r in index.search(first['generation'], [1., 0.], ['allowed'])} == {'kept'}
with index.connection() as db:
assert db.execute("SELECT count(*) FROM members WHERE source='deleted'").fetchone()[0] == 0
await service.stop()
asyncio.run(run())


@pytest.mark.parametrize('change,expected_mode,expected_start', [
({}, 'incremental', 9400),
({'start': 200}, 'incremental', 0),
Expand Down
4 changes: 3 additions & 1 deletion tests/test_local_search_totals.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,11 +21,13 @@ def snapshot(tmp_path, **values):

def test_snapshot_is_deduplicated_complete_and_immutable(tmp_path):
plan = snapshot(tmp_path)
plan.append(0, {'messages': [message('a'), message('a'), message('b')], 'has_more': True}, lambda: None)
plan.append(0, {'messages': [message('a'), message('a'), message('b')], 'has_more': True,
'warning': '第一页读取不完整'}, lambda: None)
with pytest.raises(ValueError):
plan.freeze(1)
plan.append(0, {'messages': [message('b'), message('c')], 'has_more': False}, lambda: None)
assert plan.freeze(1)['total'] == 3
assert plan.segment(0)['warning'] == '第一页读取不完整'
reopened = snapshot(tmp_path, processed=2, offset=2)
assert reopened.metadata()['total'] == 3
assert [m['source'] for m in reopened.page(0, 1, 100, lambda: None)['messages']] == ['b', 'c']
Expand Down