Add MCP search unit tests
This commit is contained in:
787
mcp_server.py
787
mcp_server.py
@@ -5,9 +5,10 @@ Based on FastMCP (stdio transport), reuses existing decryption.
|
||||
Runs on Windows Python (needs access to D:\ WeChat databases).
|
||||
"""
|
||||
|
||||
import os, sys, json, time, sqlite3, tempfile, struct, hashlib, atexit, re
|
||||
import hmac as hmac_mod
|
||||
from datetime import datetime
|
||||
import os, sys, json, time, sqlite3, tempfile, struct, hashlib, atexit, re
|
||||
import hmac as hmac_mod
|
||||
from contextlib import closing
|
||||
from datetime import datetime
|
||||
import xml.etree.ElementTree as ET
|
||||
from Crypto.Cipher import AES
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
@@ -219,9 +220,10 @@ atexit.register(_cache.cleanup)
|
||||
|
||||
_contact_names = None # {username: display_name}
|
||||
_contact_full = None # [{username, nick_name, remark}]
|
||||
_self_username = None
|
||||
_XML_UNSAFE_RE = re.compile(r'<!DOCTYPE|<!ENTITY', re.IGNORECASE)
|
||||
_XML_PARSE_MAX_LEN = 20000
|
||||
_self_username = None
|
||||
_XML_UNSAFE_RE = re.compile(r'<!DOCTYPE|<!ENTITY', re.IGNORECASE)
|
||||
_XML_PARSE_MAX_LEN = 20000
|
||||
_QUERY_LIMIT_MAX = 500
|
||||
|
||||
|
||||
def _load_contacts_from(db_path):
|
||||
@@ -568,12 +570,12 @@ MSG_DB_KEYS = sorted([
|
||||
])
|
||||
|
||||
|
||||
def _find_msg_table_for_user(username):
|
||||
"""在所有 message_N.db 中查找用户的消息表,返回 (db_path, table_name)"""
|
||||
table_hash = hashlib.md5(username.encode()).hexdigest()
|
||||
table_name = f"Msg_{table_hash}"
|
||||
if not _is_safe_msg_table_name(table_name):
|
||||
return None, None
|
||||
def _find_msg_table_for_user(username):
|
||||
"""在所有 message_N.db 中查找用户的消息表,返回 (db_path, table_name)"""
|
||||
table_hash = hashlib.md5(username.encode()).hexdigest()
|
||||
table_name = f"Msg_{table_hash}"
|
||||
if not _is_safe_msg_table_name(table_name):
|
||||
return None, None
|
||||
|
||||
for rel_key in MSG_DB_KEYS:
|
||||
path = _cache.get(rel_key)
|
||||
@@ -592,15 +594,54 @@ def _find_msg_table_for_user(username):
|
||||
pass
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
return None, None
|
||||
|
||||
return None, None
|
||||
|
||||
|
||||
def _find_msg_tables_for_user(username):
|
||||
"""返回用户在所有 message_N.db 中对应的消息表,按最新消息时间倒序排列。"""
|
||||
table_hash = hashlib.md5(username.encode()).hexdigest()
|
||||
table_name = f"Msg_{table_hash}"
|
||||
if not _is_safe_msg_table_name(table_name):
|
||||
return []
|
||||
|
||||
matches = []
|
||||
for rel_key in MSG_DB_KEYS:
|
||||
path = _cache.get(rel_key)
|
||||
if not path:
|
||||
continue
|
||||
conn = sqlite3.connect(path)
|
||||
try:
|
||||
exists = conn.execute(
|
||||
"SELECT 1 FROM sqlite_master WHERE type='table' AND name=?",
|
||||
(table_name,)
|
||||
).fetchone()
|
||||
if not exists:
|
||||
continue
|
||||
max_create_time = conn.execute(
|
||||
f"SELECT MAX(create_time) FROM [{table_name}]"
|
||||
).fetchone()[0] or 0
|
||||
matches.append({
|
||||
'db_path': path,
|
||||
'table_name': table_name,
|
||||
'max_create_time': max_create_time,
|
||||
})
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
matches.sort(key=lambda item: item['max_create_time'], reverse=True)
|
||||
return matches
|
||||
|
||||
|
||||
def _validate_pagination(limit, offset=0):
|
||||
if limit <= 0:
|
||||
raise ValueError("limit 必须大于 0")
|
||||
if offset < 0:
|
||||
raise ValueError("offset 不能小于 0")
|
||||
def _validate_pagination(limit, offset=0):
|
||||
if limit <= 0:
|
||||
raise ValueError("limit 必须大于 0")
|
||||
if limit > _QUERY_LIMIT_MAX:
|
||||
raise ValueError(f"limit 不能大于 {_QUERY_LIMIT_MAX}")
|
||||
if offset < 0:
|
||||
raise ValueError("offset 不能小于 0")
|
||||
|
||||
|
||||
def _parse_time_value(value, field_name, is_end=False):
|
||||
@@ -650,49 +691,54 @@ def _build_message_filters(start_ts=None, end_ts=None, keyword=''):
|
||||
return clauses, params
|
||||
|
||||
|
||||
def _query_messages(conn, table_name, start_ts=None, end_ts=None, keyword='', limit=20, offset=0):
|
||||
if not _is_safe_msg_table_name(table_name):
|
||||
raise ValueError(f'非法消息表名: {table_name}')
|
||||
|
||||
clauses, params = _build_message_filters(start_ts, end_ts, keyword)
|
||||
where_sql = f"WHERE {' AND '.join(clauses)}" if clauses else ''
|
||||
sql = f"""
|
||||
SELECT local_id, local_type, create_time, real_sender_id, message_content,
|
||||
WCDB_CT_message_content
|
||||
FROM [{table_name}]
|
||||
{where_sql}
|
||||
ORDER BY create_time DESC
|
||||
LIMIT ? OFFSET ?
|
||||
"""
|
||||
return conn.execute(sql, (*params, limit, offset)).fetchall()
|
||||
def _query_messages(conn, table_name, start_ts=None, end_ts=None, keyword='', limit=20, offset=0):
|
||||
if not _is_safe_msg_table_name(table_name):
|
||||
raise ValueError(f'非法消息表名: {table_name}')
|
||||
|
||||
clauses, params = _build_message_filters(start_ts, end_ts, keyword)
|
||||
where_sql = f"WHERE {' AND '.join(clauses)}" if clauses else ''
|
||||
sql = f"""
|
||||
SELECT local_id, local_type, create_time, real_sender_id, message_content,
|
||||
WCDB_CT_message_content
|
||||
FROM [{table_name}]
|
||||
{where_sql}
|
||||
ORDER BY create_time DESC
|
||||
"""
|
||||
if limit is None:
|
||||
return conn.execute(sql, params).fetchall()
|
||||
sql += "\n LIMIT ? OFFSET ?"
|
||||
return conn.execute(sql, (*params, limit, offset)).fetchall()
|
||||
|
||||
|
||||
def _resolve_chat_context(chat_name):
|
||||
username = resolve_username(chat_name)
|
||||
if not username:
|
||||
return None
|
||||
|
||||
names = get_contact_names()
|
||||
display_name = names.get(username, username)
|
||||
db_path, table_name = _find_msg_table_for_user(username)
|
||||
if not db_path:
|
||||
return {
|
||||
'query': chat_name,
|
||||
'username': username,
|
||||
'display_name': display_name,
|
||||
'db_path': None,
|
||||
'table_name': None,
|
||||
'is_group': '@chatroom' in username,
|
||||
}
|
||||
|
||||
return {
|
||||
'query': chat_name,
|
||||
'username': username,
|
||||
'display_name': display_name,
|
||||
'db_path': db_path,
|
||||
'table_name': table_name,
|
||||
'is_group': '@chatroom' in username,
|
||||
}
|
||||
def _resolve_chat_context(chat_name):
|
||||
username = resolve_username(chat_name)
|
||||
if not username:
|
||||
return None
|
||||
|
||||
names = get_contact_names()
|
||||
display_name = names.get(username, username)
|
||||
message_tables = _find_msg_tables_for_user(username)
|
||||
if not message_tables:
|
||||
return {
|
||||
'query': chat_name,
|
||||
'username': username,
|
||||
'display_name': display_name,
|
||||
'db_path': None,
|
||||
'table_name': None,
|
||||
'message_tables': [],
|
||||
'is_group': '@chatroom' in username,
|
||||
}
|
||||
|
||||
primary = message_tables[0]
|
||||
return {
|
||||
'query': chat_name,
|
||||
'username': username,
|
||||
'display_name': display_name,
|
||||
'db_path': primary['db_path'],
|
||||
'table_name': primary['table_name'],
|
||||
'message_tables': message_tables,
|
||||
'is_group': '@chatroom' in username,
|
||||
}
|
||||
|
||||
|
||||
def _resolve_chat_contexts(chat_names):
|
||||
@@ -713,11 +759,11 @@ def _resolve_chat_contexts(chat_names):
|
||||
if not ctx:
|
||||
unresolved.append(name)
|
||||
continue
|
||||
if not ctx['db_path']:
|
||||
missing_tables.append(ctx['display_name'])
|
||||
continue
|
||||
if ctx['username'] in seen:
|
||||
continue
|
||||
if not ctx['message_tables']:
|
||||
missing_tables.append(ctx['display_name'])
|
||||
continue
|
||||
if ctx['username'] in seen:
|
||||
continue
|
||||
seen.add(ctx['username'])
|
||||
resolved.append(ctx)
|
||||
|
||||
@@ -743,31 +789,20 @@ def _normalize_chat_names(chat_name):
|
||||
return [value] if value else []
|
||||
|
||||
|
||||
def _format_history_lines(rows, username, display_name, is_group, names, id_to_username):
|
||||
lines = []
|
||||
for local_id, local_type, create_time, real_sender_id, content, ct in reversed(rows):
|
||||
time_str = datetime.fromtimestamp(create_time).strftime('%Y-%m-%d %H:%M')
|
||||
content = _decompress_content(content, ct)
|
||||
if content is None:
|
||||
content = '(无法解压)'
|
||||
|
||||
sender, text = _format_message_text(
|
||||
local_id, local_type, content, is_group, username, display_name, names
|
||||
)
|
||||
if text and len(text) > 500:
|
||||
text = text[:500] + '...'
|
||||
|
||||
sender_label = _resolve_sender_label(
|
||||
real_sender_id, sender, is_group, username, display_name, names, id_to_username
|
||||
)
|
||||
if sender_label:
|
||||
lines.append(f'[{time_str}] {sender_label}: {text}')
|
||||
else:
|
||||
lines.append(f'[{time_str}] {text}')
|
||||
return lines
|
||||
def _format_history_lines(rows, username, display_name, is_group, names, id_to_username):
|
||||
lines = []
|
||||
ctx = {
|
||||
'username': username,
|
||||
'display_name': display_name,
|
||||
'is_group': is_group,
|
||||
}
|
||||
for row in reversed(rows):
|
||||
_, line = _build_history_line(row, ctx, names, id_to_username)
|
||||
lines.append(line)
|
||||
return lines
|
||||
|
||||
|
||||
def _build_search_entry(row, ctx, names, id_to_username):
|
||||
def _build_search_entry(row, ctx, names, id_to_username):
|
||||
local_id, local_type, create_time, real_sender_id, content, ct = row
|
||||
content = _decompress_content(content, ct)
|
||||
if content is None:
|
||||
@@ -792,8 +827,291 @@ def _build_search_entry(row, ctx, names, id_to_username):
|
||||
entry = f"[{time_str}] [{ctx['display_name']}]"
|
||||
if sender_label:
|
||||
entry += f" {sender_label}:"
|
||||
entry += f" {text}"
|
||||
return create_time, entry
|
||||
entry += f" {text}"
|
||||
return create_time, entry
|
||||
|
||||
|
||||
def _build_history_line(row, ctx, names, id_to_username):
|
||||
local_id, local_type, create_time, real_sender_id, content, ct = row
|
||||
time_str = datetime.fromtimestamp(create_time).strftime('%Y-%m-%d %H:%M')
|
||||
content = _decompress_content(content, ct)
|
||||
if content is None:
|
||||
content = '(无法解压)'
|
||||
|
||||
sender, text = _format_message_text(
|
||||
local_id, local_type, content, ctx['is_group'], ctx['username'], ctx['display_name'], names
|
||||
)
|
||||
if text and len(text) > 500:
|
||||
text = text[:500] + '...'
|
||||
|
||||
sender_label = _resolve_sender_label(
|
||||
real_sender_id, sender, ctx['is_group'], ctx['username'], ctx['display_name'], names, id_to_username
|
||||
)
|
||||
if sender_label:
|
||||
return create_time, f'[{time_str}] {sender_label}: {text}'
|
||||
return create_time, f'[{time_str}] {text}'
|
||||
|
||||
|
||||
def _get_chat_message_tables(ctx):
|
||||
if ctx.get('message_tables'):
|
||||
return ctx['message_tables']
|
||||
if ctx.get('db_path') and ctx.get('table_name'):
|
||||
return [{'db_path': ctx['db_path'], 'table_name': ctx['table_name']}]
|
||||
return []
|
||||
|
||||
|
||||
def _iter_table_contexts(ctx):
|
||||
for table in _get_chat_message_tables(ctx):
|
||||
yield {
|
||||
'query': ctx['query'],
|
||||
'username': ctx['username'],
|
||||
'display_name': ctx['display_name'],
|
||||
'db_path': table['db_path'],
|
||||
'table_name': table['table_name'],
|
||||
'is_group': ctx['is_group'],
|
||||
}
|
||||
|
||||
|
||||
def _collect_chat_history_lines(ctx, names, start_ts=None, end_ts=None, limit=20, offset=0):
|
||||
collected = []
|
||||
failures = []
|
||||
|
||||
for table_ctx in _iter_table_contexts(ctx):
|
||||
try:
|
||||
with closing(sqlite3.connect(table_ctx['db_path'])) as conn:
|
||||
id_to_username = _load_name2id_maps(conn)
|
||||
rows = _query_messages(
|
||||
conn,
|
||||
table_ctx['table_name'],
|
||||
start_ts=start_ts,
|
||||
end_ts=end_ts,
|
||||
limit=None,
|
||||
)
|
||||
for row in rows:
|
||||
collected.append(_build_history_line(row, table_ctx, names, id_to_username))
|
||||
except Exception as e:
|
||||
failures.append(f"{table_ctx['db_path']}: {e}")
|
||||
|
||||
ordered = sorted(collected, key=lambda item: item[0], reverse=True)
|
||||
paged = ordered[offset:offset + limit]
|
||||
paged.sort(key=lambda item: item[0])
|
||||
return [line for _, line in paged], failures
|
||||
|
||||
|
||||
def _collect_chat_search_entries(ctx, names, keyword, start_ts=None, end_ts=None):
|
||||
collected = []
|
||||
failures = []
|
||||
contexts_by_db = {}
|
||||
for table_ctx in _iter_table_contexts(ctx):
|
||||
contexts_by_db.setdefault(table_ctx['db_path'], []).append(table_ctx)
|
||||
|
||||
for db_path, db_contexts in contexts_by_db.items():
|
||||
try:
|
||||
with closing(sqlite3.connect(db_path)) as conn:
|
||||
db_entries, db_failures = _collect_search_entries(
|
||||
conn,
|
||||
db_contexts,
|
||||
names,
|
||||
keyword,
|
||||
start_ts=start_ts,
|
||||
end_ts=end_ts,
|
||||
)
|
||||
collected.extend(db_entries)
|
||||
failures.extend(db_failures)
|
||||
except Exception as e:
|
||||
failures.extend(f"{table_ctx['display_name']}: {e}" for table_ctx in db_contexts)
|
||||
|
||||
return collected, failures
|
||||
|
||||
|
||||
def _load_search_contexts_from_db(conn, db_path, names):
|
||||
tables = conn.execute(
|
||||
"SELECT name FROM sqlite_master WHERE type='table' AND name LIKE 'Msg_%'"
|
||||
).fetchall()
|
||||
|
||||
table_to_username = {}
|
||||
try:
|
||||
for (user_name,) in conn.execute("SELECT user_name FROM Name2Id").fetchall():
|
||||
if not user_name:
|
||||
continue
|
||||
table_hash = hashlib.md5(user_name.encode()).hexdigest()
|
||||
table_to_username[f"Msg_{table_hash}"] = user_name
|
||||
except sqlite3.Error:
|
||||
pass
|
||||
|
||||
contexts = []
|
||||
for (table_name,) in tables:
|
||||
username = table_to_username.get(table_name, '')
|
||||
display_name = names.get(username, username) if username else table_name
|
||||
contexts.append({
|
||||
'query': display_name,
|
||||
'username': username,
|
||||
'display_name': display_name,
|
||||
'db_path': db_path,
|
||||
'table_name': table_name,
|
||||
'is_group': '@chatroom' in username,
|
||||
})
|
||||
return contexts
|
||||
|
||||
|
||||
def _collect_search_entries(conn, contexts, names, keyword, start_ts=None, end_ts=None):
|
||||
collected = []
|
||||
failures = []
|
||||
id_to_username = _load_name2id_maps(conn)
|
||||
|
||||
for ctx in contexts:
|
||||
try:
|
||||
rows = _query_messages(
|
||||
conn,
|
||||
ctx['table_name'],
|
||||
start_ts=start_ts,
|
||||
end_ts=end_ts,
|
||||
keyword=keyword,
|
||||
limit=None,
|
||||
)
|
||||
for row in rows:
|
||||
formatted = _build_search_entry(row, ctx, names, id_to_username)
|
||||
if formatted:
|
||||
collected.append(formatted)
|
||||
except Exception as e:
|
||||
failures.append(f"{ctx['display_name']}: {e}")
|
||||
|
||||
return collected, failures
|
||||
|
||||
|
||||
def _page_search_entries(entries, limit, offset):
|
||||
ordered = sorted(entries, key=lambda x: x[0], reverse=True)
|
||||
paged = ordered[offset:offset + limit]
|
||||
paged.sort(key=lambda x: x[0])
|
||||
return paged
|
||||
|
||||
|
||||
def _search_single_chat(ctx, keyword, start_ts, end_ts, start_time, end_time, limit, offset):
|
||||
names = get_contact_names()
|
||||
|
||||
entries, failures = _collect_chat_search_entries(
|
||||
ctx,
|
||||
names,
|
||||
keyword,
|
||||
start_ts=start_ts,
|
||||
end_ts=end_ts,
|
||||
)
|
||||
|
||||
paged = _page_search_entries(entries, limit, offset)
|
||||
|
||||
if not paged:
|
||||
if failures:
|
||||
return "查询失败: " + ";".join(failures)
|
||||
return f"未在 {ctx['display_name']} 中找到包含 \"{keyword}\" 的消息"
|
||||
|
||||
header = f"在 {ctx['display_name']} 中搜索 \"{keyword}\" 找到 {len(paged)} 条结果(offset={offset}, limit={limit})"
|
||||
if start_time or end_time:
|
||||
header += f"\n时间范围: {start_time or '最早'} ~ {end_time or '最新'}"
|
||||
if failures:
|
||||
header += "\n查询失败: " + ";".join(failures)
|
||||
return header + ":\n\n" + "\n\n".join(item[1] for item in paged)
|
||||
|
||||
|
||||
def _search_multiple_chats(chat_names, keyword, start_ts, end_ts, start_time, end_time, limit, offset):
|
||||
try:
|
||||
resolved_contexts, unresolved, missing_tables = _resolve_chat_contexts(chat_names)
|
||||
except ValueError as e:
|
||||
return f"错误: {e}"
|
||||
|
||||
if not resolved_contexts:
|
||||
details = []
|
||||
if unresolved:
|
||||
details.append("未找到联系人: " + "、".join(unresolved))
|
||||
if missing_tables:
|
||||
details.append("无消息表: " + "、".join(missing_tables))
|
||||
suffix = f"\n{chr(10).join(details)}" if details else ""
|
||||
return f"错误: 没有可查询的聊天对象{suffix}"
|
||||
|
||||
names = get_contact_names()
|
||||
collected = []
|
||||
failures = []
|
||||
for ctx in resolved_contexts:
|
||||
chat_entries, chat_failures = _collect_chat_search_entries(
|
||||
ctx,
|
||||
names,
|
||||
keyword,
|
||||
start_ts=start_ts,
|
||||
end_ts=end_ts,
|
||||
)
|
||||
collected.extend(chat_entries)
|
||||
failures.extend(chat_failures)
|
||||
|
||||
paged = _page_search_entries(collected, limit, offset)
|
||||
|
||||
notes = []
|
||||
if unresolved:
|
||||
notes.append("未找到联系人: " + "、".join(unresolved))
|
||||
if missing_tables:
|
||||
notes.append("无消息表: " + "、".join(missing_tables))
|
||||
if failures:
|
||||
notes.append("查询失败: " + ";".join(failures))
|
||||
|
||||
if not paged:
|
||||
header = f"在 {len(resolved_contexts)} 个聊天对象中未找到包含 \"{keyword}\" 的消息"
|
||||
if start_time or end_time:
|
||||
header += f"\n时间范围: {start_time or '最早'} ~ {end_time or '最新'}"
|
||||
if notes:
|
||||
header += "\n" + "\n".join(notes)
|
||||
return header
|
||||
|
||||
header = (
|
||||
f"在 {len(resolved_contexts)} 个聊天对象中搜索 \"{keyword}\" 找到 {len(paged)} 条结果"
|
||||
f"(offset={offset}, limit={limit})"
|
||||
)
|
||||
if start_time or end_time:
|
||||
header += f"\n时间范围: {start_time or '最早'} ~ {end_time or '最新'}"
|
||||
if notes:
|
||||
header += "\n" + "\n".join(notes)
|
||||
return header + ":\n\n" + "\n\n".join(item[1] for item in paged)
|
||||
|
||||
|
||||
def _search_all_messages(keyword, start_ts, end_ts, start_time, end_time, limit, offset):
|
||||
names = get_contact_names()
|
||||
collected = []
|
||||
failures = []
|
||||
|
||||
for rel_key in MSG_DB_KEYS:
|
||||
path = _cache.get(rel_key)
|
||||
if not path:
|
||||
continue
|
||||
|
||||
try:
|
||||
with closing(sqlite3.connect(path)) as conn:
|
||||
contexts = _load_search_contexts_from_db(conn, path, names)
|
||||
db_entries, db_failures = _collect_search_entries(
|
||||
conn,
|
||||
contexts,
|
||||
names,
|
||||
keyword,
|
||||
start_ts=start_ts,
|
||||
end_ts=end_ts,
|
||||
)
|
||||
collected.extend(db_entries)
|
||||
failures.extend(db_failures)
|
||||
except Exception as e:
|
||||
failures.append(f"{rel_key}: {e}")
|
||||
|
||||
paged = _page_search_entries(collected, limit, offset)
|
||||
|
||||
if not paged:
|
||||
header = f"未找到包含 \"{keyword}\" 的消息"
|
||||
if start_time or end_time:
|
||||
header += f"\n时间范围: {start_time or '最早'} ~ {end_time or '最新'}"
|
||||
if failures:
|
||||
header += "\n查询失败: " + ";".join(failures)
|
||||
return header
|
||||
|
||||
header = f"搜索 \"{keyword}\" 找到 {len(paged)} 条结果(offset={offset}, limit={limit})"
|
||||
if start_time or end_time:
|
||||
header += f"\n时间范围: {start_time or '最早'} ~ {end_time or '最新'}"
|
||||
if failures:
|
||||
header += "\n查询失败: " + ";".join(failures)
|
||||
return header + ":\n\n" + "\n\n".join(item[1] for item in paged)
|
||||
|
||||
|
||||
# ============ MCP Server ============
|
||||
@@ -864,15 +1182,15 @@ def get_recent_sessions(limit: int = 20) -> str:
|
||||
|
||||
|
||||
@mcp.tool()
|
||||
def get_chat_history(chat_name: str, limit: int = 50, offset: int = 0, start_time: str = "", end_time: str = "") -> str:
|
||||
"""获取指定聊天的消息记录。
|
||||
|
||||
Args:
|
||||
chat_name: 聊天对象的名字、备注名或wxid,自动模糊匹配
|
||||
limit: 返回的消息数量,默认50
|
||||
offset: 分页偏移量,默认0
|
||||
start_time: 起始时间,支持 YYYY-MM-DD / YYYY-MM-DD HH:MM / YYYY-MM-DD HH:MM:SS
|
||||
end_time: 结束时间,支持 YYYY-MM-DD / YYYY-MM-DD HH:MM / YYYY-MM-DD HH:MM:SS
|
||||
def get_chat_history(chat_name: str, limit: int = 50, offset: int = 0, start_time: str = "", end_time: str = "") -> str:
|
||||
"""获取指定聊天的消息记录。
|
||||
|
||||
Args:
|
||||
chat_name: 聊天对象的名字、备注名或wxid,自动模糊匹配
|
||||
limit: 返回的消息数量,默认50,最大500
|
||||
offset: 分页偏移量,默认0
|
||||
start_time: 起始时间,支持 YYYY-MM-DD / YYYY-MM-DD HH:MM / YYYY-MM-DD HH:MM:SS
|
||||
end_time: 结束时间,支持 YYYY-MM-DD / YYYY-MM-DD HH:MM / YYYY-MM-DD HH:MM:SS
|
||||
"""
|
||||
try:
|
||||
_validate_pagination(limit, offset)
|
||||
@@ -886,41 +1204,29 @@ def get_chat_history(chat_name: str, limit: int = 50, offset: int = 0, start_tim
|
||||
if not ctx['db_path']:
|
||||
return f"找不到 {ctx['display_name']} 的消息记录(可能在未解密的DB中或无消息)"
|
||||
|
||||
names = get_contact_names()
|
||||
conn = sqlite3.connect(ctx['db_path'])
|
||||
try:
|
||||
id_to_username = _load_name2id_maps(conn)
|
||||
rows = _query_messages(
|
||||
conn,
|
||||
ctx['table_name'],
|
||||
start_ts=start_ts,
|
||||
end_ts=end_ts,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
except Exception as e:
|
||||
conn.close()
|
||||
return f"查询失败: {e}"
|
||||
conn.close()
|
||||
|
||||
if not rows:
|
||||
return f"{ctx['display_name']} 无消息记录"
|
||||
|
||||
lines = _format_history_lines(
|
||||
rows,
|
||||
ctx['username'],
|
||||
ctx['display_name'],
|
||||
ctx['is_group'],
|
||||
names,
|
||||
id_to_username,
|
||||
)
|
||||
|
||||
header = f"{ctx['display_name']} 的消息记录(返回 {len(lines)} 条,offset={offset}, limit={limit})"
|
||||
if ctx['is_group']:
|
||||
header += " [群聊]"
|
||||
if start_time or end_time:
|
||||
header += f"\n时间范围: {start_time or '最早'} ~ {end_time or '最新'}"
|
||||
return header + ":\n\n" + "\n".join(lines)
|
||||
names = get_contact_names()
|
||||
lines, failures = _collect_chat_history_lines(
|
||||
ctx,
|
||||
names,
|
||||
start_ts=start_ts,
|
||||
end_ts=end_ts,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
|
||||
if not lines:
|
||||
if failures:
|
||||
return "查询失败: " + ";".join(failures)
|
||||
return f"{ctx['display_name']} 无消息记录"
|
||||
|
||||
header = f"{ctx['display_name']} 的消息记录(返回 {len(lines)} 条,offset={offset}, limit={limit})"
|
||||
if ctx['is_group']:
|
||||
header += " [群聊]"
|
||||
if start_time or end_time:
|
||||
header += f"\n时间范围: {start_time or '最早'} ~ {end_time or '最新'}"
|
||||
if failures:
|
||||
header += "\n查询失败: " + ";".join(failures)
|
||||
return header + ":\n\n" + "\n".join(lines)
|
||||
|
||||
|
||||
@mcp.tool()
|
||||
@@ -939,7 +1245,7 @@ def search_messages(
|
||||
chat_name: 聊天对象名称,可为空、单个字符串或字符串列表
|
||||
start_time: 起始时间,可为空
|
||||
end_time: 结束时间,可为空
|
||||
limit: 返回的结果数量,默认20
|
||||
limit: 返回的结果数量,默认20,最大500
|
||||
offset: 分页偏移量,默认0
|
||||
"""
|
||||
if not keyword or len(keyword) < 1:
|
||||
@@ -959,193 +1265,38 @@ def search_messages(
|
||||
return f"找不到聊天对象: {chat_names[0]}\n提示: 可以用 get_contacts(query='{chat_names[0]}') 搜索联系人"
|
||||
if not ctx['db_path']:
|
||||
return f"找不到 {ctx['display_name']} 的消息记录(可能在未解密的DB中或无消息)"
|
||||
|
||||
names = get_contact_names()
|
||||
conn = sqlite3.connect(ctx['db_path'])
|
||||
try:
|
||||
id_to_username = _load_name2id_maps(conn)
|
||||
rows = _query_messages(
|
||||
conn,
|
||||
ctx['table_name'],
|
||||
start_ts=start_ts,
|
||||
end_ts=end_ts,
|
||||
keyword=keyword,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
except Exception as e:
|
||||
conn.close()
|
||||
return f"查询失败: {e}"
|
||||
conn.close()
|
||||
|
||||
if not rows:
|
||||
return f"未在 {ctx['display_name']} 中找到包含 \"{keyword}\" 的消息"
|
||||
|
||||
entries = []
|
||||
for row in rows:
|
||||
formatted = _build_search_entry(row, ctx, names, id_to_username)
|
||||
if formatted:
|
||||
entries.append(formatted)
|
||||
|
||||
if not entries:
|
||||
return f"未在 {ctx['display_name']} 中找到包含 \"{keyword}\" 的可读消息"
|
||||
|
||||
entries.sort(key=lambda x: x[0])
|
||||
header = f"在 {ctx['display_name']} 中搜索 \"{keyword}\" 找到 {len(entries)} 条结果(offset={offset}, limit={limit})"
|
||||
if start_time or end_time:
|
||||
header += f"\n时间范围: {start_time or '最早'} ~ {end_time or '最新'}"
|
||||
return header + ":\n\n" + "\n\n".join(item[1] for item in entries)
|
||||
return _search_single_chat(
|
||||
ctx,
|
||||
keyword,
|
||||
start_ts,
|
||||
end_ts,
|
||||
start_time,
|
||||
end_time,
|
||||
limit,
|
||||
offset,
|
||||
)
|
||||
|
||||
if len(chat_names) > 1:
|
||||
try:
|
||||
resolved_contexts, unresolved, missing_tables = _resolve_chat_contexts(chat_names)
|
||||
except ValueError as e:
|
||||
return f"错误: {e}"
|
||||
|
||||
if not resolved_contexts:
|
||||
details = []
|
||||
if unresolved:
|
||||
details.append("未找到联系人: " + "、".join(unresolved))
|
||||
if missing_tables:
|
||||
details.append("无消息表: " + "、".join(missing_tables))
|
||||
suffix = f"\n{chr(10).join(details)}" if details else ""
|
||||
return f"错误: 没有可查询的聊天对象{suffix}"
|
||||
|
||||
names = get_contact_names()
|
||||
collected = []
|
||||
failures = []
|
||||
per_chat_limit = limit + offset
|
||||
|
||||
for ctx in resolved_contexts:
|
||||
conn = sqlite3.connect(ctx['db_path'])
|
||||
try:
|
||||
id_to_username = _load_name2id_maps(conn)
|
||||
rows = _query_messages(
|
||||
conn,
|
||||
ctx['table_name'],
|
||||
start_ts=start_ts,
|
||||
end_ts=end_ts,
|
||||
keyword=keyword,
|
||||
limit=per_chat_limit,
|
||||
offset=0,
|
||||
)
|
||||
for row in rows:
|
||||
formatted = _build_search_entry(row, ctx, names, id_to_username)
|
||||
if formatted:
|
||||
collected.append(formatted)
|
||||
except Exception as e:
|
||||
failures.append(f"{ctx['display_name']}: {e}")
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
collected.sort(key=lambda x: x[0], reverse=True)
|
||||
paged = collected[offset:offset + limit]
|
||||
|
||||
notes = []
|
||||
if unresolved:
|
||||
notes.append("未找到联系人: " + "、".join(unresolved))
|
||||
if missing_tables:
|
||||
notes.append("无消息表: " + "、".join(missing_tables))
|
||||
if failures:
|
||||
notes.append("查询失败: " + ";".join(failures))
|
||||
|
||||
if not paged:
|
||||
header = f"在 {len(resolved_contexts)} 个聊天对象中未找到包含 \"{keyword}\" 的消息"
|
||||
if start_time or end_time:
|
||||
header += f"\n时间范围: {start_time or '最早'} ~ {end_time or '最新'}"
|
||||
if notes:
|
||||
header += "\n" + "\n".join(notes)
|
||||
return header
|
||||
|
||||
header = (
|
||||
f"在 {len(resolved_contexts)} 个聊天对象中搜索 \"{keyword}\" 找到 {len(paged)} 条结果"
|
||||
f"(offset={offset}, limit={limit})"
|
||||
return _search_multiple_chats(
|
||||
chat_names,
|
||||
keyword,
|
||||
start_ts,
|
||||
end_ts,
|
||||
start_time,
|
||||
end_time,
|
||||
limit,
|
||||
offset,
|
||||
)
|
||||
if start_time or end_time:
|
||||
header += f"\n时间范围: {start_time or '最早'} ~ {end_time or '最新'}"
|
||||
if notes:
|
||||
header += "\n" + "\n".join(notes)
|
||||
return header + ":\n\n" + "\n\n".join(item[1] for item in paged)
|
||||
|
||||
names = get_contact_names()
|
||||
results = []
|
||||
max_results = limit + offset
|
||||
|
||||
for rel_key in MSG_DB_KEYS:
|
||||
if len(results) >= max_results:
|
||||
break
|
||||
|
||||
path = _cache.get(rel_key)
|
||||
if not path:
|
||||
continue
|
||||
|
||||
conn = sqlite3.connect(path)
|
||||
try:
|
||||
# 获取所有 Msg_ 表
|
||||
tables = conn.execute(
|
||||
"SELECT name FROM sqlite_master WHERE type='table' AND name LIKE 'Msg_%'"
|
||||
).fetchall()
|
||||
|
||||
# 获取 Name2Id 映射(hash -> username 反查)
|
||||
name2id = {}
|
||||
try:
|
||||
for r in conn.execute("SELECT user_name FROM Name2Id").fetchall():
|
||||
h = hashlib.md5(r[0].encode()).hexdigest()
|
||||
name2id[f"Msg_{h}"] = r[0]
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
for (tname,) in tables:
|
||||
if len(results) >= max_results:
|
||||
break
|
||||
username = name2id.get(tname, '')
|
||||
is_group = '@chatroom' in username
|
||||
display = names.get(username, username) if username else tname
|
||||
|
||||
try:
|
||||
clauses, params = _build_message_filters(start_ts, end_ts, keyword)
|
||||
where_sql = f"WHERE {' AND '.join(clauses)}" if clauses else ''
|
||||
rows = conn.execute(f"""
|
||||
SELECT local_type, create_time, message_content,
|
||||
WCDB_CT_message_content
|
||||
FROM [{tname}]
|
||||
{where_sql}
|
||||
ORDER BY create_time DESC
|
||||
LIMIT ? OFFSET ?
|
||||
""", (*params, max_results - len(results), 0)).fetchall()
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
for local_type, ts, content, ct in rows:
|
||||
content = _decompress_content(content, ct)
|
||||
if content is None:
|
||||
continue
|
||||
sender, text = _parse_message_content(content, local_type, is_group)
|
||||
time_str = datetime.fromtimestamp(ts).strftime('%Y-%m-%d %H:%M')
|
||||
sender_name = ''
|
||||
if is_group and sender:
|
||||
sender_name = names.get(sender, sender)
|
||||
|
||||
entry = f"[{time_str}] [{display}]"
|
||||
if sender_name:
|
||||
entry += f" {sender_name}:"
|
||||
entry += f" {text}"
|
||||
if len(entry) > 300:
|
||||
entry = entry[:300] + "..."
|
||||
results.append((ts, entry))
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
results.sort(key=lambda x: x[0], reverse=True)
|
||||
entries = [r[1] for r in results[offset:offset + limit]]
|
||||
|
||||
if not entries:
|
||||
return f"未找到包含 \"{keyword}\" 的消息"
|
||||
|
||||
header = f"搜索 \"{keyword}\" 找到 {len(entries)} 条结果(offset={offset}, limit={limit})"
|
||||
if start_time or end_time:
|
||||
header += f"\n时间范围: {start_time or '最早'} ~ {end_time or '最新'}"
|
||||
return header + ":\n\n" + "\n\n".join(entries)
|
||||
return _search_all_messages(
|
||||
keyword,
|
||||
start_ts,
|
||||
end_ts,
|
||||
start_time,
|
||||
end_time,
|
||||
limit,
|
||||
offset,
|
||||
)
|
||||
|
||||
@mcp.tool()
|
||||
def get_contacts(query: str = "", limit: int = 50) -> str:
|
||||
|
||||
Reference in New Issue
Block a user