context_manager.py 5.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169
  1. """Step 8: 上下文工程 — 对话压缩、Token 管理、多轮连贯性"""
  2. import json
  3. import os
  4. from datetime import datetime
  5. from typing import List, Dict, Optional
  6. class ContextManager:
  7. """对话上下文管理器:压缩历史、控制 Token 用量、保持连贯性"""
  8. def __init__(self, max_tokens: int = 4000, summary_trigger: int = 3000):
  9. self.max_tokens = max_tokens # 上下文最大 token 数
  10. self.summary_trigger = summary_trigger # 触发压缩的阈值
  11. self.turns: List[Dict] = [] # 对话轮次
  12. self.summary: str = "" # 压缩后的摘要
  13. self.total_turns = 0
  14. @staticmethod
  15. def _estimate_tokens(text: str) -> int:
  16. """简单 Token 估算:中文 ~1.5 字/token,英文 ~4 字/token"""
  17. chinese = sum(1 for c in text if '一' <= c <= '鿿')
  18. other = len(text) - chinese
  19. return int(chinese / 1.5 + other / 4)
  20. def add_turn(self, role: str, content: str):
  21. """添加一轮对话"""
  22. self.total_turns += 1
  23. turn = {
  24. "id": self.total_turns,
  25. "role": role,
  26. "content": content,
  27. "tokens": self._estimate_tokens(content),
  28. "time": datetime.now().strftime("%H:%M:%S"),
  29. }
  30. self.turns.append(turn)
  31. # 检查是否需要压缩
  32. total = sum(t["tokens"] for t in self.turns)
  33. if total > self.summary_trigger:
  34. self._compress()
  35. def _compress(self):
  36. """压缩早期对话为摘要"""
  37. if len(self.turns) <= 4:
  38. return # 保留最近 4 轮
  39. # 取最早的 60% 轮次进行压缩
  40. split = max(1, int(len(self.turns) * 0.6))
  41. old_turns = self.turns[:split]
  42. recent = self.turns[split:]
  43. # 生成摘要
  44. lines = []
  45. for t in old_turns:
  46. role_label = "用户" if t["role"] == "user" else "助手"
  47. snippet = t["content"][:200].replace("\n", " ")
  48. lines.append(f"[{role_label}]: {snippet}")
  49. new_summary = "对话历史摘要:\n" + "\n".join(lines)
  50. if self.summary:
  51. self.summary = self.summary[:500] + "\n...\n" + new_summary
  52. else:
  53. self.summary = new_summary
  54. # 限制摘要长度
  55. while self._estimate_tokens(self.summary) > 1500:
  56. # Drop the earliest part of the summary string by splitting on lines
  57. lines = self.summary.split('\n')
  58. if len(lines) <= 2:
  59. # If there are only a couple lines left, we must chop characters
  60. self.summary = self.summary[int(len(self.summary) * 0.8):]
  61. else:
  62. self.summary = "对话历史摘要:\n" + "\n".join(lines[2:])
  63. self.turns = recent
  64. def get_context(self, system_prompt: str = "",
  65. current_query: str = "") -> str:
  66. """构建当前上下文字符串"""
  67. parts = []
  68. # 压缩摘要
  69. if self.summary:
  70. parts.append(f"## 历史对话摘要\n{self.summary[:2000]}")
  71. # 最近对话
  72. if self.turns:
  73. parts.append("## 最近对话")
  74. for t in self.turns[-8:]: # 最近 8 轮
  75. role_label = "用户" if t["role"] == "user" else "助手"
  76. content = t["content"]
  77. if self._estimate_tokens(content) > 500:
  78. content = content[:500] + "..."
  79. parts.append(f"### {role_label}\n{content}")
  80. return "\n\n".join(parts)
  81. def get_stats(self) -> str:
  82. """获取上下文使用统计"""
  83. total = sum(t["tokens"] for t in self.turns)
  84. summary_tokens = self._estimate_tokens(self.summary) if self.summary else 0
  85. return (f"上下文: {len(self.turns)} 活跃轮次, "
  86. f"约 {total} tokens 活跃 + {summary_tokens} tokens 摘要, "
  87. f"总计 {self.total_turns} 轮对话")
  88. def clear(self):
  89. self.turns = []
  90. self.summary = ""
  91. self.total_turns = 0
  92. # ===== 上下文感知的 System Prompt 构建器 =====
  93. def build_context_aware_prompt(
  94. ctx: ContextManager,
  95. base_prompt: str,
  96. user_query: str,
  97. memory_context: str = "",
  98. kb_context: str = "",
  99. ) -> str:
  100. """构建完整上下文感知的系统消息"""
  101. parts = [base_prompt]
  102. # 对话上下文
  103. context_str = ctx.get_context()
  104. if context_str:
  105. parts.append(f"\n## 当前对话上下文\n{context_str}")
  106. # 记忆上下文
  107. if memory_context:
  108. parts.append(f"\n## 用户记忆\n{memory_context}")
  109. # 知识库上下文
  110. if kb_context:
  111. parts.append(f"\n## 相关知识\n{kb_context}")
  112. return "\n".join(parts)
  113. # 全局单例
  114. _ctx_instance: Optional[ContextManager] = None
  115. def get_context() -> ContextManager:
  116. global _ctx_instance
  117. if _ctx_instance is None:
  118. _ctx_instance = ContextManager()
  119. return _ctx_instance
  120. # ===== 工具函数 =====
  121. def context_stats(query: str = "") -> str:
  122. """查看当前上下文使用统计"""
  123. return get_context().get_stats()
  124. def context_clear(query: str = "") -> str:
  125. """清空上下文(开始新会话)"""
  126. get_context().clear()
  127. return "上下文已清空,开始新会话。"
  128. def context_summarize(query: str = "") -> str:
  129. """手动触发上下文压缩"""
  130. ctx = get_context()
  131. ctx._compress()
  132. return ctx.get_stats()