From 836d50a548e39f650654833d7f5473f9c810a703 Mon Sep 17 00:00:00 2001 From: huaiyj <8699003+huaiyj@user.noreply.gitee.com> Date: Mon, 24 Aug 2026 20:53:07 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=94=AF=E6=8C=81=E5=BC=82=E6=AD=A5?= =?UTF-8?q?=E6=BB=9A=E5=8A=A8=E4=BC=9A=E8=AF=9D=E6=91=98=E8=A6=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ocai-service/ocai/ai/agent_manager.py | 3 +- ocai-service/ocai/ai/memory/summary.py | 174 +++++++++++++++---- ocai-service/tests/ai/test_summary_memory.py | 88 ++++++++++ 3 files changed, 226 insertions(+), 39 deletions(-) create mode 100644 ocai-service/tests/ai/test_summary_memory.py diff --git a/ocai-service/ocai/ai/agent_manager.py b/ocai-service/ocai/ai/agent_manager.py index db47dfe..a19632e 100644 --- a/ocai-service/ocai/ai/agent_manager.py +++ b/ocai-service/ocai/ai/agent_manager.py @@ -555,8 +555,9 @@ class AgentManager: summary = SummaryMemory( summarize_threshold=summarize_threshold, keep_recent=max_messages // 2, + async_summarize=True, ) - return CompositeMemory(layers=[conv, summary]) + return CompositeMemory(layers=[summary, conv]) # ================================================================ # 工具管理 diff --git a/ocai-service/ocai/ai/memory/summary.py b/ocai-service/ocai/ai/memory/summary.py index 8009317..aea69a2 100644 --- a/ocai-service/ocai/ai/memory/summary.py +++ b/ocai-service/ocai/ai/memory/summary.py @@ -8,7 +8,9 @@ SummaryMemory — 中期摘要记忆 """ import logging +import threading from collections.abc import Callable +from datetime import UTC, datetime from typing import Any from ocai.ai.memory.base import BaseMemory, MemoryEntry @@ -39,6 +41,8 @@ class SummaryMemory(BaseMemory): summarize_threshold: int = 20, keep_recent: int = 10, summarize_fn: Callable[[list[dict[str, str]]], str] | None = None, + max_summary_chars: int = 4000, + async_summarize: bool = True, ): """ Args: @@ -46,14 +50,25 @@ class SummaryMemory(BaseMemory): keep_recent: 摘要时保留最近 N 条原始消息 summarize_fn: 摘要函数 (messages -> summary_text), 为 None 时使用简单拼接回退 + max_summary_chars: 滚动摘要的最大字符预算 + async_summarize: 是否在后台线程生成摘要 """ self._threshold = summarize_threshold self._keep_recent = keep_recent self._summarize_fn = summarize_fn + self._max_summary_chars = max(256, int(max_summary_chars)) + self._async_summarize = async_summarize self._entries: list[MemoryEntry] = [] self._summary: str | None = None self._summary_count: int = 0 # 已被摘要的消息数 + self._summary_version: int = 0 + self._summary_updated_at: str | None = None + self._summarizing = False + self._generation = 0 + self._lock = threading.RLock() + self._idle = threading.Event() + self._idle.set() def add( self, @@ -73,10 +88,11 @@ class SummaryMemory(BaseMemory): is_public=is_public, reasoning=reasoning, ) - self._entries.append(entry) + with self._lock: + self._entries.append(entry) + should_summarize = len(self._entries) >= self._threshold and not self._summarizing - # 检查是否需要触发摘要 - if len(self._entries) >= self._threshold: + if should_summarize: self._do_summarize() def get_context( @@ -88,17 +104,19 @@ class SummaryMemory(BaseMemory): ) -> list[dict[str, str]]: messages = [] - # 注入摘要 - if self._summary: + with self._lock: + summary = self._summary + entries = list(self._entries) + + if summary: messages.append( { "role": "system", - "content": f"[对话摘要]\n{self._summary}", + "content": f"[对话摘要]\n{summary}", } ) # 添加当前消息 - entries = self._entries if agent_name: entries = [e for e in entries if e.agent_name is None or e.agent_name == agent_name] @@ -110,50 +128,130 @@ class SummaryMemory(BaseMemory): return messages def clear(self) -> None: - self._entries.clear() - self._summary = None - self._summary_count = 0 + with self._lock: + self._generation += 1 + self._entries.clear() + self._summary = None + self._summary_count = 0 + self._summary_version = 0 + self._summary_updated_at = None + self._summarizing = False + self._idle.set() @property def size(self) -> int: - return len(self._entries) + with self._lock: + return len(self._entries) def summarize(self) -> str | None: - return self._summary + with self._lock: + return self._summary + + @property + def summary_metadata(self) -> dict[str, Any]: + with self._lock: + return { + "version": self._summary_version, + "updated_at": self._summary_updated_at, + "summarized_messages": self._summary_count, + "source_message_range": ({"start": 1, "end": self._summary_count} if self._summary_count else None), + "max_chars": self._max_summary_chars, + "pending": self._summarizing, + } + + def wait_for_summary(self, timeout: float | None = None) -> bool: + """Wait for pending background summarization, mainly for graceful shutdown.""" + return self._idle.wait(timeout) def _do_summarize(self) -> None: - """执行摘要:将旧消息压缩为摘要文本""" - to_summarize = self._entries[: -self._keep_recent] - keep = self._entries[-self._keep_recent :] + """Capture old messages and summarize them without blocking new writes.""" + with self._lock: + if self._summarizing: + return + count = len(self._entries) - self._keep_recent + if count <= 0: + return + messages = [{"role": entry.role, "content": entry.content} for entry in self._entries[:count]] + generation = self._generation + self._summarizing = True + self._idle.clear() + + if self._async_summarize: + threading.Thread( + target=self._finish_summarize, + args=(messages, count, generation), + name="ocai-summary", + daemon=True, + ).start() + else: + self._finish_summarize(messages, count, generation) - if not to_summarize: - return + def _finish_summarize( + self, + messages: list[dict[str, str]], + count: int, + generation: int, + ) -> None: + new_summary = self._generate_summary(messages) + schedule_next = False + with self._lock: + if generation != self._generation: + return + self._summary = self._merge_summaries(self._summary, new_summary) + del self._entries[:count] + self._summary_count += count + self._summary_version += 1 + self._summary_updated_at = datetime.now(UTC).isoformat() + self._summarizing = False + self._idle.set() + schedule_next = len(self._entries) >= self._threshold - msgs = [{"role": e.role, "content": e.content} for e in to_summarize] + logger.debug( + "Summarized %d messages, total=%d, version=%d", + count, + self._summary_count, + self._summary_version, + ) + if schedule_next: + self._do_summarize() + def _generate_summary(self, messages: list[dict[str, str]]) -> str: if self._summarize_fn: try: - new_summary = self._summarize_fn(msgs) - except Exception as e: - logger.warning(f"Summarize failed, using fallback: {e}") - new_summary = self._fallback_summarize(msgs) - else: - new_summary = self._fallback_summarize(msgs) - - # 合并旧摘要 - if self._summary: - self._summary = f"{self._summary}\n\n{new_summary}" - else: - self._summary = new_summary - - self._summary_count += len(to_summarize) - self._entries = keep + generated = self._summarize_fn(messages) + if generated: + return str(generated) + except Exception as exc: + logger.warning("Summarize failed, using fallback: %s", exc) + return self._fallback_summarize(messages) + + def _merge_summaries(self, old: str | None, new: str) -> str: + combined = f"{old}\n\n{new}" if old else new + if len(combined) <= self._max_summary_chars: + return combined - logger.debug( - f"Summarized {len(to_summarize)} messages, " - f"keeping {len(keep)} recent. " - f"Total summarized: {self._summary_count}" - ) + if self._summarize_fn: + try: + compacted = self._summarize_fn( + [ + { + "role": "system", + "content": ( + "将下面的历史摘要重写为一份滚动摘要,保留关键结论、" + f"操作和未决事项,长度不超过 {self._max_summary_chars} 字符。" + ), + }, + {"role": "user", "content": combined}, + ] + ) + if compacted: + return compacted[: self._max_summary_chars] + except Exception as exc: + logger.warning("Rolling summary compression failed: %s", exc) + + head_size = self._max_summary_chars // 3 + tail_size = self._max_summary_chars - head_size - 7 + return f"{combined[:head_size]}\n[...]\n{combined[-tail_size:]}" @staticmethod def _fallback_summarize(messages: list[dict[str, str]]) -> str: diff --git a/ocai-service/tests/ai/test_summary_memory.py b/ocai-service/tests/ai/test_summary_memory.py new file mode 100644 index 0000000..de6e61d --- /dev/null +++ b/ocai-service/tests/ai/test_summary_memory.py @@ -0,0 +1,88 @@ +import threading +import time + +from ocai.ai.memory.summary import SummaryMemory +from ocai.ai.agent_manager import AgentManager + + +def test_background_summary_does_not_block_add(): + release = threading.Event() + + def summarize(messages): + release.wait(1) + return "summary" + + memory = SummaryMemory( + summarize_threshold=3, + keep_recent=1, + summarize_fn=summarize, + async_summarize=True, + ) + memory.add("user", "one") + memory.add("assistant", "two") + + started = time.monotonic() + memory.add("user", "three") + elapsed = time.monotonic() - started + + assert elapsed < 0.2 + assert memory.summary_metadata["pending"] is True + release.set() + assert memory.wait_for_summary(1) + assert memory.summarize() == "summary" + assert memory.summary_metadata["summarized_messages"] == 2 + assert memory.summary_metadata["source_message_range"] == {"start": 1, "end": 2} + + +def test_rolling_summary_stays_within_budget_and_tracks_version(): + memory = SummaryMemory( + summarize_threshold=3, + keep_recent=1, + max_summary_chars=256, + async_summarize=False, + summarize_fn=lambda messages: "x" * 180, + ) + + for index in range(5): + memory.add("user", f"message-{index}") + + assert len(memory.summarize()) <= 256 + assert memory.summary_metadata["version"] == 2 + assert memory.summary_metadata["updated_at"] is not None + + +def test_clear_discards_inflight_summary(): + release = threading.Event() + memory = SummaryMemory( + summarize_threshold=2, + keep_recent=1, + summarize_fn=lambda messages: (release.wait(1), "late")[1], + ) + memory.add("user", "one") + memory.add("assistant", "two") + + memory.clear() + release.set() + + assert memory.wait_for_summary(1) + assert memory.summarize() is None + assert memory.size == 0 + + +def test_agent_manager_injects_generated_summary_into_primary_context(): + memory = AgentManager().create_memory( + max_messages=4, + enable_summary=True, + summarize_threshold=3, + ) + memory.add("user", "one") + memory.add("assistant", "two") + memory.add("user", "three") + + deadline = time.monotonic() + 1 + while memory.summarize() is None and time.monotonic() < deadline: + time.sleep(0.01) + + context = memory.get_context() + assert context[0]["role"] == "system" + assert "对话摘要" in context[0]["content"] -- Gitee