session_routes.py 6.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228
  1. # -*- coding: utf-8 -*-
  2. """
  3. 会话管理接口 —— 聊天历史查询、会话列表、删除会话、权限模式管理
  4. """
  5. from fastapi import APIRouter, Depends
  6. from pydantic import BaseModel, Field
  7. from api.auth import get_current_user
  8. from db.chat_store import (
  9. get_messages, get_sessions, delete_session,
  10. get_session_mode, set_session_mode, get_latest_context_usage,
  11. pin_session, is_session_pinned, update_session_title,
  12. )
  13. router = APIRouter()
  14. class ModeRequest(BaseModel):
  15. """权限模式切换请求"""
  16. mode: str = Field(..., description="权限模式", examples=["plan", "full"])
  17. class TitleRequest(BaseModel):
  18. """会话标题编辑请求"""
  19. title: str = Field(..., description="新标题", max_length=100, examples=["15216工作面需风量计算"])
  20. @router.get("/chat/history/{session_id}")
  21. async def get_chat_history(
  22. session_id: str,
  23. limit: int = 100,
  24. offset: int = 0,
  25. user_info: dict = Depends(get_current_user),
  26. ):
  27. """查询指定会话的聊天历史
  28. Args:
  29. session_id: 会话ID
  30. limit: 返回条数上限,默认100
  31. offset: 偏移量,默认0
  32. Returns:
  33. { "session_id": "...", "messages": [...], "total": N }
  34. """
  35. messages = get_messages(session_id, limit=limit, offset=offset)
  36. # 计算总数
  37. all_msgs = get_messages(session_id, limit=10000)
  38. total = len(all_msgs)
  39. return {
  40. "session_id": session_id,
  41. "messages": messages,
  42. "total": total,
  43. }
  44. @router.get("/sessions")
  45. async def list_sessions(
  46. limit: int = 20,
  47. offset: int = 0,
  48. user_info: dict = Depends(get_current_user),
  49. ):
  50. """查询当前用户的会话列表(按更新时间倒序)
  51. Returns:
  52. [{ "session_id": "...", "title": "...", "message_count": N, ... }, ...]
  53. """
  54. user_name = user_info.get("username", "admin")
  55. sessions = get_sessions(limit=limit, offset=offset, user_name=user_name)
  56. return {"sessions": sessions, "total": len(sessions)}
  57. @router.delete("/sessions/{session_id}")
  58. async def remove_session(
  59. session_id: str,
  60. user_info: dict = Depends(get_current_user),
  61. ):
  62. """删除指定会话及其所有消息"""
  63. delete_session(session_id)
  64. return {"status": "ok", "message": f"会话 {session_id} 已删除"}
  65. # ── 权限模式管理 ──
  66. @router.get("/sessions/{session_id}/mode")
  67. async def get_mode(
  68. session_id: str,
  69. user_info: dict = Depends(get_current_user),
  70. ):
  71. """获取会话的权限模式"""
  72. return {
  73. "session_id": session_id,
  74. "mode": get_session_mode(session_id),
  75. }
  76. @router.put("/sessions/{session_id}/mode")
  77. async def set_mode(
  78. session_id: str,
  79. req: ModeRequest,
  80. user_info: dict = Depends(get_current_user),
  81. ):
  82. """设置会话的权限模式
  83. mode 取值:
  84. - plan: 计划模式(先生成执行计划,用户确认后再执行)
  85. - full: 完全访问(无限制自动执行)
  86. """
  87. valid = {"plan", "full"}
  88. if req.mode not in valid:
  89. return {
  90. "error": f"无效的模式: {req.mode}",
  91. "valid_modes": sorted(valid),
  92. }
  93. set_session_mode(session_id, req.mode)
  94. return {
  95. "session_id": session_id,
  96. "mode": req.mode,
  97. "message": f"已切换为「{req.mode}」模式",
  98. }
  99. # ── 会话置顶 ──
  100. @router.put("/sessions/{session_id}/pin")
  101. async def toggle_pin(
  102. session_id: str,
  103. user_info: dict = Depends(get_current_user),
  104. ):
  105. """切换会话置顶状态(置顶 ⇄ 取消置顶)。
  106. Returns:
  107. { "session_id": "...", "pinned": true, "message": "..." }
  108. """
  109. pinned = not is_session_pinned(session_id)
  110. pin_session(session_id, pinned)
  111. return {
  112. "session_id": session_id,
  113. "pinned": pinned,
  114. "message": "已置顶" if pinned else "已取消置顶",
  115. }
  116. # ── 会话标题编辑 ──
  117. @router.put("/sessions/{session_id}/title")
  118. async def edit_title(
  119. session_id: str,
  120. req: TitleRequest,
  121. user_info: dict = Depends(get_current_user),
  122. ):
  123. """编辑会话标题。
  124. Args:
  125. req.title: 新标题,最长100字
  126. """
  127. title = req.title.strip()
  128. if not title:
  129. return {"error": "标题不能为空"}
  130. if len(title) > 100:
  131. return {"error": f"标题过长({len(title)}字),最多100字"}
  132. update_session_title(session_id, title)
  133. return {
  134. "session_id": session_id,
  135. "title": title,
  136. "message": "标题已更新",
  137. }
  138. # ── 上下文容量查询 ──
  139. @router.get("/sessions/{session_id}/context")
  140. async def get_context_usage(
  141. session_id: str,
  142. user_info: dict = Depends(get_current_user),
  143. ):
  144. """查询会话最近一次的上下文用量详情。
  145. Returns:
  146. {
  147. "session_id": "...",
  148. "total_limit": 131072,
  149. "current_usage": 12345,
  150. "usage_pct": 9.4,
  151. "breakdown": {
  152. "messages": 5000,
  153. "mcp": 4000,
  154. "skills": 800,
  155. "system_prompt": 600,
  156. "other": 2745
  157. },
  158. "updated_at": "2026-08-05 14:30:00"
  159. }
  160. 若无记录则返回 null 值。
  161. """
  162. row = get_latest_context_usage(session_id)
  163. if not row:
  164. return {
  165. "session_id": session_id,
  166. "total_limit": None,
  167. "current_usage": None,
  168. "usage_pct": None,
  169. "breakdown": None,
  170. "updated_at": None,
  171. }
  172. total_limit = row.get("total_limit", 131072)
  173. current = row.get("current_usage", 0)
  174. pct = round(current / total_limit * 100, 1) if total_limit else 0
  175. return {
  176. "session_id": session_id,
  177. "total_limit": total_limit,
  178. "current_usage": current,
  179. "usage_pct": pct,
  180. "breakdown": {
  181. "messages": row.get("messages_tokens", 0),
  182. "mcp": row.get("mcp_tokens", 0),
  183. "skills": row.get("skills_tokens", 0),
  184. "system_prompt": row.get("system_prompt_tokens", 0),
  185. "other": row.get("other_tokens", 0),
  186. },
  187. "updated_at": row.get("updated_at"),
  188. }