feat(mcp): get_chat_history 加 msg_types 按类型过滤
LLM 用 \`get_chat_history\` 查"和 X 的所有图片消息"时, 只能拉 50 条 混合消息再客户端过滤 —— 大部分 token 浪费在不需要的文本上。同样 "只看转账记录" / "只看语音" 的场景, 没有原生过滤手段。 \`get_chat_history\` 加一个可选 kwarg \`msg_types: list[str] | None = None\`: - 接受 \`['text', 'image', 'voice', 'video', 'file', 'emoji', 'location', 'namecard', 'voip', 'system']\` 子集 - \`'file'\` 是 alias → 'app' (WeChat 把文件归到 \`local_type=49\`, 俗称 file) - 输入大小写不敏感, 自动 strip - 未知类型立即报错并列出可选值 (不偷偷过滤合法部分) - None 或 \`[]\` 表示不过滤, 完全等价于旧行为 (向后兼容) 实现上拆 3 件: 1. \`_MSG_TYPE_MAP\` 常量 (字符串 → \`local_type\` 整数列表) 2. \`_resolve_msg_types()\` helper 做输入校验 + 翻译 3. \`_build_message_filters\` / \`_query_messages\` / \`_collect_chat_history_lines\` 链路加 \`type_filter=None\` 透传, SQL 注入 \`local_type IN (?,?,...)\` clause \`tests/test_msg_types_filter.py\` 12 个 case: - None / 空 → 不过滤 - 单类型 / 多类型解析 - \`file\` alias → app - 大小写 + strip 不敏感 - 未知类型报错且不放过合法的 - SQL 生成: 无过滤时 clauses 不含 \`local_type\`, 单类型生成 \`IN (?)\`, 多类型生成 \`IN (?,?,?)\` - 与 time / keyword 组合时 param 顺序正确 全量 \`pytest tests/\` 212/212 通过。 新参数默认 None, **既有调用方零修改**。 类型映射表 (\`_MSG_TYPE_MAP\`) 命名是有立场的判断 (比如 \`'app'\` 这一 桶实际混了文件 / 分享卡 / 小程序 / 转账 / 引用回复), 如果维护者 不同意具体 label 或想拆细, 改 dict 就行, 不影响接口。 与 #103 (\`_pagination_hint\`) 触碰同一文件 \`mcp_server.py\`, 后合的 rebase 即可, 无逻辑冲突。
This commit is contained in:
@@ -1160,7 +1160,42 @@ def _pagination_hint(count, limit, offset):
|
|||||||
return ""
|
return ""
|
||||||
|
|
||||||
|
|
||||||
def _build_message_filters(start_ts=None, end_ts=None, keyword=''):
|
_MSG_TYPE_MAP = {
|
||||||
|
'text': [1],
|
||||||
|
'image': [3],
|
||||||
|
'voice': [34],
|
||||||
|
'namecard': [42],
|
||||||
|
'video': [43],
|
||||||
|
'emoji': [47],
|
||||||
|
'location': [48],
|
||||||
|
'app': [49],
|
||||||
|
'voip': [50],
|
||||||
|
'system': [10000],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_msg_types(msg_types):
|
||||||
|
"""把 ['text', 'image'] 风格的输入翻成 local_type 整数列表。
|
||||||
|
|
||||||
|
返回 (type_filter_list, error_msg); 任一项无效返回 (None, error)。
|
||||||
|
None / 空列表表示不过滤。
|
||||||
|
"""
|
||||||
|
if not msg_types:
|
||||||
|
return None, None
|
||||||
|
type_filter = []
|
||||||
|
for t in msg_types:
|
||||||
|
key = t.strip().lower()
|
||||||
|
if key == 'file':
|
||||||
|
key = 'app' # 'file' 是常见叫法; WeChat 把文件归到 type=49 (app message)
|
||||||
|
if key not in _MSG_TYPE_MAP:
|
||||||
|
return None, (
|
||||||
|
f"未知消息类型 \"{t}\"。可选: " + ", ".join(sorted(_MSG_TYPE_MAP))
|
||||||
|
)
|
||||||
|
type_filter.extend(_MSG_TYPE_MAP[key])
|
||||||
|
return type_filter, None
|
||||||
|
|
||||||
|
|
||||||
|
def _build_message_filters(start_ts=None, end_ts=None, keyword='', type_filter=None):
|
||||||
clauses = []
|
clauses = []
|
||||||
params = []
|
params = []
|
||||||
if start_ts is not None:
|
if start_ts is not None:
|
||||||
@@ -1172,14 +1207,18 @@ def _build_message_filters(start_ts=None, end_ts=None, keyword=''):
|
|||||||
if keyword:
|
if keyword:
|
||||||
clauses.append('message_content LIKE ?')
|
clauses.append('message_content LIKE ?')
|
||||||
params.append(f'%{keyword}%')
|
params.append(f'%{keyword}%')
|
||||||
|
if type_filter:
|
||||||
|
placeholders = ','.join('?' * len(type_filter))
|
||||||
|
clauses.append(f'local_type IN ({placeholders})')
|
||||||
|
params.extend(type_filter)
|
||||||
return clauses, params
|
return clauses, params
|
||||||
|
|
||||||
|
|
||||||
def _query_messages(conn, table_name, start_ts=None, end_ts=None, keyword='', limit=20, offset=0, oldest_first=False):
|
def _query_messages(conn, table_name, start_ts=None, end_ts=None, keyword='', limit=20, offset=0, oldest_first=False, type_filter=None):
|
||||||
if not _is_safe_msg_table_name(table_name):
|
if not _is_safe_msg_table_name(table_name):
|
||||||
raise ValueError(f'非法消息表名: {table_name}')
|
raise ValueError(f'非法消息表名: {table_name}')
|
||||||
|
|
||||||
clauses, params = _build_message_filters(start_ts, end_ts, keyword)
|
clauses, params = _build_message_filters(start_ts, end_ts, keyword, type_filter)
|
||||||
where_sql = f"WHERE {' AND '.join(clauses)}" if clauses else ''
|
where_sql = f"WHERE {' AND '.join(clauses)}" if clauses else ''
|
||||||
order = 'ASC' if oldest_first else 'DESC'
|
order = 'ASC' if oldest_first else 'DESC'
|
||||||
sql = f"""
|
sql = f"""
|
||||||
@@ -1376,7 +1415,7 @@ def _page_ranked_entries(entries, limit, offset, oldest_first=False):
|
|||||||
return paged
|
return paged
|
||||||
|
|
||||||
|
|
||||||
def _collect_chat_history_lines(ctx, names, start_ts=None, end_ts=None, limit=20, offset=0, oldest_first=False):
|
def _collect_chat_history_lines(ctx, names, start_ts=None, end_ts=None, limit=20, offset=0, oldest_first=False, type_filter=None):
|
||||||
collected = []
|
collected = []
|
||||||
failures = []
|
failures = []
|
||||||
candidate_limit = _candidate_page_size(limit, offset)
|
candidate_limit = _candidate_page_size(limit, offset)
|
||||||
@@ -1398,6 +1437,7 @@ def _collect_chat_history_lines(ctx, names, start_ts=None, end_ts=None, limit=20
|
|||||||
limit=batch_size,
|
limit=batch_size,
|
||||||
offset=fetch_offset,
|
offset=fetch_offset,
|
||||||
oldest_first=oldest_first,
|
oldest_first=oldest_first,
|
||||||
|
type_filter=type_filter,
|
||||||
)
|
)
|
||||||
if not rows:
|
if not rows:
|
||||||
break
|
break
|
||||||
@@ -1724,7 +1764,7 @@ def get_recent_sessions(limit: int = 20) -> str:
|
|||||||
|
|
||||||
|
|
||||||
@mcp.tool()
|
@mcp.tool()
|
||||||
def get_chat_history(chat_name: str, limit: int = 50, offset: int = 0, start_time: str = "", end_time: str = "", oldest_first: bool = False) -> str:
|
def get_chat_history(chat_name: str, limit: int = 50, offset: int = 0, start_time: str = "", end_time: str = "", oldest_first: bool = False, msg_types: list[str] | None = None) -> str:
|
||||||
"""获取指定聊天的消息记录。
|
"""获取指定聊天的消息记录。
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -1734,6 +1774,8 @@ def get_chat_history(chat_name: str, limit: int = 50, offset: int = 0, start_tim
|
|||||||
start_time: 起始时间,支持 YYYY-MM-DD / YYYY-MM-DD HH:MM / YYYY-MM-DD HH:MM:SS
|
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
|
end_time: 结束时间,支持 YYYY-MM-DD / YYYY-MM-DD HH:MM / YYYY-MM-DD HH:MM:SS
|
||||||
oldest_first: 为 True 时返回最早的消息(默认 False 返回最新消息)
|
oldest_first: 为 True 时返回最早的消息(默认 False 返回最新消息)
|
||||||
|
msg_types: 按消息类型过滤,可选值: text, image, voice, video, file(=app),
|
||||||
|
emoji, location, namecard, voip, system。传 None 或不传表示不过滤
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
_validate_pagination(limit, offset, limit_max=None)
|
_validate_pagination(limit, offset, limit_max=None)
|
||||||
@@ -1741,6 +1783,10 @@ def get_chat_history(chat_name: str, limit: int = 50, offset: int = 0, start_tim
|
|||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
return f"错误: {e}"
|
return f"错误: {e}"
|
||||||
|
|
||||||
|
type_filter, type_err = _resolve_msg_types(msg_types)
|
||||||
|
if type_err:
|
||||||
|
return f"错误: {type_err}"
|
||||||
|
|
||||||
ctx = _resolve_chat_context(chat_name)
|
ctx = _resolve_chat_context(chat_name)
|
||||||
if not ctx:
|
if not ctx:
|
||||||
return f"找不到聊天对象: {chat_name}\n提示: 可以用 get_contacts(query='{chat_name}') 搜索联系人"
|
return f"找不到聊天对象: {chat_name}\n提示: 可以用 get_contacts(query='{chat_name}') 搜索联系人"
|
||||||
@@ -1756,6 +1802,7 @@ def get_chat_history(chat_name: str, limit: int = 50, offset: int = 0, start_tim
|
|||||||
limit=limit,
|
limit=limit,
|
||||||
offset=offset,
|
offset=offset,
|
||||||
oldest_first=oldest_first,
|
oldest_first=oldest_first,
|
||||||
|
type_filter=type_filter,
|
||||||
)
|
)
|
||||||
|
|
||||||
if not lines:
|
if not lines:
|
||||||
@@ -1768,6 +1815,8 @@ def get_chat_history(chat_name: str, limit: int = 50, offset: int = 0, start_tim
|
|||||||
header += " [群聊]"
|
header += " [群聊]"
|
||||||
if start_time or end_time:
|
if start_time or end_time:
|
||||||
header += f"\n时间范围: {start_time or '最早'} ~ {end_time or '最新'}"
|
header += f"\n时间范围: {start_time or '最早'} ~ {end_time or '最新'}"
|
||||||
|
if msg_types:
|
||||||
|
header += f"\n类型过滤: {', '.join(msg_types)}"
|
||||||
if failures:
|
if failures:
|
||||||
header += "\n查询失败: " + ";".join(failures)
|
header += "\n查询失败: " + ";".join(failures)
|
||||||
return header + ":\n\n" + "\n".join(lines) + _pagination_hint(len(lines), limit, offset)
|
return header + ":\n\n" + "\n".join(lines) + _pagination_hint(len(lines), limit, offset)
|
||||||
|
|||||||
85
tests/test_msg_types_filter.py
Normal file
85
tests/test_msg_types_filter.py
Normal file
@@ -0,0 +1,85 @@
|
|||||||
|
"""测试 _resolve_msg_types 和 _build_message_filters 的 type_filter 路径。"""
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
|
||||||
|
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||||
|
|
||||||
|
import mcp_server
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolve_none_returns_no_filter():
|
||||||
|
assert mcp_server._resolve_msg_types(None) == (None, None)
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolve_empty_returns_no_filter():
|
||||||
|
assert mcp_server._resolve_msg_types([]) == (None, None)
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolve_single_text_type():
|
||||||
|
type_filter, err = mcp_server._resolve_msg_types(['text'])
|
||||||
|
assert err is None
|
||||||
|
assert type_filter == [1]
|
||||||
|
|
||||||
|
|
||||||
|
def test_resolve_multiple_types():
|
||||||
|
type_filter, err = mcp_server._resolve_msg_types(['image', 'voice', 'video'])
|
||||||
|
assert err is None
|
||||||
|
assert sorted(type_filter) == [3, 34, 43]
|
||||||
|
|
||||||
|
|
||||||
|
def test_file_alias_maps_to_app():
|
||||||
|
"""'file' 是常见叫法, 实际是 type=49 (app message)。"""
|
||||||
|
type_filter, err = mcp_server._resolve_msg_types(['file'])
|
||||||
|
assert err is None
|
||||||
|
assert type_filter == [49]
|
||||||
|
|
||||||
|
|
||||||
|
def test_case_insensitive_and_strip():
|
||||||
|
type_filter, err = mcp_server._resolve_msg_types([' Text ', 'IMAGE'])
|
||||||
|
assert err is None
|
||||||
|
assert sorted(type_filter) == [1, 3]
|
||||||
|
|
||||||
|
|
||||||
|
def test_unknown_type_returns_error():
|
||||||
|
type_filter, err = mcp_server._resolve_msg_types(['unknown'])
|
||||||
|
assert type_filter is None
|
||||||
|
assert err is not None
|
||||||
|
assert 'unknown' in err
|
||||||
|
assert 'text' in err # 错误提示列出可选值
|
||||||
|
|
||||||
|
|
||||||
|
def test_partial_unknown_aborts_whole():
|
||||||
|
"""混入一个未知类型时整体失败, 不偷偷过滤合法的。"""
|
||||||
|
type_filter, err = mcp_server._resolve_msg_types(['text', 'invalid_type'])
|
||||||
|
assert type_filter is None
|
||||||
|
assert 'invalid_type' in err
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_filters_without_type_filter():
|
||||||
|
"""type_filter=None 时 SQL 不包含 local_type 子句。"""
|
||||||
|
clauses, params = mcp_server._build_message_filters()
|
||||||
|
assert not any('local_type' in c for c in clauses)
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_filters_with_single_type():
|
||||||
|
clauses, params = mcp_server._build_message_filters(type_filter=[1])
|
||||||
|
assert any('local_type IN (?)' == c for c in clauses)
|
||||||
|
assert 1 in params
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_filters_with_multiple_types():
|
||||||
|
clauses, params = mcp_server._build_message_filters(type_filter=[1, 3, 34])
|
||||||
|
type_clause = [c for c in clauses if 'local_type' in c][0]
|
||||||
|
assert type_clause == 'local_type IN (?,?,?)'
|
||||||
|
assert params == [1, 3, 34]
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_filters_combines_with_time_and_keyword():
|
||||||
|
clauses, params = mcp_server._build_message_filters(
|
||||||
|
start_ts=1000, end_ts=2000, keyword='hello', type_filter=[1]
|
||||||
|
)
|
||||||
|
assert 'create_time >= ?' in clauses
|
||||||
|
assert 'create_time <= ?' in clauses
|
||||||
|
assert 'message_content LIKE ?' in clauses
|
||||||
|
assert any('local_type' in c for c in clauses)
|
||||||
|
assert params == [1000, 2000, '%hello%', 1]
|
||||||
Reference in New Issue
Block a user