diff --git a/src/wechat_decrypt_tool/local_search/frozen.py b/src/wechat_decrypt_tool/local_search/frozen.py index 55c3b376..653e9647 100644 --- a/src/wechat_decrypt_tool/local_search/frozen.py +++ b/src/wechat_decrypt_tool/local_search/frozen.py @@ -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'], @@ -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))) @@ -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) diff --git a/src/wechat_decrypt_tool/local_search/index.py b/src/wechat_decrypt_tool/local_search/index.py index 1583c1f3..43da1c7d 100644 --- a/src/wechat_decrypt_tool/local_search/index.py +++ b/src/wechat_decrypt_tool/local_search/index.py @@ -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'] @@ -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} diff --git a/src/wechat_decrypt_tool/local_search/service.py b/src/wechat_decrypt_tool/local_search/service.py index b4022acb..61731eb9 100644 --- a/src/wechat_decrypt_tool/local_search/service.py +++ b/src/wechat_decrypt_tool/local_search/service.py @@ -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] @@ -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, @@ -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']) diff --git a/tests/test_local_search.py b/tests/test_local_search.py index 34cd1ca6..0cf2c4e9 100644 --- a/tests/test_local_search.py +++ b/tests/test_local_search.py @@ -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), diff --git a/tests/test_local_search_totals.py b/tests/test_local_search_totals.py index 00be0b85..82814796 100644 --- a/tests/test_local_search_totals.py +++ b/tests/test_local_search_totals.py @@ -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']