chat_history_tools.py 4.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142
  1. # -*- coding: utf-8 -*-
  2. """
  3. 聊天记录查询工具模块
  4. 提供 Agent 直接查询 SQLite 聊天记录的能力,解决 Agent 用文件系统工具
  5. 找不到聊天记录的问题(聊天记录存在 SQLite,不在文件系统中)。
  6. 工具列表:
  7. - search_chat_history: 按标题模糊搜索当前用户的会话
  8. - get_session_chat: 获取指定会话的完整消息历史
  9. """
  10. import json
  11. from tools.vent_tools import _current_user, _current_session
  12. async def search_chat_history(query: str, max_results: int = 5) -> str:
  13. """搜索当前用户的聊天历史记录(按会话标题匹配)。
  14. 当用户询问"之前聊过什么""搜索历史记录""找一下关于xxx的对话"时调用。
  15. 返回匹配的会话列表,包含标题、消息数量、会话 ID。
  16. 如果需要查看某个会话的具体对话内容,再调用 get_session_chat。
  17. Args:
  18. query: 搜索关键词,匹配会话标题
  19. max_results: 最大返回条数,默认 5
  20. Returns:
  21. JSON 格式,包含 matching_sessions 列表。
  22. 每条含 session_id、title、message_count、updated_at。
  23. """
  24. from db.chat_store import search_sessions as _db_search
  25. if not query or not query.strip():
  26. return json.dumps({"error": "搜索关键词不能为空"}, ensure_ascii=False)
  27. user_name = _current_user.get()
  28. try:
  29. sessions = _db_search(user_name, query.strip(), limit=max_results)
  30. results = [
  31. {
  32. "session_id": s["session_id"],
  33. "title": s.get("title", ""),
  34. "message_count": s.get("message_count", 0),
  35. "updated_at": s.get("updated_at", ""),
  36. }
  37. for s in sessions
  38. ]
  39. return json.dumps(
  40. {"query": query.strip(), "matching_sessions": results},
  41. ensure_ascii=False,
  42. )
  43. except Exception as e:
  44. return json.dumps({"error": f"搜索失败: {str(e)}"}, ensure_ascii=False)
  45. async def get_session_chat(session_id: str, limit: int = 50) -> str:
  46. """获取指定会话的完整消息历史。
  47. 当用户要求"查看那个会话的内容""把之前讨论的内容找出来""总结一下上次的对话",
  48. 且已知 session_id 时调用。通常先通过 search_chat_history 获取 session_id,
  49. 再调用本工具读取具体内容。
  50. Args:
  51. session_id: 会话 ID(UUID 格式)
  52. limit: 最大返回消息条数,默认 50
  53. Returns:
  54. JSON 格式,包含 session_info(标题、创建时间等)和 messages 列表。
  55. 每条消息含 role(user/assistant)、content、created_at。
  56. """
  57. from db.chat_store import get_messages as _db_msgs
  58. from db.chat_store import get_session_info as _db_info
  59. if not session_id or not session_id.strip():
  60. return json.dumps({"error": "会话 ID 不能为空"}, ensure_ascii=False)
  61. sid = session_id.strip()
  62. try:
  63. info = _db_info(sid)
  64. if not info:
  65. return json.dumps(
  66. {"error": f"会话不存在: {sid}", "session_id": sid},
  67. ensure_ascii=False,
  68. )
  69. msgs = _db_msgs(sid, limit=limit)
  70. messages = [
  71. {
  72. "role": m["role"],
  73. "content": m.get("content", ""),
  74. "created_at": m.get("created_at", ""),
  75. }
  76. for m in msgs
  77. ]
  78. return json.dumps(
  79. {
  80. "session_id": sid,
  81. "session_info": {
  82. "title": info.get("title", ""),
  83. "created_at": info.get("created_at", ""),
  84. "updated_at": info.get("updated_at", ""),
  85. },
  86. "message_count": len(messages),
  87. "messages": messages,
  88. },
  89. ensure_ascii=False,
  90. )
  91. except Exception as e:
  92. return json.dumps(
  93. {"error": f"读取失败: {str(e)}", "session_id": sid},
  94. ensure_ascii=False,
  95. )
  96. async def get_current_session_id() -> str:
  97. """获取当前本轮对话的数据库会话ID(session_id)。
  98. 当需要知道当前正在进行的会话的唯一标识符时调用此工具。
  99. 该 ID 可用于配合 get_session_chat 工具查看本轮会话的完整消息历史,
  100. 或配合 search_chat_history 跨会话检索。
  101. 无需任何参数,自动从请求上下文中获取。
  102. Returns:
  103. JSON 格式,包含当前会话的 session_id 和对应的 user_name。
  104. """
  105. user_name = _current_user.get()
  106. session_id = _current_session.get()
  107. if not session_id:
  108. return json.dumps({
  109. "error": "无法获取当前会话ID,可能不在对话上下文中",
  110. "session_id": "",
  111. "user_name": user_name,
  112. }, ensure_ascii=False)
  113. return json.dumps({
  114. "session_id": session_id,
  115. "user_name": user_name,
  116. }, ensure_ascii=False)