diff --git a/.gitignore b/.gitignore index 91426e394..60da44eaa 100644 --- a/.gitignore +++ b/.gitignore @@ -28,3 +28,7 @@ pyrightconfig.json # spec-workflow tool artifacts .spec-workflow + +# Local-only examples +examples/session_service_with_advanced_memory_sql/ +examples/session_service_with_advanced_memory_redis/ diff --git a/README.zh_CN.md b/README.zh_CN.md index 29f3ea24b..5a8e7fd0e 100644 --- a/README.zh_CN.md +++ b/README.zh_CN.md @@ -497,7 +497,7 @@ skill_tool_set = SkillToolSet(repository=repository, run_tool_kwargs=tool_kwargs 建议先看: -- Session:[examples/session_service_with_in_memory](./examples/session_service_with_in_memory/README.md) / [examples/session_service_with_redis](./examples/session_service_with_redis/README.md) / [examples/session_service_with_sql](./examples/session_service_with_sql/README.md) / [examples/session_summarizer](./examples/session_summarizer/README.md) / [examples/session_state](./examples/session_state/README.md) +- Session:[examples/session_service_with_in_memory](./examples/session_service_with_in_memory/README.md) / [examples/session_service_with_redis](./examples/session_service_with_redis/README.md) / [examples/session_service_with_sql](./examples/session_service_with_sql/README.md) / [Advanced Memory Redis 压缩](./examples/session_service_with_advanced_memory_redis/README.md) / [Advanced Memory SQL 压缩](./examples/session_service_with_advanced_memory_sql/README.md) / [examples/session_summarizer](./examples/session_summarizer/README.md) / [examples/session_state](./examples/session_state/README.md) - Memory: [examples/memory_service_with_in_memory](./examples/memory_service_with_in_memory/README.md) / [examples/memory_service_with_redis](./examples/memory_service_with_redis/README.md) / [examples/memory_service_with_sql](./examples/memory_service_with_sql/README.md) / [examples/memory_service_with_mem0](./examples/memory_service_with_mem0/README.md) / [examples/memory_service_with_mempalace](./examples/memory_service_with_mempalace/README.md) - Knowledge:[examples/knowledge_with_documentloader](./examples/knowledge_with_documentloader/README.md) / [examples/knowledge_with_vectorstore](./examples/knowledge_with_vectorstore/README.md) / [examples/knowledge_with_rag_agent](./examples/knowledge_with_rag_agent/README.md) / [examples/knowledge_with_searchtool_rag_agent](./examples/knowledge_with_searchtool_rag_agent/README.md) / [examples/knowledge_with_prompt_template](./examples/knowledge_with_prompt_template/README.md) / [examples/knowledge_with_custom_components](./examples/knowledge_with_custom_components/README.md) diff --git a/examples/memory_service_with_advanced_memory/.env b/examples/memory_service_with_advanced_memory/.env index 2da17e1ce..8061a2bc8 100644 --- a/examples/memory_service_with_advanced_memory/.env +++ b/examples/memory_service_with_advanced_memory/.env @@ -1,8 +1,4 @@ # Set TRPC_AGENT_API_KEY, TRPC_AGENT_BASE_URL, and TRPC_AGENT_MODEL_NAME. TRPC_AGENT_API_KEY= TRPC_AGENT_BASE_URL= -TRPC_AGENT_MODEL_NAME= -# Optional: enable token-based context budgeting for Advanced Memory. -# Set both model limits to enable token-based context budgeting. -TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS= -TRPC_AGENT_MAX_OUTPUT_TOKENS= +TRPC_AGENT_MODEL_NAME= \ No newline at end of file diff --git a/examples/memory_service_with_advanced_memory/README.md b/examples/memory_service_with_advanced_memory/README.md index 0b430210f..e81a4aea4 100644 --- a/examples/memory_service_with_advanced_memory/README.md +++ b/examples/memory_service_with_advanced_memory/README.md @@ -1,223 +1,111 @@ -# Advanced Memory - -## Advanced Memory 简介 - -`Advanced Memory` 是一套面向 Agent 的本地化记忆与上下文管理机制,重点增强 -Agent 在长期信息沉淀和超长对话处理方面的能力: - -- **本地化持久存储**:记忆和上下文数据以本地文件形式持久化,存储位置、数据边界 - 和组织方式清晰可控,适合本地开发、调试、迁移和审计。 -- **更强的长期记忆能力**:支持将对话中的稳定事实、用户偏好和重要经验主动沉淀为 - 可组织、可更新、可跨 Session 使用的长期记忆,而不是简单堆积历史消息。 -- **分层记忆管理**:分别管理原始对话、Session 级记忆和跨 Session 长期记忆,让不同 - 类型的信息以合适的粒度参与后续推理。 -- **上下文管理**:根据上下文规模、信息类型和使用情况,对历史消息、工具结果及记忆 - 内容进行统一治理,在保留关键信息的同时控制模型输入规模。 -- **上下文压缩**:支持对历史上下文和工具结果进行渐进式裁剪、压缩和摘要,降低长 - 对话导致的上下文膨胀以及超出模型窗口限制的风险。 -- **结构化记忆提取**:从持续增长的对话中提取结构化信息,形成更稳定、更易维护的 - Session Memory,提升后续对话对历史信息的利用效率。 - -本示例演示如何使用 `AdvancedMemorySessionService`。它把 Session 持久化和 -Advanced Memory 上下文管理整合到一个 SessionService 中,用户不需要显式调用 -`setup_advanced_memory()`,也不需要再创建 `InMemorySessionService`。 - -## 示例流程 - -脚本使用同一个 Runner 执行多个 Session: - -1. `session-1` 连续输入多轮 Python 开发偏好。 -2. 当累计上下文和工具调用达到配置阈值后,系统会提取 session memory,并写入 - `session_memory.md`。 -3. `session-1` 请求总结已经学习到的开发偏好。 -4. `session-2` 查询长期记忆,验证不同 Session 共享同一个 `MEMORY/`。 - -## 使用方式 - -```python -from pathlib import Path - -from trpc_agent_sdk.memory import AdvancedMemoryConfig -from trpc_agent_sdk.sessions import AdvancedMemorySessionService -from trpc_agent_sdk.runners import Runner - -session_service = AdvancedMemorySessionService( - config=AdvancedMemoryConfig( - root_dir=Path(__file__).resolve().parent, - ) -) - -runner = Runner( - app_name="advanced_memory_demo", - agent=agent, - session_service=session_service, -) +# Advanced Memory 本地持久化示例 + +本示例演示如何使用 `AdvancedMemoryService` 在本地实现持久化的跨会话记忆。 +Agent 可以主动保存用户的重要信息,并在后续会话中根据记忆索引查找和读取相关内容。 + +## 关键特性 + +- **主动式记忆**:由 Agent 根据对话内容主动调用工具保存长期有效的信息。 +- **记忆分类**:每条记忆包含名称、描述、类型、摘要和详细内容。 +- **基于记忆索引的记忆召回**:先读取 `MEMORY.md` 索引,再读取匹配的记忆文件,避免检索全部记忆内容。 +- **跨会话持久化**:本地记忆默认保存在示例目录下,并按应用和用户进行隔离。(支持 Redis,SQL 存储) + +## Agent 层级结构说明 + +`AdvancedMemoryService` 通过 `Runner` 绑定到 Agent。它负责初始化长期记忆运行时、注入记忆相关指令并安装工具;具体的记忆保存和读取由 Agent 根据工具描述主动完成。 + +## 关键代码解释 + +### `save_memory` + +保存或更新一条长期记忆,并同步更新 `MEMORY.md` 索引。适合保存用户的稳定偏好、 +习惯和其他未来会话仍然有价值的信息。 + +### `list_memory_index` + +读取当前用户的记忆索引。Agent 在需要回忆信息时应先调用这个工具,了解有哪些可用记忆。 + +### `read_memory` + +根据索引中的文件名读取完整记忆内容。这样可以只读取与当前问题相关的记忆。 + +## 环境要求 + +- Python 3.10 或更高版本 +- 已安装项目依赖 +- 一个可访问的 OpenAI 兼容模型服务 + +在 `.env` 中配置: + +```dotenv +TRPC_AGENT_API_KEY=your-api-key +TRPC_AGENT_BASE_URL=https://your-llm-endpoint/v1 +TRPC_AGENT_MODEL_NAME=your-model-name ``` -`Runner` 检测到 `AdvancedMemorySessionService` 后会自动完成 Advanced Memory -绑定,包括: - -- transcript 持久化 -- session memory 提取 -- 长期记忆 tools:`save_memory`、`read_memory`、`list_memory_index` -- `HistorySnip` -- `Microcompact` -- `AutoCompact` -- `ToolResultBudget` - -`AdvancedMemoryConfig` 默认已经启用这些能力,本示例直接使用默认配置。 - -## 数据目录 - -运行后,数据默认写入当前示例目录: - -```text -MEMORY/ -├── MEMORY.md -└── *.md # 长期记忆详情 - -SESSION/ -├── _state.json # app/user 级 state -├── session-1/ -│ ├── session.json # Session 元数据和 session state -│ ├── transcript.jsonl # 原始 Events 和 checkpoint -│ ├── session_memory.md # 结构化 Session 记忆 -│ └── tool-results/ # 超大工具结果 -└── session-2/ - ├── session.json - ├── transcript.jsonl - └── session_memory.md +## 代码构建 + +```bash +git clone https://github.com/trpc-group/trpc-agent-python.git +cd trpc-agent-python +./build.sh +source .venv/bin/activate ``` -其中: +如果已经在当前项目中创建了 Python 3.10+ 虚拟环境,也可以直接安装: -- `session.json` 保存 Session 元数据和状态,不保存完整 Events。 -- `transcript.jsonl` 是追加写入的原始事件日志,可用于恢复 Session。 -- `session_memory.md` 是根据 transcript 提取的结构化摘要。 -- `MEMORY/` 保存跨 Session 使用的长期记忆。 +```bash +python -m pip install -e . +``` ## 运行 -先在本目录创建 `.env`,然后填写模型配置: - ```bash cd examples/memory_service_with_advanced_memory -python3 run_agent.py +python run_agent.py ``` -需要的环境变量: - -- `TRPC_AGENT_API_KEY` -- `TRPC_AGENT_BASE_URL` -- `TRPC_AGENT_MODEL_NAME` -- `TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS`(可选,模型总上下文窗口大小,单位为 token) -- `TRPC_AGENT_MAX_OUTPUT_TOKENS`(可选,模型最大输出窗口大小,单位为 token) - -`.env` 中留空的变量不会覆盖默认值;如果同时在 Python 中传入 -`model_context_window_tokens` 或 `max_output_tokens`,Python 显式配置优先。 - -如果配置了模型上下文窗口,Advanced Memory 会用 -`TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS - TRPC_AGENT_MAX_OUTPUT_TOKENS` -作为可用于输入内容的窗口;两个变量都留空时使用字符数阈值。 - -## `AdvancedMemoryConfig` 配置项 - -下面列出当前所有可直接传入 `AdvancedMemoryConfig` 的配置项。**没有特殊需求时, -只设置 `root_dir` 即可**;示例中的值均为默认值。 - -```python -session_service = AdvancedMemorySessionService( - config=AdvancedMemoryConfig( - root_dir=Path(__file__).resolve().parent, # 当前示例目录 - # Optional - enabled=True, # 总开关和存储路径 - memory_dir_name="MEMORY", # 长期记忆目录 - session_dir_name="SESSION", # Session 数据目录 - memory_index_name="MEMORY.md", # 长期记忆索引文件 - transcript_name="transcript.jsonl", # transcript 文件 - session_memory_name="session_memory.md", # Session 摘要文件 - encoding="utf-8", # 文件编码 - transcript_fsync=False, # transcript 写入后是否 fsync - - # 长期记忆 - memory_index_max_lines=200, # 注入 prompt 的索引最大行数 - memory_index_max_bytes=25_000, # 注入 prompt 的索引最大字节数 - long_term_memory_injection_enabled=True, # 是否注入 MEMORY.md - - # 工具结果 - tool_result_max_chars=50_000, # 单个工具结果最大字符数 - tool_results_per_message_max_chars=200_000, # 单条消息工具结果总上限 - tool_result_preview_chars=2_000, # 超限结果的预览字符数 - - # HistorySnip - history_snip_enabled=True, # 是否压缩过长历史 - history_snip_trigger_chars=600_000, # 触发阈值 - history_snip_target_chars=400_000, # 压缩目标 - history_snip_keep_recent=5, # 保留最近的完整消息数 - history_snip_tool_names=( # 可处理的工具名称 - "Read", "Bash", "Grep", "Glob", - "WebSearch", "WebFetch", "Edit", "Write", - ), - - # Token 上下文预算 - # 这两个值也可以通过 .env 配置;显式传参优先于环境变量。 - # model_context_window_tokens=131072, # 显式设置后覆盖环境变量 - # max_output_tokens=8192, # 显式设置后覆盖环境变量 - # 如果省略这两行,则分别读取 .env;未配置时默认 None 和 0。 - token_warning_ratio=0.85, # 告警比例 - token_autocompact_ratio=0.90, # 自动压缩比例 - token_blocking_ratio=0.95, # 阻止继续增加上下文的比例 - token_estimator=None, # 可选:自定义 token 估算器 - context_window_resolver=None, # 可选:自定义窗口解析器 - - # Session Memory - session_memory_enabled=True, # 是否启用 Session 摘要 - session_memory_initial_chars=40_000, # 首次提取字符阈值 - session_memory_update_chars=20_000, # 后续更新字符阈值 - session_memory_initial_tokens=10_000, # 首次提取 token 阈值 - session_memory_update_tokens=5_000, # 后续更新 token 阈值 - session_memory_tool_calls_between_updates=3, # 两次更新间的工具调用数 - session_memory_prompt_max_chars=200_000, # 摘要请求最大字符数 - session_memory_request_overhead_tokens=2_048, # 请求预留 token - session_memory_section_max_chars=8_000, # 单个摘要 section 最大字符数 - session_memory_total_max_chars=54_000, # 摘要总最大字符数 - session_memory_wait_timeout_seconds=15.0, # 等待摘要 Agent 的超时时间 - - # AutoCompact - autocompact_enabled=True, # 是否启用自动压缩 - autocompact_trigger_chars=700_000, # 触发阈值 - autocompact_target_chars=350_000, # 压缩目标 - autocompact_blocking_chars=780_000, # 阻止继续增加上下文的阈值 - autocompact_keep_recent_contents=8, # 保留最近内容数 - autocompact_max_failures=3, # 最大连续失败次数 - autocompact_summary_input_max_chars=600_000, # 摘要 Agent 输入上限 - autocompact_summary_retries=3, # 摘要 Agent 重试次数 - - # Microcompact - microcompact_enabled=True, # 是否启用工具结果微压缩 - microcompact_gap_seconds=3_600.0, # 工具结果时间间隔阈值 - microcompact_trigger_count=20, # 触发工具结果数量 - microcompact_keep_recent=5, # 保留最近工具结果数 - microcompact_tool_names=( # 可处理的工具名称 - "Read", "Bash", "Grep", "Glob", - "WebSearch", "WebFetch", "Edit", "Write", - ), - - # Advanced Memory preload - preload_memory_enabled=False, # 是否自动预加载相关 topic - preload_memory_max_topics=5, # 一次最多加载的 topic 数 - preload_memory_max_chars=50_000, # 预加载内容总字符上限 - preload_memory_candidate_limit=200, # 筛选模型的候选 topic 数 - ), -) -``` +示例会使用同一用户运行多个会话,验证长期记忆可以在不同会话之间复用。 + +## 运行结果(实测) + +```txt +👤 [session-1] Please remember that my favorite programming language is Python. Save this as a user preference. +[2026-09-11 13:17:51][INFO][trpc_agent_sdk][trpc_agent_sdk/sessions/_in_memory_session_service.py:398][4086343] Cleanup task started with interval: 5.0s +🔧 save_memory({'filename': 'favorite_programming_language.md', 'name': 'Favorite programming language', 'description': "The user's favorite programming language and related preference.", 'memory_type': 'user', 'summary': "The user's favorite programming language is Python.", 'content': "The user's favorite programming language is Python. When discussing languages, choosing examples, or recommending tech stacks, prefer Python."}) +📊 {'saved': True, 'filename': 'favorite_programming_language.md', 'path': '/data/workspace/trpc-agent-python-am-service2/examples/memory_service_with_advanced_memory/tenants/advanced_memory_demo/demo-user/MEMORY/favorite_programming_language.md', 'memory_type': 'user', 'updated_at': '2026-09-11T05:17:54.071665+00:00'} +🤖 Done! ✅ I've saved your preference that **Python is your favorite programming language** as a long-term user memory. + +I'll keep this in mind across future conversations — for example, when recommending libraries, writing code examples, or discussing tech stacks, I'll lean toward Python by default. + +👤 [session-1] I use Python mainly for backend services and data processing. +🤖 That's useful context — let me update your existing memory to include it. + + +🔧 save_memory({'content': "The user's favorite programming language is Python. They use Python mainly for backend services and data processing. When discussing languages, choosing examples, or recommending tech stacks, prefer Python and focus on backend development and data processing use cases.", 'description': "The user's favorite programming language and how they use it.", 'filename': 'favorite_programming_language.md', 'memory_type': 'user', 'summary': "The user's favorite programming language is Python, used mainly for backend services and data processing.", 'name': 'Favorite programming language'}) +📊 {'saved': True, 'filename': 'favorite_programming_language.md', 'path': '/data/workspace/trpc-agent-python-am-service2/examples/memory_service_with_advanced_memory/tenants/advanced_memory_demo/demo-user/MEMORY/favorite_programming_language.md', 'memory_type': 'user', 'updated_at': '2026-09-11T05:17:56.582892+00:00'} +🤖 Got it — I've updated your memory. ✅ + +...... + +👤 [session-2] What do you remember about my favorite programming language? +🤖 + + +🔧 read_memory({'filename': 'favorite_programming_language.md'}) +📊 {'found': True, 'filename': 'favorite_programming_language.md', 'content': "---\nname: Favorite programming language\ndescription: The user's favorite programming language and how they use it.\ntype: user\nupdated_at: 2026-09-11T05:17:56.582892+00:00\n---\nThe user's favorite programming language is Python. They use Python mainly for backend services and data processing. When discussing languages, choosing examples, or recommending tech stacks, prefer Python and focus on backend development and data processing use cases.\n", 'updated_at': '2026-09-11T05:17:56.582892+00:00', 'freshness': 'today', 'freshness_notice': 'This memory was last updated today. It is a point-in-time observation and may no longer reflect the current state. Verify it when necessary, and update this memory if it is outdated or incorrect.'} +🤖 Here's what I remember about your favorite programming language: + +**Python** 🐍 + +From my long-term memory: +- **Python is your favorite programming language**, and you use it mainly for **backend services** and **data processing**. +- When discussing languages, choosing examples, or recommending tech stacks, I should prefer Python and focus on backend development and data processing use cases. -`preload_memory_model` 不是 `AdvancedMemoryConfig` 字段,而是 -`AdvancedMemorySessionService` 的可选参数,用于指定轻量筛选模型: +Related preferences I also have on file: +- You like **typed Python code** with clear dataclasses and small, focused modules. +- You prefer **pytest and focused unit tests** for Python testing. +- You like **concise documentation** with runnable commands and examples. -```python -session_service = AdvancedMemorySessionService( - config=AdvancedMemoryConfig(preload_memory_enabled=True), - preload_memory_model=small_model, # 不传时复用主 Agent 的模型 -) +Is there anything else you'd like me to remember or clarify about your language preferences? ``` diff --git a/examples/memory_service_with_advanced_memory/agent/agent.py b/examples/memory_service_with_advanced_memory/agent/agent.py index 8f93758ec..2165a5f26 100644 --- a/examples/memory_service_with_advanced_memory/agent/agent.py +++ b/examples/memory_service_with_advanced_memory/agent/agent.py @@ -7,6 +7,7 @@ from trpc_agent_sdk.agents import LlmAgent from trpc_agent_sdk.models import OpenAIModel +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerFilter from .config import get_model_config from .prompts import INSTRUCTION @@ -22,6 +23,7 @@ def create_agent() -> LlmAgent: model_name=model_name, api_key=api_key, base_url=base_url, + filters=[AdvancedAutoCompactSummarizerFilter()], ), instruction=INSTRUCTION, ) diff --git a/examples/memory_service_with_advanced_memory/run_agent.py b/examples/memory_service_with_advanced_memory/run_agent.py index 097a3271d..b1027950c 100644 --- a/examples/memory_service_with_advanced_memory/run_agent.py +++ b/examples/memory_service_with_advanced_memory/run_agent.py @@ -11,26 +11,50 @@ from pathlib import Path from dotenv import load_dotenv -from trpc_agent_sdk.memory import AdvancedMemoryConfig -from trpc_agent_sdk.sessions import AdvancedMemorySessionService +from trpc_agent_sdk.memory.advanced_memory import AdvancedMemoryServiceConfig +from trpc_agent_sdk.memory import AdvancedMemoryService +from trpc_agent_sdk.sessions import InMemorySessionService from trpc_agent_sdk.sessions import SessionServiceConfig +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizer +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerConfig +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerManager from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import Part from agent.agent import create_agent -load_dotenv() +load_dotenv(Path(__file__).with_name(".env"), override=True) -def create_session_service() -> AdvancedMemorySessionService: - """Create the persistent Advanced Memory session service.""" - return AdvancedMemorySessionService( - config=AdvancedMemoryConfig(root_dir=Path(__file__).resolve().parent), - session_config=SessionServiceConfig(ttl=SessionServiceConfig.create_ttl_config( - ttl_seconds=60, - cleanup_interval_seconds=5, - )), +def create_session_service() -> InMemorySessionService: + """Create the session service with the independent Compact manager.""" + compact_manager = AdvancedAutoCompactSummarizerManager( + summarizer=AdvancedAutoCompactSummarizer( + config=AdvancedAutoCompactSummarizerConfig(), + ), ) + return InMemorySessionService( + session_config=SessionServiceConfig( + ttl=SessionServiceConfig.create_ttl_config( + enable=True, + ttl_seconds=60, + cleanup_interval_seconds=5, + ), + store_historical_events=True, + ), + summarizer_manager=compact_manager, + ) + + +def create_memory_service() -> AdvancedMemoryService: + """Create the independent long-term Advanced Memory service.""" + memory_config = AdvancedMemoryServiceConfig( + root_dir=Path(__file__).resolve().parent, + memory_ttl_seconds=120, + memory_focus_instruction=("特别关注并主动记住用户长期稳定的兴趣爱好、" + "编程语言偏好、开发习惯和测试习惯。"), + ) + return AdvancedMemoryService(config=memory_config) async def run_turn(runner, *, user_id: str, session_id: str, prompt: str) -> None: @@ -57,12 +81,14 @@ async def main() -> None: """Run two independent sessions sharing Advanced Memory.""" agent = create_agent() session_service = create_session_service() + memory_service = create_memory_service() from trpc_agent_sdk.runners import Runner runner = Runner( app_name="advanced_memory_demo", agent=agent, session_service=session_service, + memory_service=memory_service, ) try: session_one_prompts = [ @@ -95,9 +121,6 @@ async def main() -> None: prompt="What do you remember about my favorite programming language?", ) - print("\n⏳ Waiting for the session TTL cleanup...") - await asyncio.sleep(125) - print("🧹 Expired Advanced Memory sessions should now be removed.") finally: await runner.close() diff --git a/examples/memory_service_with_advanced_memory_redis/.env b/examples/memory_service_with_advanced_memory_redis/.env new file mode 100644 index 000000000..52b372762 --- /dev/null +++ b/examples/memory_service_with_advanced_memory_redis/.env @@ -0,0 +1,6 @@ +REDIS_URL= + +# Set TRPC_AGENT_API_KEY, TRPC_AGENT_BASE_URL, and TRPC_AGENT_MODEL_NAME. +TRPC_AGENT_API_KEY= +TRPC_AGENT_BASE_URL= +TRPC_AGENT_MODEL_NAME= \ No newline at end of file diff --git a/examples/memory_service_with_advanced_memory_redis/README.md b/examples/memory_service_with_advanced_memory_redis/README.md new file mode 100644 index 000000000..dcd6a2a34 --- /dev/null +++ b/examples/memory_service_with_advanced_memory_redis/README.md @@ -0,0 +1,190 @@ +# Advanced Memory Redis 持久化示例 + +本示例演示如何使用 `AdvancedMemoryService` 将长期记忆保存到 Redis,实现跨会话、跨 Python 进程的持久化记忆。 + +## 关键特性 + +- **主动式记忆**:Agent 根据对话内容主动调用工具保存长期有效的信息。 +- **记忆分类**:每条记忆包含名称、描述、类型、摘要和详细内容。 +- **基于记忆索引的记忆召回**:先读取 Redis 中的 `MEMORY.md` 索引, + 再读取与问题相关的记忆内容。 +- **Redis 持久化**:多个进程或实例使用相同的 Redis、应用名和用户 ID时,可以访问同一份长期记忆。 + +## Agent 层级结构说明 + +`AdvancedMemoryService` 通过 `Runner` 绑定到 Agent,并根据配置使用 Redis 保存记忆索引和记忆主题。Agent 通过三个工具主动管理长期记忆。 + +## 关键代码解释 + +### `save_memory` + +保存或更新一条长期记忆,同时更新 Redis 中的记忆索引。 + +### `list_memory_index` + +读取当前用户的记忆索引,帮助 Agent 找到与当前问题相关的记忆文件。 + +### `read_memory` + +根据索引中的文件名读取完整记忆内容。 + +## 环境要求 + +- Python 3.10 或更高版本 +- 可访问的 Redis 服务 +- 一个可访问的 OpenAI 兼容模型服务 + +**启动本地 Redis:** + +```bash +docker run --name advanced-memory-redis \ + -p 6379:6379 \ + -d redis:7-alpine +``` + +然后在当前目录的 `.env` 中配置: + +```dotenv +REDIS_URL=redis://localhost:6379/0 +``` + +如果容器已经存在,执行: + +```bash +docker start advanced-memory-redis +``` + +检查 Redis: + +```bash +docker exec advanced-memory-redis redis-cli PING +# PONG +``` + +如果使用已有的**远程 Redis 服务**,不需要执行 Docker 命令,只需要在当前目录的`.env` 中配置 Redis 连接信息: + +```dotenv +REDIS_URL=redis://:password@redis.example.com:6379/0 +``` + +如果 Redis 使用 ACL 用户名和密码: + +```dotenv +REDIS_URL=redis://username:password@redis.example.com:6379/0 +``` + +启用 TLS 时使用 `rediss` 协议: + +```dotenv +REDIS_URL=rediss://username:password@redis.example.com:6380/0 +``` + +也可以拆分配置: + +```dotenv +REDIS_HOST=redis.example.com +REDIS_PORT=6379 +REDIS_DB=0 +REDIS_USER=your-user +REDIS_PASSWORD=your-password +REDIS_TLS=false +``` + +代码会优先使用 `REDIS_URL`;未设置时,才会根据这些字段构造连接串。密码包含 `@`、`:`、`/`、`#` 等特殊字符时,需要进行 URL 编码。 + +## 模型配置 + +在当前目录的 `.env` 中配置: + +```dotenv +TRPC_AGENT_API_KEY=your-api-key +TRPC_AGENT_BASE_URL=https://your-llm-endpoint/v1 +TRPC_AGENT_MODEL_NAME=your-model-name +``` + +Redis 配置请参考上面的本地 Redis 或远程 Redis 配置方式。 + +## 代码构建 + +```bash +git clone https://github.com/trpc-group/trpc-agent-python.git +cd trpc-agent-python +./build.sh +source .venv/bin/activate +``` + +如果已经在当前项目中创建了 Python 3.10+ 虚拟环境,也可以直接安装: + +```bash +python -m pip install -e . +``` + +## 运行 + +```bash +cd examples/memory_service_with_advanced_memory_redis +source ../../.venv/bin/activate +python run_agent.py +``` + +脚本会依次启动写入和读取两个独立进程,验证 Redis 中的记忆可以跨进程和不同会话读取。也可以单独运行某个阶段: + +```bash +python run_agent.py --phase write +python run_agent.py --phase read +``` + +## Redis 中的存储 + +记忆索引和主题内容会以 Redis key 保存,key 前缀为: + +```text +advanced-memory-redis-demo:v1:* +``` + +查看本示例写入的 key: + +```bash +docker exec advanced-memory-redis redis-cli --scan \ + --pattern 'advanced-memory-redis-demo:v1:*' +``` + +示例中的记忆 TTL 在代码的 `AdvancedMemoryServiceConfig` 中配置为 `memory_ttl_seconds=120`。 + +## 运行结果(实测) + +```txt +==================== WRITE PROCESS ==================== + +----- Runner A, query 1 ----- + +📝 user: Do you remember my name? +🤖 Assistant: + + +🔧 tool call: list_memory_index({}) +📊 Tool Result: {'index_path': 'advanced-memory://redis/advanced-memory-redis-demo:v1:{advanced-memory-redis-demo:redis-demo-user}:memory:MEMORY.md', 'index': ''} +🤖 Assistant: I checked my long-term memory, but it looks like I don't have any record of your name yet — my memory index is currently empty. + +If you'd like, tell me your name (or anything else you'd like me to remember about you), and I'll save it for future conversations. 😊 + +...... + +==================== READ PROCESS ==================== + +----- Runner B, query 1 ----- + +📝 user: Do you remember my name? +🔧 tool call: read_memory({'filename': 'user-identity.md'}) +📊 Tool Result: {'found': True, 'filename': 'user-identity.md', 'content': "---\nname: User identity\ndescription: Alice's name and basic identity for personalization.\ntype: user\nupdated_at: 2026-09-11T05:54:20.441889+00:00\n---\nThe user's name is Alice. She introduced herself on first contact. Use this name for personalized responses.\n", 'updated_at': '2026-09-11T05:54:20.441889+00:00', 'freshness': 'today', 'freshness_notice': 'This memory was last updated today. It is a point-in-time observation and may no longer reflect the current state. Verify it when necessary, and update this memory if it is outdated or incorrect.'} +🤖 Assistant: Yes, I remember! Your name is **Alice**. 😊 + +I've stored that in my long-term memory so I can personalize my responses for you. Is there anything else I can help you with? + +----- Runner B, query 2 ----- + +📝 user: Do you remember my favorite color? +🔧 tool call: read_memory({'filename': 'favorite-color.md'}) +📊 Tool Result: {'found': True, 'filename': 'favorite-color.md', 'content': "---\nname: Favorite color\ndescription: Alice's favorite color.\ntype: user\nupdated_at: 2026-09-11T05:54:24.620584+00:00\n---\nAlice's favorite color is blue.\n", 'updated_at': '2026-09-11T05:54:24.620584+00:00', 'freshness': 'today', 'freshness_notice': 'This memory was last updated today. It is a point-in-time observation and may no longer reflect the current state. Verify it when necessary, and update this memory if it is outdated or incorrect.'} +🤖 Assistant: Yes, I remember! Your favorite color is **blue**. 💙 +``` \ No newline at end of file diff --git a/examples/memory_service_with_advanced_memory_redis/agent/__init__.py b/examples/memory_service_with_advanced_memory_redis/agent/__init__.py new file mode 100644 index 000000000..ee02e466a --- /dev/null +++ b/examples/memory_service_with_advanced_memory_redis/agent/__init__.py @@ -0,0 +1 @@ +"""Agent package for the Redis Advanced Memory example.""" diff --git a/examples/memory_service_with_advanced_memory_redis/agent/agent.py b/examples/memory_service_with_advanced_memory_redis/agent/agent.py new file mode 100644 index 000000000..633f5009e --- /dev/null +++ b/examples/memory_service_with_advanced_memory_redis/agent/agent.py @@ -0,0 +1,27 @@ +"""Agent definition for the Redis Advanced Memory example.""" + +import os + +from trpc_agent_sdk.agents import LlmAgent +from trpc_agent_sdk.models import OpenAIModel +from trpc_agent_sdk.tools import FunctionTool + +from .tools import get_weather_report + + +def create_agent() -> LlmAgent: + """Create an agent whose Runner installs Advanced Memory tools.""" + api_key = os.getenv("TRPC_AGENT_API_KEY", "") + base_url = os.getenv("TRPC_AGENT_BASE_URL", "") + model_name = os.getenv("TRPC_AGENT_MODEL_NAME", "") + if not api_key or not base_url or not model_name: + raise ValueError("TRPC_AGENT_API_KEY, TRPC_AGENT_BASE_URL, and TRPC_AGENT_MODEL_NAME must be set") + return LlmAgent( + name="advanced_memory_redis_assistant", + description="A Redis-backed Advanced Memory demonstration assistant", + model=OpenAIModel(model_name=model_name, api_key=api_key, base_url=base_url), + instruction=("When the user asks you to remember a durable personal preference or fact, use save_memory. " + "When the user asks what you remember, use list_memory_index first and read_memory for the " + "relevant file. Always answer using the tool result."), + tools=[FunctionTool(get_weather_report)], + ) diff --git a/examples/memory_service_with_advanced_memory_redis/agent/tools.py b/examples/memory_service_with_advanced_memory_redis/agent/tools.py new file mode 100644 index 000000000..98f84225e --- /dev/null +++ b/examples/memory_service_with_advanced_memory_redis/agent/tools.py @@ -0,0 +1,21 @@ +"""Tools for the Advanced Memory Redis example.""" + + +def get_weather_report(city: str) -> dict: + """Return a small deterministic weather report for a city.""" + if city.lower() == "london": + return { + "status": + "success", + "report": ("The current weather in London is cloudy with a temperature of " + "18 degrees Celsius and a chance of rain."), + } + if city.lower() == "paris": + return { + "status": "success", + "report": "The weather in Paris is sunny with a temperature of 25 degrees Celsius.", + } + return { + "status": "error", + "error_message": f"Weather information for '{city}' is not available.", + } diff --git a/examples/memory_service_with_advanced_memory_redis/run_agent.py b/examples/memory_service_with_advanced_memory_redis/run_agent.py new file mode 100644 index 000000000..d33e031af --- /dev/null +++ b/examples/memory_service_with_advanced_memory_redis/run_agent.py @@ -0,0 +1,130 @@ +#!/usr/bin/env python3 +"""Run twice to verify Redis Advanced Memory survives process restarts.""" + +from __future__ import annotations + +import asyncio +import argparse +import os +import subprocess +import sys +from pathlib import Path +from urllib.parse import quote + +from dotenv import load_dotenv + +from agent.agent import create_agent +from trpc_agent_sdk.memory.advanced_memory import AdvancedMemoryServiceConfig +from trpc_agent_sdk.memory import AdvancedMemoryService +from trpc_agent_sdk.runners import Runner +from trpc_agent_sdk.sessions import InMemorySessionService +from trpc_agent_sdk.types import Content, Part + +load_dotenv(Path(__file__).with_name(".env"), override=True) + +RUNNER_A_QUERIES = [ + "Do you remember my name?", + "Do you remember my favorite color?", + "what is the weather like in paris?", + "Hello! My name is Alice. What's your name?", + "Do you remember my name?", + "Hello! My favorite color is blue. What's your favorite color?", + "Do you remember my favorite color?", +] + +RUNNER_B_QUERIES = [ + "Do you remember my name?", + "Do you remember my favorite color?", +] + + +def build_redis_url_from_environment() -> str: + """Use REDIS_URL directly, or construct it from standard Redis variables.""" + redis_url = os.getenv("REDIS_URL") + if redis_url: + return redis_url + + host = os.getenv("REDIS_HOST", "127.0.0.1") + port = os.getenv("REDIS_PORT", "6379") + database = os.getenv("REDIS_DB", "0") + username = os.getenv("REDIS_USER", "") + password = os.getenv("REDIS_PASSWORD", "") + scheme = "rediss" if os.getenv("REDIS_TLS", "").lower() in {"1", "true", "yes"} else "redis" + + if username and password: + auth = f"{quote(username, safe='')}:{quote(password, safe='')}@" + elif password: + auth = f":{quote(password, safe='')}@" + else: + auth = "" + return f"{scheme}://{auth}{host}:{port}/{database}" + + +def create_advanced_memory_service(redis_url: str) -> AdvancedMemoryService: + """Create the long-term Advanced Memory service backed by Redis.""" + config = AdvancedMemoryServiceConfig( + storage_backend="redis", + redis_url=redis_url, + redis_key_prefix="advanced-memory-redis-demo:v1", + memory_ttl_seconds=120, + ) + return AdvancedMemoryService(config=config) + + +async def ask(runner: Runner, session_id: str, prompt: str) -> None: + """Send one message through the shared app and user identity.""" + print(f"\n📝 user: {prompt}") + async for event in runner.run_async( + user_id="redis-demo-user", + session_id=session_id, + new_message=Content(parts=[Part.from_text(text=prompt)]), + ): + if not event.content or not event.content.parts: + continue + for part in event.content.parts: + if part.function_call: + print(f"🔧 tool call: {part.function_call.name}({part.function_call.args})") + elif part.function_response: + print(f"📊 Tool Result: {part.function_response.response}") + elif not event.partial and part.text and not part.thought: + print(f"🤖 Assistant: {part.text}") + + +async def run_phase(phase: str) -> None: + """Run Runner A or Runner B against the same Redis user.""" + app_name = "advanced-memory-redis-demo" + redis_url = build_redis_url_from_environment() + runner = Runner( + app_name=app_name, + agent=create_agent(), + session_service=InMemorySessionService(), + memory_service=create_advanced_memory_service(redis_url), + ) + try: + queries = RUNNER_A_QUERIES if phase == "write" else RUNNER_B_QUERIES + runner_name = "A" if phase == "write" else "B" + for index, prompt in enumerate(queries): + print(f"\n----- Runner {runner_name}, query {index + 1} -----") + await ask(runner, f"redis-{phase}-session-{index}", prompt) + finally: + await runner.close() + + +def run_two_processes() -> None: + """Start fresh writer and reader processes to prove Redis persistence.""" + for phase in ("write", "read"): + print(f"\n{'=' * 20} {phase.upper()} PROCESS {'=' * 20}", flush=True) + subprocess.run( + [sys.executable, str(Path(__file__).resolve()), "--phase", phase], + check=True, + ) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--phase", choices=("write", "read")) + arguments = parser.parse_args() + if arguments.phase: + asyncio.run(run_phase(arguments.phase)) + else: + run_two_processes() diff --git a/examples/memory_service_with_advanced_memory_sql/.env b/examples/memory_service_with_advanced_memory_sql/.env new file mode 100644 index 000000000..5e8b0ded0 --- /dev/null +++ b/examples/memory_service_with_advanced_memory_sql/.env @@ -0,0 +1,12 @@ +# Model configuration +TRPC_AGENT_API_KEY= +TRPC_AGENT_BASE_URL= +TRPC_AGENT_MODEL_NAME= + +# Easy local test with SQLite. SQL_IS_ASYNC=false uses the built-in sqlite driver. +SQL_URL=sqlite:///advanced-memory-sql-demo.db +SQL_IS_ASYNC=false + +# For MySQL, replace SQL_URL and set SQL_IS_ASYNC=true: +# SQL_URL= +# SQL_IS_ASYNC=true diff --git a/examples/memory_service_with_advanced_memory_sql/README.md b/examples/memory_service_with_advanced_memory_sql/README.md new file mode 100644 index 000000000..86f927347 --- /dev/null +++ b/examples/memory_service_with_advanced_memory_sql/README.md @@ -0,0 +1,119 @@ +# Advanced Memory SQL 持久化示例 + +本示例演示如何使用 `AdvancedMemoryService` 将长期记忆保存到 SQL 数据库,实现跨会话、跨 Python 进程的持久化记忆。 + +## 关键特性 + +- **主动式记忆**:Agent 根据对话内容主动调用工具保存长期有效的信息。 +- **记忆分类**:每条记忆包含名称、描述、类型、摘要和详细内容。 +- **基于记忆索引的记忆召回**:先读取数据库中的记忆索引,再读取与问题相关的记忆内容。 +- **SQL 持久化**:多个进程或实例使用相同的数据库、应用名和用户 ID 时,可以访问同一份长期记忆。 + +## Agent 层级结构说明 + +`AdvancedMemoryService` 通过 `Runner` 绑定到 Agent,并根据配置使用 SQL 保存记忆索引和记忆主题。Agent 通过三个工具主动管理长期记忆。 + +## 关键代码解释 + +### `save_memory` + +保存或更新一条长期记忆,同时更新数据库中的记忆索引。 + +### `list_memory_index` + +读取当前用户的记忆索引,帮助 Agent 找到与当前问题相关的记忆文件。 + +### `read_memory` + +根据索引中的文件名读取完整记忆内容。 + +## 环境要求 + +- Python 3.10 或更高版本 +- SQLite 或可访问的 MySQL 数据库 +- 一个可访问的 OpenAI 兼容模型服务 + +默认使用 **SQLite**,不需要额外启动数据库: + +```dotenv +SQL_URL=sqlite:///advanced-memory-sql-demo.db +SQL_IS_ASYNC=false +``` + +使用 **MySQL** 时: + +```dotenv +SQL_URL=mysql+aiomysql://user:password@localhost:3306/trpc_agent_advanced_memory?charset=utf8mb4 +SQL_IS_ASYNC=true +``` + +## SQL 配置 + +在当前目录的 `.env` 中配置数据库和模型: + +```dotenv +SQL_URL=sqlite:///advanced-memory-sql-demo.db +SQL_IS_ASYNC=false +TRPC_AGENT_API_KEY=your-api-key +TRPC_AGENT_BASE_URL=https://your-llm-endpoint/v1 +TRPC_AGENT_MODEL_NAME=your-model-name +``` + +也可以使用 `MYSQL_USER`、`MYSQL_PASSWORD`、`MYSQL_HOST`、`MYSQL_PORT` 和 `MYSQL_DB` 由脚本构造 MySQL 连接串。 + +## 代码构建 + +```bash +git clone https://github.com/trpc-group/trpc-agent-python.git +cd trpc-agent-python +./build.sh +source .venv/bin/activate +``` + +如果已经在当前项目中创建了 Python 3.10+ 虚拟环境,也可以直接安装: + +```bash +python -m pip install -e . +``` + +## 运行 + +```bash +cd examples/memory_service_with_advanced_memory_sql +source ../../.venv/bin/activate +python run_agent.py +``` + +脚本会依次启动写入和读取两个独立进程,验证数据库中的记忆可以跨进程和不同会话读取。也可以单独运行某个阶段: + +```bash +python run_agent.py --phase write +python run_agent.py --phase read +``` + +首次运行时,SQLite 数据库文件和 Advanced Memory 数据表会自动创建。示例中的记忆 TTL 在代码的 `AdvancedMemoryServiceConfig` 中配置为`memory_ttl_seconds=120`。 + +## 运行结果(实测) + +```txt +=================== WRITE PROCESS ==================== + +----- Runner A, query 1 ----- +📝 user: Do you remember my name? +🔧 tool call: list_memory_index({}) +📊 Tool Result: {'index_path': 'advanced-memory://sql/advanced-memory-sql-demo/sql-demo-user/memory/MEMORY.md', 'index': ''} +🤖 Assistant: Let me check my long-term memory. +🔧 tool call: list_memory_index({}) +📊 Tool Result: {'index_path': 'advanced-memory://sql/advanced-memory-sql-demo/sql-demo-user/memory/MEMORY.md', 'index': ''} +🤖 Assistant: I checked my long-term memory, but it's currently empty — I don't have any stored details about you yet, including your name. 😊 + +If you'd like me to remember it for future conversations, just tell me your name (and anything else you'd like me to keep in mind, like preferences or context), and I'll save it right away. + +... + +----- Runner B, query 2 ----- +📝 user: Do you remember my favorite color? +🔧 tool call: read_memory({'filename': 'user-profile.md'}) +📊 Tool Result: {'found': True, 'filename': 'user-profile.md', 'content': "---\nname: User profile\ndescription: Basic identity and preferences of the user.\ntype: user\nupdated_at: 2026-09-11T06:00:15.755769+00:00\n---\n---\nname: User profile\ndescription: Basic identity and preferences of the user.\ntype: user\nupdated_at: 2026-09-11T06:00:09.820671+00:00\n---\nThe user's name is Alice. She introduced herself at the start of ourfirst conversation. Her favorite color is blue, which she shared in a later conversation.\n", 'updated_at': '2026-09-11T06:00:15.755769+00:00', 'freshness': 'today', 'freshness_notice': 'This memory was last updated today. It is a point-in-time observation and may no longer reflect the current state. Verify it when necessary, and update this memory if it is outdated or incorrect.'} +🤖 Assistant: Yes, I remember! Your favorite color is **blue** — you shared that with me in a later conversation, Alice. 💙 +``` \ No newline at end of file diff --git a/examples/memory_service_with_advanced_memory_sql/agent/__init__.py b/examples/memory_service_with_advanced_memory_sql/agent/__init__.py new file mode 100644 index 000000000..3b7ed6716 --- /dev/null +++ b/examples/memory_service_with_advanced_memory_sql/agent/__init__.py @@ -0,0 +1 @@ +"""Agent package for the Advanced Memory SQL example.""" diff --git a/examples/memory_service_with_advanced_memory_sql/agent/agent.py b/examples/memory_service_with_advanced_memory_sql/agent/agent.py new file mode 100644 index 000000000..94532abcb --- /dev/null +++ b/examples/memory_service_with_advanced_memory_sql/agent/agent.py @@ -0,0 +1,25 @@ +"""Agent definition for the Advanced Memory SQL example.""" + +from trpc_agent_sdk.agents import LlmAgent +from trpc_agent_sdk.models import OpenAIModel +from trpc_agent_sdk.tools import FunctionTool + +from .config import get_model_config +from .prompts import INSTRUCTION +from .tools import get_weather_report + + +def create_agent() -> LlmAgent: + """Create an agent; Runner installs the Advanced Memory tools.""" + api_key, base_url, model_name = get_model_config() + return LlmAgent( + name="advanced_memory_sql_assistant", + description="A minimal Advanced Memory SQL demonstration assistant", + model=OpenAIModel( + model_name=model_name, + api_key=api_key, + base_url=base_url, + ), + instruction=INSTRUCTION, + tools=[FunctionTool(get_weather_report)], + ) diff --git a/examples/memory_service_with_advanced_memory_sql/agent/config.py b/examples/memory_service_with_advanced_memory_sql/agent/config.py new file mode 100644 index 000000000..a9ef0c1bf --- /dev/null +++ b/examples/memory_service_with_advanced_memory_sql/agent/config.py @@ -0,0 +1,14 @@ +"""Model configuration for the Advanced Memory SQL example.""" + +import os + + +def get_model_config() -> tuple[str, str, str]: + """Read the model configuration from the environment.""" + api_key = os.getenv("TRPC_AGENT_API_KEY", "") + base_url = os.getenv("TRPC_AGENT_BASE_URL", "") + model_name = os.getenv("TRPC_AGENT_MODEL_NAME", "") + if not api_key or not base_url or not model_name: + raise ValueError("TRPC_AGENT_API_KEY, TRPC_AGENT_BASE_URL, and " + "TRPC_AGENT_MODEL_NAME must be set") + return api_key, base_url, model_name diff --git a/examples/memory_service_with_advanced_memory_sql/agent/prompts.py b/examples/memory_service_with_advanced_memory_sql/agent/prompts.py new file mode 100644 index 000000000..93966f933 --- /dev/null +++ b/examples/memory_service_with_advanced_memory_sql/agent/prompts.py @@ -0,0 +1,8 @@ +"""Prompt for the Advanced Memory SQL example.""" + +INSTRUCTION = """You are a helpful assistant demonstrating Advanced Memory. + +When the user asks you to remember a durable personal preference or fact, use +save_memory. When the user asks what you remember, use list_memory_index first +and read_memory for the relevant file. Always answer using the tool result. +""" diff --git a/examples/memory_service_with_advanced_memory_sql/agent/tools.py b/examples/memory_service_with_advanced_memory_sql/agent/tools.py new file mode 100644 index 000000000..cb75e7e0b --- /dev/null +++ b/examples/memory_service_with_advanced_memory_sql/agent/tools.py @@ -0,0 +1,21 @@ +"""Tools for the Advanced Memory SQL example.""" + + +def get_weather_report(city: str) -> dict: + """Return a small deterministic weather report for a city.""" + if city.lower() == "london": + return { + "status": + "success", + "report": ("The current weather in London is cloudy with a temperature of " + "18 degrees Celsius and a chance of rain."), + } + if city.lower() == "paris": + return { + "status": "success", + "report": "The weather in Paris is sunny with a temperature of 25 degrees Celsius.", + } + return { + "status": "error", + "error_message": f"Weather information for '{city}' is not available.", + } diff --git a/examples/memory_service_with_advanced_memory_sql/run_agent.py b/examples/memory_service_with_advanced_memory_sql/run_agent.py new file mode 100644 index 000000000..d0484cb6a --- /dev/null +++ b/examples/memory_service_with_advanced_memory_sql/run_agent.py @@ -0,0 +1,121 @@ +#!/usr/bin/env python3 +"""Run the Advanced Memory SQL persistence example.""" + +from __future__ import annotations + +import argparse +import asyncio +import os +import subprocess +import sys +from pathlib import Path +from urllib.parse import quote + +from dotenv import load_dotenv + +from agent.agent import create_agent +from trpc_agent_sdk.memory.advanced_memory import AdvancedMemoryServiceConfig +from trpc_agent_sdk.memory import AdvancedMemoryService +from trpc_agent_sdk.runners import Runner +from trpc_agent_sdk.sessions import InMemorySessionService +from trpc_agent_sdk.types import Content, Part + +load_dotenv(Path(__file__).with_name(".env"), override=True) + +RUNNER_A_QUERIES = [ + "Do you remember my name?", + "Do you remember my favorite color?", + "what is the weather like in paris?", + "Hello! My name is Alice. What's your name?", + "Do you remember my name?", + "Hello! My favorite color is blue. What's your favorite color?", + "Do you remember my favorite color?", +] + +RUNNER_B_QUERIES = [ + "Do you remember my name?", + "Do you remember my favorite color?", +] + + +def build_sql_url_from_environment() -> str: + """Use SQL_URL or build a MySQL URL from standard environment variables.""" + sql_url = os.getenv("SQL_URL") + if sql_url: + return sql_url + + user = quote(os.getenv("MYSQL_USER", "root"), safe="") + password = quote(os.getenv("MYSQL_PASSWORD", ""), safe="") + host = os.getenv("MYSQL_HOST", "127.0.0.1") + port = os.getenv("MYSQL_PORT", "3306") + database = os.getenv("MYSQL_DB", "trpc_agent_advanced_memory") + return f"mysql+aiomysql://{user}:{password}@{host}:{port}/{database}?charset=utf8mb4" + + +def sql_is_async() -> bool: + """Return whether the configured SQL driver is asynchronous.""" + return os.getenv("SQL_IS_ASYNC", "true").lower() in {"1", "true", "yes"} + + +def create_advanced_memory_service(sql_url: str) -> AdvancedMemoryService: + """Create the long-term Advanced Memory service backed by SQL.""" + config = AdvancedMemoryServiceConfig( + storage_backend="sql", + sql_url=sql_url, + sql_is_async=sql_is_async(), + memory_ttl_seconds=120, + ) + return AdvancedMemoryService(config=config) + + +async def run_phase(phase: str) -> None: + """Run Runner A or Runner B against the same SQL database.""" + sql_url = build_sql_url_from_environment() + runner = Runner( + app_name="advanced-memory-sql-demo", + agent=create_agent(), + session_service=InMemorySessionService(), + memory_service=create_advanced_memory_service(sql_url), + ) + try: + queries = RUNNER_A_QUERIES if phase == "write" else RUNNER_B_QUERIES + runner_name = "A" if phase == "write" else "B" + for index, prompt in enumerate(queries): + print(f"\n----- Runner {runner_name}, query {index + 1} -----") + print(f"📝 user: {prompt}") + async for event in runner.run_async( + user_id="sql-demo-user", + session_id=f"sql-{phase}-session-{index}", + new_message=Content(parts=[Part.from_text(text=prompt)]), + ): + if not event.content or not event.content.parts: + continue + for part in event.content.parts: + if part.function_call: + print(f"🔧 tool call: {part.function_call.name}({part.function_call.args})") + elif part.function_response: + print(f"📊 Tool Result: {part.function_response.response}") + elif not event.partial and part.text and not part.thought: + print(f"🤖 Assistant: {part.text}") + finally: + await runner.close() + + +def run_two_processes() -> None: + """Start independent writer and reader processes.""" + for phase in ("write", "read"): + print(f"\n{'=' * 20} {phase.upper()} PROCESS {'=' * 20}", flush=True) + subprocess.run( + [sys.executable, str(Path(__file__).resolve()), "--phase", phase], + check=True, + ) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--phase", choices=("write", "read")) + args = parser.parse_args() + if args.phase: + asyncio.run(run_phase(args.phase)) + else: + run_two_processes() diff --git a/examples/session_summarizer_with_advanced/.env b/examples/session_summarizer_with_advanced/.env new file mode 100644 index 000000000..dc791393a --- /dev/null +++ b/examples/session_summarizer_with_advanced/.env @@ -0,0 +1,4 @@ +# Set TRPC_AGENT_API_KEY、TRPC_AGENT_BASE_URL、TRPC_AGENT_MODEL_NAME +TRPC_AGENT_API_KEY=your-api-key +TRPC_AGENT_BASE_URL=your-base-url +TRPC_AGENT_MODEL_NAME=your-model-name diff --git a/examples/session_summarizer_with_advanced/README.md b/examples/session_summarizer_with_advanced/README.md new file mode 100644 index 000000000..77cd9d96e --- /dev/null +++ b/examples/session_summarizer_with_advanced/README.md @@ -0,0 +1,91 @@ +# Advanced Session Summarizer 示例 + +本示例验证 `AdvancedAutoCompactSummarizer` 的模型调用前压缩流程。示例执行 5 轮真实多轮对话,随着历史增长,内置 Model Filter 会在每次请求模型前检查并压缩上下文。 + +## 验证内容 + +- 模型显式配置 `AdvancedAutoCompactSummarizerFilter` +- 使用较小的字符阈值,在前几轮内触发自动压缩 +- 压缩后的模型可见窗口以 summary event 开头 +- 被替换的原始 events 移入 `historical_events` +- Session Memory 写入 `session.state` +- 压缩后对话继续进行,模型仍可依据摘要回答历史事实 + +## 为什么必须用真实多轮对话 + +压缩边界基于模型请求中的 content 数量计算(`keep_recent_contents`)。框架会把非当前 Agent 产出的历史事件转换为 user 角色,并合并相邻同角色 content。手工塞入 `author="assistant"` 的伪造事件会被合并成单条 user content,导致找不到压缩边界并报 `Not enough model contents to compact`。因此示例通过真实回合累积 user/model 交替历史。 + +## 组件关系 + +```text +OpenAIModel +└── AdvancedAutoCompactSummarizerFilter(模型调用前触发) + +InMemorySessionService +└── AdvancedAutoCompactSummarizerManager + └── AdvancedAutoCompactSummarizer +``` + +Filter 必须显式安装到模型上,Manager 不会自动修改 Agent 或 Model。 + +## 环境要求 + +- Python3.10+,推荐 Python3.12 + +## 构建步骤 + +```bash +git clone https://github.com/trpc-group/trpc-agent-python.git +cd trpc-agent-python +./build.sh +source .venv/bin/activate +``` + +## 运行步骤 + +### 配置环境变量 + +在当前目录的 `.env` 中配置(或通过 `export` 设置): + +```dotenv +TRPC_AGENT_API_KEY=your-api-key +TRPC_AGENT_BASE_URL=https://your-openai-compatible-endpoint/v1 +TRPC_AGENT_MODEL_NAME=your-model-name +``` + +### 运行命令 + +```bash +cd examples/session_summarizer_with_advanced +python3 run_agent.py +``` + +## 预期结果 + +```text +After turn 2 + active events: 4 + historical events: 0 + active window starts with summary: False + session memory persisted: True + +After turn 3 + active events: 5 + historical events: 3 + active window starts with summary: True + session memory persisted: True + +After turn 5 + active events: 5 + historical events: 10 + active window starts with summary: True + session memory persisted: True + +PASS: compaction ran on turn(s) [3, 4, 5]. +``` + +关键现象是 `active events` 稳定在一个小窗口,而 `historical events` 持续增长——说明模型可见上下文被压缩,原始事件仍可完整追溯。具体轮次和数量取决于模型回复长度。若模型未按压缩提示词返回 `` 块,示例会在结束时明确报错。 + +## 调整阈值 + +示例为便于观察而关闭 token 模式,使用 `trigger_chars=4000`。生产环境可启用 `TokenContextTrackerConfig` 并设置模型上下文窗口,或根据业务规模提高字符阈值。 diff --git a/examples/session_summarizer_with_advanced/agent/__init__.py b/examples/session_summarizer_with_advanced/agent/__init__.py new file mode 100644 index 000000000..bc6e483f9 --- /dev/null +++ b/examples/session_summarizer_with_advanced/agent/__init__.py @@ -0,0 +1,5 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. diff --git a/examples/session_summarizer_with_advanced/agent/agent.py b/examples/session_summarizer_with_advanced/agent/agent.py new file mode 100644 index 000000000..13bb361a2 --- /dev/null +++ b/examples/session_summarizer_with_advanced/agent/agent.py @@ -0,0 +1,38 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Agent used by the advanced session compaction example.""" + +from trpc_agent_sdk.agents import LlmAgent +from trpc_agent_sdk.models import LLMModel +from trpc_agent_sdk.models import OpenAIModel +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerFilter + +from .config import get_model_config +from .prompts import INSTRUCTION + + +def _create_model() -> LLMModel: + """Create a model with explicit before-model compaction filtering.""" + api_key, url, model_name = get_model_config() + return OpenAIModel( + model_name=model_name, + api_key=api_key, + base_url=url, + filters=[AdvancedAutoCompactSummarizerFilter()], + ) + + +def create_agent() -> LlmAgent: + """Create the Python tutor agent.""" + return LlmAgent( + name="python_tutor", + description="Python programming tutor that helps users learn Python", + model=_create_model(), + instruction=INSTRUCTION, + ) + + +root_agent = create_agent() diff --git a/examples/session_summarizer_with_advanced/agent/config.py b/examples/session_summarizer_with_advanced/agent/config.py new file mode 100644 index 000000000..db0d491b8 --- /dev/null +++ b/examples/session_summarizer_with_advanced/agent/config.py @@ -0,0 +1,19 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +""" Agent config module""" + +import os + + +def get_model_config() -> tuple[str, str, str]: + """Get model config from environment variables""" + api_key = os.getenv('TRPC_AGENT_API_KEY', '') + url = os.getenv('TRPC_AGENT_BASE_URL', '') + model_name = os.getenv('TRPC_AGENT_MODEL_NAME', '') + if not api_key or not url or not model_name: + raise ValueError('''TRPC_AGENT_API_KEY, TRPC_AGENT_BASE_URL, + and TRPC_AGENT_MODEL_NAME must be set in environment variables''') + return api_key, url, model_name diff --git a/examples/session_summarizer_with_advanced/agent/prompts.py b/examples/session_summarizer_with_advanced/agent/prompts.py new file mode 100644 index 000000000..c64e864f5 --- /dev/null +++ b/examples/session_summarizer_with_advanced/agent/prompts.py @@ -0,0 +1,15 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +""" prompts for agent""" + +INSTRUCTION = """You are a professional Python programming tutor. Your tasks are: +1. Answer users' Python-related questions patiently +2. Provide clear explanations and example code +3. Adjust teaching difficulty based on the user's progress +4. Encourage practice and questions + +Communicate in a friendly, professional manner. +""" diff --git a/examples/session_summarizer_with_advanced/agent/tools.py b/examples/session_summarizer_with_advanced/agent/tools.py new file mode 100644 index 000000000..16e188b49 --- /dev/null +++ b/examples/session_summarizer_with_advanced/agent/tools.py @@ -0,0 +1,6 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Tools for the advanced session compaction example.""" diff --git a/examples/session_summarizer_with_advanced/run_agent.py b/examples/session_summarizer_with_advanced/run_agent.py new file mode 100644 index 000000000..bc82216e3 --- /dev/null +++ b/examples/session_summarizer_with_advanced/run_agent.py @@ -0,0 +1,173 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Demonstrate Advanced session compaction running before each model call.""" + +from __future__ import annotations + +import asyncio +import os +import uuid +from pathlib import Path + +from dotenv import load_dotenv + +from trpc_agent_sdk.models import LLMModel +from trpc_agent_sdk.runners import Runner +from trpc_agent_sdk.sessions import InMemorySessionService +from trpc_agent_sdk.sessions import Session +from trpc_agent_sdk.sessions import SessionServiceConfig +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizer +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerConfig +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerManager +from trpc_agent_sdk.sessions.compact import AutoCompactSummarizerConfig +from trpc_agent_sdk.sessions.compact import SessionMemoryExtractorConfig +from trpc_agent_sdk.sessions.compact import TokenContextTrackerConfig +from trpc_agent_sdk.types import Content +from trpc_agent_sdk.types import Part + +load_dotenv(Path(__file__).with_name(".env")) + +SESSION_MEMORY_STATE_KEY = "_trpc_agent:summary" + +PROJECT_BRIEF = """I am building Project Apollo and I want you to remember its constraints: +- Runtime: Python 3.12 with asyncio everywhere, no blocking calls in request paths. +- Web layer: FastAPI with Pydantic models for every request and response body. +- Storage: SQLite through SQLAlchemy, and every failed write must roll back. +- Testing: pytest with async tests covering each persistence path. +- Style: full type hints, small functions, no bare except. +- Deployment: Linux containers, released every Friday afternoon. +""" + +# Multi-turn conversation. Real turns are what build alternating user/model +# history, which is the shape the compaction boundary is computed from. +CONVERSATIONS = ( + PROJECT_BRIEF + "\nAcknowledge the constraints and outline the module layout you would use.", + "Show me the SQLAlchemy session helper for Project Apollo, " + "including how rollback is handled on a failed write.", + "Now show the pytest fixtures and one async test that proves the rollback path works.", + "Explain how I should structure the FastAPI routers and dependency injection for this project.", + "Recap Project Apollo: its stack, storage rules, testing rules, and release cadence.", +) + + +def create_compact_config() -> AdvancedAutoCompactSummarizerConfig: + """Use small character budgets so the example compacts within a few turns.""" + return AdvancedAutoCompactSummarizerConfig( + # Character thresholds keep the demo independent of any model's + # context window. Production setups usually enable the token tracker. + token_context_tracker=TokenContextTrackerConfig(enabled=False), + session_memory=SessionMemoryExtractorConfig( + initial_chars=1_000, + update_chars=500, + ), + auto_compact=AutoCompactSummarizerConfig( + trigger_chars=4_000, + target_chars=2_000, + blocking_chars=20_000, + keep_recent_contents=2, + summary_input_max_chars=12_000, + ), + ) + + +def create_summarizer_manager(model: LLMModel) -> AdvancedAutoCompactSummarizerManager: + """Create the Advanced summarizer that the model filter drives.""" + summarizer = AdvancedAutoCompactSummarizer( + config=create_compact_config(), + model=model, + ) + return AdvancedAutoCompactSummarizerManager(summarizer=summarizer) + + +def print_session_state(label: str, session: Session) -> None: + """Print the state needed to verify compaction behavior.""" + summary_anchor = bool(session.events and session.events[0].is_summary_event()) + print(f"\n{label}") + print(f" active events: {len(session.events)}") + print(f" historical events: {len(session.historical_events)}") + print(f" active window starts with summary: {summary_anchor}") + print(f" session memory persisted: {SESSION_MEMORY_STATE_KEY in session.state}") + + +async def run_turn(runner: Runner, user_id: str, session_id: str, prompt: str) -> None: + """Send one user message and stream the answer.""" + print(f"\nUser: {prompt.splitlines()[0]}") + print("Assistant: ", end="", flush=True) + async for event in runner.run_async( + user_id=user_id, + session_id=session_id, + new_message=Content(parts=[Part.from_text(text=prompt)]), + ): + if not event.content or not event.content.parts: + continue + for part in event.content.parts: + if part.text and not part.thought: + print(part.text, end="" if event.partial else "\n", flush=True) + + +async def main() -> None: + """Run a multi-turn conversation and verify before-model Advanced AutoCompact.""" + app_name = "advanced-session-summarizer-demo" + user_id = "demo-user" + session_id = os.getenv("SESSION_ID", str(uuid.uuid4())) + + # Import after load_dotenv so the module-level Agent can read model settings. + from agent.agent import root_agent + + manager = create_summarizer_manager(root_agent.model) + session_service = InMemorySessionService( + summarizer_manager=manager, + session_config=SessionServiceConfig(store_historical_events=True), + ) + runner = Runner( + app_name=app_name, + agent=root_agent, + session_service=session_service, + ) + + try: + await session_service.create_session( + app_name=app_name, + user_id=user_id, + session_id=session_id, + ) + print(f"Session: {app_name}/{user_id}/{session_id}") + + compaction_turns: list[int] = [] + historical_count = 0 + for index, prompt in enumerate(CONVERSATIONS, start=1): + await run_turn(runner, user_id, session_id, prompt) + stored = await session_service.get_session( + app_name=app_name, + user_id=user_id, + session_id=session_id, + ) + if stored is None: + raise RuntimeError("Session disappeared during the conversation") + print_session_state(f"After turn {index}", stored) + if len(stored.historical_events) > historical_count: + compaction_turns.append(index) + historical_count = len(stored.historical_events) + + if not compaction_turns: + raise RuntimeError( + "Compaction never ran. Confirm that the model installs " + "AdvancedAutoCompactSummarizerFilter, that the conversation " + "exceeds auto_compact.trigger_chars, and that the model returns " + "the requested block." + ) + if not stored.events[0].is_summary_event(): + raise RuntimeError("Compaction ran but the active window lost its summary anchor") + + print(f"\nPASS: compaction ran on turn(s) {compaction_turns}.") + print(f" archived events: {len(stored.historical_events)}") + print(f" model-visible events: {len(stored.events)}") + finally: + await runner.close() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/advanced_memory/test_advanced_memory_session_service.py b/tests/advanced_memory/test_advanced_memory_session_service.py deleted file mode 100644 index 3854dd2e1..000000000 --- a/tests/advanced_memory/test_advanced_memory_session_service.py +++ /dev/null @@ -1,221 +0,0 @@ -"""Tests for the standalone Advanced Memory SessionService.""" - -from __future__ import annotations - -import asyncio -import json -from pathlib import Path -from types import SimpleNamespace - -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig -from trpc_agent_sdk.events import Event -from trpc_agent_sdk.runners import Runner -from trpc_agent_sdk.sessions import AdvancedMemorySessionService -from trpc_agent_sdk.sessions import SessionServiceConfig -from trpc_agent_sdk.types import Content -from trpc_agent_sdk.types import Part - - -def _event(event_id: str, text: str) -> Event: - """Create a deterministic event for persistence tests.""" - return Event( - id=event_id, - invocation_id=f"invocation-{event_id}", - author="user", - content=Content(parts=[Part.from_text(text=text)]), - ) - - -def _config(root_dir: Path) -> AdvancedMemoryConfig: - """Disable model-driven background work for storage-only tests.""" - return AdvancedMemoryConfig( - root_dir=root_dir, - session_memory_enabled=False, - history_snip_enabled=False, - microcompact_enabled=False, - autocompact_enabled=False, - ) - - -async def test_session_service_persists_and_restores_events(tmp_path: Path) -> None: - """Ensure a new service instance can restore a complete transcript.""" - first = AdvancedMemorySessionService(config=_config(tmp_path)) - session = await first.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="demo-session", - state={ - "session-key": "session-value", - "app:theme": "dark", - "user:name": "alice", - }, - ) - await first.append_event(session, _event("event-1", "hello")) - metadata = json.loads((first.runtime.paths.session_dir(session.id) / "session.json").read_text(encoding="utf-8")) - assert metadata["state"] == {"session-key": "session-value"} - - second = AdvancedMemorySessionService(config=_config(tmp_path)) - restored = await second.get_session( - app_name="demo-app", - user_id="demo-user", - session_id="demo-session", - ) - - assert restored is not None - assert [event.id for event in restored.events] == ["event-1"] - assert restored.state["session-key"] == "session-value" - assert restored.state["app:theme"] == "dark" - assert restored.state["user:name"] == "alice" - - -async def test_session_id_collision_between_users_is_rejected(tmp_path: Path) -> None: - """Prevent different users from silently sharing one session directory.""" - service = AdvancedMemorySessionService(config=_config(tmp_path)) - await service.create_session( - app_name="demo-app", - user_id="user-a", - session_id="shared-session", - ) - - try: - await service.create_session( - app_name="demo-app", - user_id="user-b", - session_id="shared-session", - ) - except ValueError as exc: - assert "already used" in str(exc) - else: - raise AssertionError("Expected a cross-user session ID collision to fail") - - -async def test_delete_session_removes_persistent_session_data(tmp_path: Path) -> None: - """Ensure deleting a session removes its metadata and transcript directory.""" - service = AdvancedMemorySessionService(config=_config(tmp_path)) - session = await service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="delete-me", - state={ - "app:theme": "dark", - "user:name": "alice" - }, - ) - await service.append_event(session, _event("event-1", "hello")) - metadata = json.loads((service.runtime.paths.session_dir(session.id) / "session.json").read_text(encoding="utf-8")) - assert metadata["state"] == {} - - await service.delete_session( - app_name="demo-app", - user_id="demo-user", - session_id=session.id, - ) - - assert await service.get_session( - app_name="demo-app", - user_id="demo-user", - session_id=session.id, - ) is None - assert not service.runtime.paths.session_dir(session.id).exists() - - -async def test_ttl_cleanup_removes_expired_persistent_sessions(tmp_path: Path) -> None: - """Ensure configured session TTL removes idle session directories.""" - session_config = SessionServiceConfig(ttl=SessionServiceConfig.create_ttl_config( - ttl_seconds=1, - cleanup_interval_seconds=0.05, - )) - service = AdvancedMemorySessionService( - config=_config(tmp_path), - session_config=session_config, - ) - session = await service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="expires", - ) - - await asyncio.sleep(1.1) - - assert await service.get_session( - app_name="demo-app", - user_id="demo-user", - session_id=session.id, - ) is None - await service.close() - - -async def test_runner_binds_standalone_session_service(tmp_path: Path) -> None: - """Ensure Runner installs Advanced callbacks without a memory service.""" - service = AdvancedMemorySessionService(config=_config(tmp_path)) - agent = SimpleNamespace(name="test-agent", tools=[], before_model_callback=None) - - runner = Runner( - app_name="demo-app", - agent=agent, - session_service=service, - ) - - assert runner.session_service.delegate is service - assert service.integration is not None - - -def test_session_service_can_cross_event_loops(tmp_path: Path) -> None: - """Ensure deferred Runner work can share the service's file lock.""" - service = AdvancedMemorySessionService(config=_config(tmp_path)) - - async def create() -> None: - await service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="demo-session", - ) - - async def append() -> None: - session = await service.get_session( - app_name="demo-app", - user_id="demo-user", - session_id="demo-session", - ) - assert session is not None - await service.append_event(session, _event("event-1", "hello")) - - asyncio.run(create()) - asyncio.run(append()) - - -def test_transcript_decorator_can_cross_event_loops(tmp_path: Path) -> None: - """Ensure the wrapped service is safe for deferred-worker event loops.""" - service = AdvancedMemorySessionService(config=_config(tmp_path)) - runner = Runner( - app_name="demo-app", - agent=SimpleNamespace( - name="test-agent", - tools=[], - before_model_callback=None, - get_subagents=lambda: [], - ), - session_service=service, - ) - - async def create() -> None: - await runner.session_service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="wrapped-session", - ) - - async def append() -> None: - session = await runner.session_service.get_session( - app_name="demo-app", - user_id="demo-user", - session_id="wrapped-session", - ) - assert session is not None - await runner.session_service.append_event(session, _event("event-1", "hello")) - - asyncio.run(create()) - asyncio.run(append()) - records = asyncio.run(service.runtime.transcripts.read_all("wrapped-session")) - assert [record["event_id"] for record in records if record.get("kind") == "event"] == ["event-1"] - asyncio.run(runner.close()) diff --git a/tests/advanced_memory/test_advanced_memory_tools.py b/tests/advanced_memory/test_advanced_memory_tools.py index 12c83705b..59124377e 100644 --- a/tests/advanced_memory/test_advanced_memory_tools.py +++ b/tests/advanced_memory/test_advanced_memory_tools.py @@ -3,10 +3,13 @@ from __future__ import annotations from pathlib import Path +from types import SimpleNamespace +from unittest.mock import AsyncMock import pytest -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig +from trpc_agent_sdk.advanced_memory import AdvancedMemoryServiceConfig +from trpc_agent_sdk.advanced_memory import AdvancedMemoryPaths from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime from trpc_agent_sdk.tools import AdvancedMemoryTools from trpc_agent_sdk.tools import create_advanced_memory_tools @@ -14,10 +17,10 @@ def _runtime(tmp_path: Path) -> AdvancedMemoryRuntime: """Create a test runtime with long-term memory enabled.""" - return AdvancedMemoryRuntime.create(AdvancedMemoryConfig( + return AdvancedMemoryRuntime.create(AdvancedMemoryServiceConfig( enabled=True, root_dir=tmp_path, - )) + )).for_scope("demo-app", "demo-user") async def test_save_read_and_update_memory_index(tmp_path: Path) -> None: @@ -68,6 +71,34 @@ async def test_save_memory_rejects_unknown_type(tmp_path: Path) -> None: ) +@pytest.mark.parametrize( + ("storage_backend", "expected_prefix"), + (("redis", "advanced-memory://redis/"), ("sql", "advanced-memory://sql/")), +) +async def test_list_memory_index_reports_backend_storage_reference( + storage_backend: str, + expected_prefix: str, +) -> None: + """Avoid exposing a local filesystem path for external memory stores.""" + config = AdvancedMemoryServiceConfig( + storage_backend=storage_backend, + redis_url="redis://localhost:6379/0" if storage_backend == "redis" else None, + sql_url="sqlite:///advanced-memory.db" if storage_backend == "sql" else None, + ) + paths = AdvancedMemoryPaths(config).for_scope("demo-app", "demo-user") + runtime = SimpleNamespace( + config=config, + paths=paths, + scope=paths.scope, + long_term_memory=SimpleNamespace(read_index=AsyncMock(return_value="")), + ) + + result = await AdvancedMemoryTools(runtime).list_memory_index() + + assert result["index_path"].startswith(expected_prefix) + assert str(paths.memory_index_path) not in result["index_path"] + + def test_factory_returns_three_named_tools(tmp_path: Path) -> None: """Ensure the factory returns the three installable tools.""" tools = create_advanced_memory_tools(_runtime(tmp_path)) diff --git a/tests/advanced_memory/test_autocompact.py b/tests/advanced_memory/test_autocompact.py deleted file mode 100644 index af4e77a33..000000000 --- a/tests/advanced_memory/test_autocompact.py +++ /dev/null @@ -1,454 +0,0 @@ -"""Unit tests for automatic compaction, replay, and circuit breaking.""" - -from __future__ import annotations - -from pathlib import Path -from types import SimpleNamespace - -from trpc_agent_sdk.advanced_memory import AutoCompact -from trpc_agent_sdk.advanced_memory import AutoCompactCallback -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig -from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime -from trpc_agent_sdk.advanced_memory import HistorySnipCallback -from trpc_agent_sdk.advanced_memory import MicrocompactCallback -from trpc_agent_sdk.advanced_memory import SessionMemoryDocument -from trpc_agent_sdk.advanced_memory import setup_autocompact -from trpc_agent_sdk.advanced_memory import setup_history_snip -from trpc_agent_sdk.advanced_memory import setup_microcompact -from trpc_agent_sdk.advanced_memory import setup_tool_result_budget -from trpc_agent_sdk.advanced_memory import ToolResultBudgetCallback -from trpc_agent_sdk.events import Event -from trpc_agent_sdk.models import LlmRequest -from trpc_agent_sdk.sessions import InMemorySessionService -from trpc_agent_sdk.types import Content -from trpc_agent_sdk.types import Part - - -class FakeSummaryGenerator: - """Return a fixed summary or fail according to configuration.""" - - def __init__(self, *, fail: bool = False) -> None: - """Initialize call tracking and the failure switch.""" - self.fail = fail - self.histories: list[str] = [] - - async def generate(self, history: str, ctx) -> str: - """Record history and return a short summary.""" - self.histories.append(history) - if self.fail: - raise RuntimeError("summary failed") - return "## 压缩摘要\n\n保留用户目标、关键文件和当前状态。" - - -def _runtime( - tmp_path: Path, - *, - enabled: bool = True, - trigger: int = 4_000, - target: int = 3_000, - blocking: int = 5_000, - keep_recent: int = 2, - max_failures: int = 3, -) -> AdvancedMemoryRuntime: - """Create an isolated runtime with small automatic-compaction limits.""" - return AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( - enabled=enabled, - root_dir=tmp_path, - autocompact_trigger_chars=trigger, - autocompact_target_chars=target, - autocompact_blocking_chars=blocking, - autocompact_keep_recent_contents=keep_recent, - autocompact_max_failures=max_failures, - autocompact_summary_input_max_chars=10_000, - autocompact_summary_retries=2, - )) - - -def _request(count: int, *, text_size: int = 800) -> LlmRequest: - """Create a model request with multiple text Contents.""" - return LlmRequest( - model="test-model", - contents=[ - Content( - role="user" if index % 2 == 0 else "model", - parts=[Part.from_text(text=f"message-{index}-" + chr(97 + index) * text_size)], - ) for index in range(count) - ], - ) - - -def _ctx(session_id: str = "session-a"): - """Create the minimal context stand-in required by AutoCompact.""" - return SimpleNamespace( - session_id=session_id, - app_name="demo-app", - agent=SimpleNamespace(model="fake-model"), - ) - - -async def test_legacy_compact_replaces_old_prefix_and_keeps_recent(tmp_path: Path) -> None: - """Ensure missing session memory invokes the summary generator.""" - runtime = _runtime(tmp_path) - generator = FakeSummaryGenerator() - request = _request(5) - - result = await AutoCompact(runtime, generator).apply( - request, - session_id="session-a", - ctx=_ctx(), - force=True, - ) - - assert result.compacted is True - assert result.source == "legacy" - assert result.request_chars_after < result.request_chars_before - assert len(request.contents) == 3 - assert "This session is being continued" in request.contents[0].parts[0].text - assert "message-3-" in request.contents[1].parts[0].text - assert len(generator.histories) == 1 - - -async def test_token_budget_triggers_autocompact_and_records_diagnostics(tmp_path: Path) -> None: - """Ensure token thresholds replace character thresholds and persist diagnostics.""" - runtime = AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( - enabled=True, - root_dir=tmp_path, - autocompact_trigger_chars=100_000, - autocompact_target_chars=50_000, - autocompact_blocking_chars=120_000, - autocompact_keep_recent_contents=2, - autocompact_summary_input_max_chars=10_000, - model_context_window_tokens=1_100, - max_output_tokens=100, - )) - result = await AutoCompact(runtime, FakeSummaryGenerator()).apply( - _request(5), - session_id="session-a", - ctx=_ctx(), - ) - - assert result.compacted - assert result.request_tokens_before is not None - assert result.request_tokens_after is not None - assert result.request_tokens_after < result.request_tokens_before - records = await runtime.transcripts.read_all("session-a") - assert records[-1]["request_tokens_before"] == result.request_tokens_before - - -async def test_session_memory_compact_avoids_summary_model_call(tmp_path: Path) -> None: - """Ensure available session memory takes priority over legacy summaries.""" - runtime = _runtime( - tmp_path, - trigger=8_000, - target=7_000, - blocking=9_000, - ) - service = InMemorySessionService() - session = await service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="session-a", - ) - request = _request(5) - parent_event_id = None - for index, content in enumerate(request.contents): - event = Event( - id=f"event-{index}", - invocation_id="invocation-1", - author="agent", - content=content.model_copy(deep=True), - ) - await runtime.transcripts.append( - session.id, - { - "schema_version": 1, - "kind": "event", - "event_id": event.id, - "parent_event_id": parent_event_id, - "session": { - "id": session.id - }, - "event": event.model_dump(mode="json", by_alias=True, exclude_none=True), - }, - ) - parent_event_id = event.id - await runtime.session_memory.write( - session.id, - SessionMemoryDocument( - session_title="已有会话记忆", - current_state="正在继续实现自动压缩。", - ), - ) - await runtime.transcripts.append( - session.id, - { - "kind": "session-memory-checkpoint", - "checkpoint_id": "session-memory:event-2", - "first_event_id": "event-0", - "last_event_id": "event-2", - }, - ) - generator = FakeSummaryGenerator() - - result = await AutoCompact(runtime, generator).apply( - request, - session_id=session.id, - ctx=_ctx(session.id), - force=True, - ) - - assert result.compacted is True - assert result.source == "session-memory" - assert generator.histories == [] - assert "已有会话记忆" in request.contents[0].parts[0].text - assert "message-3-" in request.contents[1].parts[0].text - - -async def test_session_memory_compact_drops_all_contents_through_checkpoint(tmp_path: Path, ) -> None: - """Ensure session-memory compaction does not retain pre-checkpoint contents.""" - runtime = _runtime( - tmp_path, - trigger=8_000, - target=7_000, - blocking=9_000, - ) - service = InMemorySessionService() - session = await service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="session-a", - ) - request = _request(5) - parent_event_id = None - for index, content in enumerate(request.contents): - event = Event( - id=f"event-{index}", - invocation_id="invocation-1", - author="agent", - content=content.model_copy(deep=True), - ) - await runtime.transcripts.append( - session.id, - { - "schema_version": 1, - "kind": "event", - "event_id": event.id, - "parent_event_id": parent_event_id, - "session": { - "id": session.id - }, - "event": event.model_dump(mode="json", by_alias=True, exclude_none=True), - }, - ) - parent_event_id = event.id - await runtime.session_memory.write( - session.id, - SessionMemoryDocument( - session_title="已有会话记忆", - current_state="已总结到最后一个 event。", - ), - ) - await runtime.transcripts.append( - session.id, - { - "kind": "session-memory-checkpoint", - "checkpoint_id": "session-memory:event-4", - "first_event_id": "event-0", - "last_event_id": "event-4", - }, - ) - - result = await AutoCompact(runtime, FakeSummaryGenerator()).apply( - request, - session_id=session.id, - ctx=_ctx(session.id), - force=True, - ) - - assert result.compacted is True - assert result.source == "session-memory" - assert len(request.contents) == 1 - assert "已有会话记忆" in request.contents[0].parts[0].text - - -async def test_successful_compaction_is_reapplied_after_restart(tmp_path: Path) -> None: - """Ensure restart restores and replays the same compaction boundary.""" - first_runtime = _runtime(tmp_path, trigger=20_000, target=10_000, blocking=30_000) - first_request = _request(5) - first = await AutoCompact(first_runtime, FakeSummaryGenerator()).apply( - first_request, - session_id="session-a", - ctx=_ctx(), - force=True, - ) - first_payload = [content.model_dump(exclude_none=True) for content in first_request.contents] - - second_runtime = _runtime(tmp_path, trigger=20_000, target=10_000, blocking=30_000) - second_request = _request(5) - second = await AutoCompact(second_runtime, FakeSummaryGenerator()).apply( - second_request, - session_id="session-a", - ctx=_ctx(), - ) - - assert first.compacted is True - assert second.compacted is False - assert second.reapplied is True - assert [content.model_dump(exclude_none=True) for content in second_request.contents] == first_payload - - -async def test_reapplied_boundary_preserves_all_new_unsummarized_contents(tmp_path: Path) -> None: - """Ensure replay does not discard new history beyond recent contents.""" - runtime = _runtime(tmp_path, trigger=50_000, target=20_000, blocking=60_000) - initial_request = _request(5, text_size=300) - await AutoCompact(runtime, FakeSummaryGenerator()).apply( - initial_request, - session_id="session-a", - ctx=_ctx(), - force=True, - ) - expanded_request = _request(9, text_size=300) - - result = await AutoCompact(runtime, FakeSummaryGenerator()).apply( - expanded_request, - session_id="session-a", - ctx=_ctx(), - ) - - visible_text = "\n".join(part.text or "" for content in expanded_request.contents for part in content.parts or []) - assert result.reapplied is True - for index in range(3, 9): - assert f"message-{index}-" in visible_text - - -async def test_reapplied_boundary_uses_signature_occurrence_not_last_match(tmp_path: Path, ) -> None: - """Ensure duplicate Content signatures do not skip unsummarized messages.""" - runtime = _runtime(tmp_path) - duplicate = "duplicate-" + "d" * 500 - original_contents = [ - Content(role="user", parts=[Part.from_text(text=duplicate)]), - Content(role="model", parts=[Part.from_text(text="middle-" + "m" * 500)]), - Content(role="user", parts=[Part.from_text(text=duplicate)]), - Content(role="user", parts=[Part.from_text(text=duplicate)]), - Content(role="model", parts=[Part.from_text(text="last-" + "l" * 500)]), - ] - await AutoCompact(runtime, FakeSummaryGenerator()).apply( - LlmRequest(model="test-model", contents=original_contents), - session_id="session-a", - ctx=_ctx(), - force=True, - ) - second_request = LlmRequest( - model="test-model", - contents=[ - *[content.model_copy(deep=True) for content in original_contents], - Content(role="user", parts=[Part.from_text(text="new-message")]), - ], - ) - - result = await AutoCompact( - _runtime(tmp_path), - FakeSummaryGenerator(), - ).apply( - second_request, - session_id="session-a", - ctx=_ctx(), - ) - - assert result.reapplied is True - assert len(second_request.contents) == 4 - assert second_request.contents[1].parts[0].text == duplicate - - -async def test_failures_retry_internally_then_trip_circuit_breaker(tmp_path: Path) -> None: - """Ensure failures persist and trigger blocking at the hard limit.""" - runtime = _runtime( - tmp_path, - trigger=2_000, - target=1_000, - blocking=3_000, - max_failures=3, - ) - generator = FakeSummaryGenerator(fail=True) - compact = AutoCompact(runtime, generator) - results = [] - for _ in range(3): - results.append(await compact.apply( - _request(4, text_size=1_000), - session_id="session-a", - ctx=_ctx(), - force=True, - )) - - assert [result.consecutive_failures for result in results] == [1, 2, 3] - assert results[-1].blocked is True - assert len(generator.histories) == 6 - records = await runtime.transcripts.read_all("session-a") - assert len([record for record in records if record["kind"] == "autocompact-failure"]) == 3 - - -async def test_circuit_breaker_skips_further_summary_calls_below_hard_limit(tmp_path: Path) -> None: - """Ensure the circuit breaker avoids summary calls below the hard limit.""" - runtime = _runtime( - tmp_path, - trigger=2_000, - target=1_000, - blocking=10_000, - max_failures=1, - ) - generator = FakeSummaryGenerator(fail=True) - compact = AutoCompact(runtime, generator) - await compact.apply( - _request(4, text_size=800), - session_id="session-a", - ctx=_ctx(), - force=True, - ) - calls_after_failure = len(generator.histories) - - result = await compact.apply( - _request(4, text_size=800), - session_id="session-a", - ctx=_ctx(), - force=True, - ) - - assert result.blocked is False - assert result.consecutive_failures == 1 - assert len(generator.histories) == calls_after_failure - - -async def test_disabled_autocompact_does_not_copy_request(tmp_path: Path) -> None: - """Ensure disabled mode preserves requests and disk state.""" - runtime = _runtime(tmp_path, enabled=False) - request = _request(5) - original_content = request.contents[0] - - result = await AutoCompact(runtime, FakeSummaryGenerator()).apply( - request, - session_id="session-a", - ctx=_ctx(), - force=True, - ) - - assert result.compacted is False - assert request.contents[0] is original_content - assert not (tmp_path / "SESSION").exists() - - -def test_setup_orders_full_context_pipeline(tmp_path: Path) -> None: - """Ensure any setup order yields the expected callback pipeline.""" - runtime = _runtime(tmp_path) - agent = SimpleNamespace(before_model_callback=None) - - setup_autocompact(agent, runtime, FakeSummaryGenerator()) - setup_microcompact(agent, runtime) - setup_history_snip(agent, runtime) - setup_tool_result_budget(agent, runtime) - - assert isinstance(agent.before_model_callback[0], ToolResultBudgetCallback) - assert isinstance(agent.before_model_callback[1], HistorySnipCallback) - assert isinstance(agent.before_model_callback[2], MicrocompactCallback) - assert isinstance(agent.before_model_callback[3], AutoCompactCallback) diff --git a/tests/advanced_memory/test_history_snip.py b/tests/advanced_memory/test_history_snip.py deleted file mode 100644 index 9d13a966c..000000000 --- a/tests/advanced_memory/test_history_snip.py +++ /dev/null @@ -1,225 +0,0 @@ -"""Unit tests for history snip under context pressure.""" - -from __future__ import annotations - -from pathlib import Path -from types import SimpleNamespace - -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig -from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime -from trpc_agent_sdk.advanced_memory import HistorySnip -from trpc_agent_sdk.advanced_memory import HistorySnipCallback -from trpc_agent_sdk.advanced_memory import Microcompact -from trpc_agent_sdk.advanced_memory import MicrocompactCallback -from trpc_agent_sdk.advanced_memory import setup_history_snip -from trpc_agent_sdk.advanced_memory import setup_microcompact -from trpc_agent_sdk.advanced_memory import setup_tool_result_budget -from trpc_agent_sdk.advanced_memory import ToolResultBudgetCallback -from trpc_agent_sdk.advanced_memory import ToolResultBudget -from trpc_agent_sdk.models import LlmRequest -from trpc_agent_sdk.types import Content -from trpc_agent_sdk.types import FunctionResponse -from trpc_agent_sdk.types import Part - - -def _runtime( - tmp_path: Path, - *, - enabled: bool = True, - snip_enabled: bool = True, - trigger_chars: int = 1_000, - target_chars: int = 600, - keep_recent: int = 2, -) -> AdvancedMemoryRuntime: - """Create an isolated runtime with small history-snip limits.""" - return AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( - enabled=enabled, - root_dir=tmp_path, - tool_result_max_chars=5_000, - tool_result_preview_chars=100, - history_snip_enabled=snip_enabled, - history_snip_trigger_chars=trigger_chars, - history_snip_target_chars=target_chars, - history_snip_keep_recent=keep_recent, - )) - - -def _request(count: int, *, output_size: int = 400) -> tuple[LlmRequest, list[Part]]: - """Create a model request with sized tool results.""" - parts = [ - Part(function_response=FunctionResponse( - id=f"result-{index}", - name="Read", - response={"output": chr(97 + index) * output_size}, - )) for index in range(count) - ] - return LlmRequest(model="test-model", contents=[Content(role="user", parts=parts)]), parts - - -def _outputs(request: LlmRequest) -> list[str]: - """Extract all tool outputs from a model request.""" - return [part.function_response.response["output"] for part in request.contents[0].parts] - - -async def test_pressure_snips_old_results_and_keeps_recent(tmp_path: Path) -> None: - """Ensure oversized requests clean old results and keep recent work.""" - request, original_parts = _request(4) - original_response = original_parts[0].function_response.response.copy() - - result = await HistorySnip(_runtime(tmp_path)).apply( - request, - session_id="session-a", - ) - - outputs = _outputs(request) - assert result.trigger == "pressure" - assert result.snipped_count == 2 - assert result.request_chars_after < result.request_chars_before - assert outputs[:2] == ["[Older tool result removed by history snip]"] * 2 - assert outputs[2:] == ["c" * 400, "d" * 400] - assert original_parts[0].function_response.response == original_response - - -async def test_token_budget_triggers_snip_without_character_pressure(tmp_path: Path) -> None: - """Ensure a configured model window triggers cleanup by token warning.""" - request, _ = _request(4, output_size=1_000) - runtime = AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( - enabled=True, - root_dir=tmp_path, - tool_result_max_chars=10_000, - tool_result_preview_chars=100, - history_snip_trigger_chars=100_000, - history_snip_target_chars=50_000, - history_snip_keep_recent=2, - model_context_window_tokens=1_000, - max_output_tokens=100, - )) - - result = await HistorySnip(runtime).apply(request, session_id="session-a") - - assert result.trigger == "pressure" - assert result.snipped_count == 2 - assert result.request_tokens_before is not None - assert result.request_tokens_after is not None - assert result.request_tokens_after < result.request_tokens_before - - -async def test_request_below_trigger_remains_unchanged(tmp_path: Path) -> None: - """Ensure cleanup does not run below the configured threshold.""" - request, _ = _request(2, output_size=50) - - result = await HistorySnip(_runtime(tmp_path)).apply( - request, - session_id="session-a", - ) - - assert result.trigger is None - assert result.snipped_count == 0 - assert _outputs(request) == ["a" * 50, "b" * 50] - - -async def test_force_snip_runs_below_pressure_threshold(tmp_path: Path) -> None: - """Ensure force mode cleans results before the recent working set.""" - request, _ = _request(4, output_size=100) - - result = await HistorySnip(_runtime(tmp_path, trigger_chars=10_000, target_chars=5_000)).apply( - request, - session_id="session-a", - force=True, - ) - - assert result.trigger == "force" - assert result.snipped_count == 2 - assert _outputs(request)[:2] == ["[Older tool result removed by history snip]"] * 2 - - -async def test_snipped_results_are_reapplied_after_restart(tmp_path: Path) -> None: - """Ensure history-snip decisions can be restored from the transcript.""" - first_request, _ = _request(4) - await HistorySnip(_runtime(tmp_path)).apply( - first_request, - session_id="session-a", - ) - - second_request, _ = _request(2) - result = await HistorySnip(_runtime(tmp_path)).apply( - second_request, - session_id="session-a", - ) - records = await _runtime(tmp_path).transcripts.read_all("session-a") - - assert result.trigger is None - assert result.reapplied_count == 2 - assert _outputs(second_request) == ["[Older tool result removed by history snip]"] * 2 - assert len([record for record in records if record["kind"] == "history-snip"]) == 2 - - -async def test_budget_recovery_pointer_survives_later_shrink_stages(tmp_path: Path, ) -> None: - """Ensure snip and Microcompact preserve budget-generated result paths.""" - runtime = AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( - enabled=True, - root_dir=tmp_path, - tool_result_max_chars=200, - tool_results_per_message_max_chars=10_000, - tool_result_preview_chars=40, - history_snip_trigger_chars=1_000, - history_snip_target_chars=500, - history_snip_keep_recent=1, - microcompact_trigger_count=2, - microcompact_keep_recent=1, - )) - request, _ = _request(5, output_size=100) - request.contents[0].parts[0].function_response.response = {"output": "oversized" * 100} - - await ToolResultBudget(runtime).apply(request, session_id="session-a") - recovery_response = request.contents[0].parts[0].function_response.response - recovery_path = recovery_response["persisted_output"]["path"] - await HistorySnip(runtime).apply( - request, - session_id="session-a", - force=True, - ) - await Microcompact(runtime).apply( - request, - session_id="session-a", - last_assistant_timestamp=None, - ) - - final_response = request.contents[0].parts[0].function_response.response - assert final_response["persisted_output"]["path"] == recovery_path - records = await runtime.transcripts.read_all("session-a") - assert not any( - record.get("result_id") == "result-0" and record.get("kind") in {"history-snip", "microcompact-clear"} - for record in records) - - -async def test_disabled_history_snip_does_not_copy_or_persist(tmp_path: Path) -> None: - """Ensure disabled history snip does not copy requests or create storage.""" - request, _ = _request(4) - original_content = request.contents[0] - - result = await HistorySnip(_runtime(tmp_path, snip_enabled=False)).apply( - request, - session_id="session-a", - ) - - assert result.snipped_count == 0 - assert request.contents[0] is original_content - assert not (tmp_path / "SESSION").exists() - - -def test_setup_orders_all_context_callbacks_by_stage(tmp_path: Path) -> None: - """Ensure any installation order yields the fixed callback order.""" - runtime = _runtime(tmp_path) - agent = SimpleNamespace(before_model_callback=None) - - setup_microcompact(agent, runtime) - setup_history_snip(agent, runtime) - setup_tool_result_budget(agent, runtime) - - assert isinstance(agent.before_model_callback[0], ToolResultBudgetCallback) - assert isinstance(agent.before_model_callback[1], HistorySnipCallback) - assert isinstance(agent.before_model_callback[2], MicrocompactCallback) diff --git a/tests/advanced_memory/test_memory_context.py b/tests/advanced_memory/test_memory_context.py index 965be4f65..a2385d9d3 100644 --- a/tests/advanced_memory/test_memory_context.py +++ b/tests/advanced_memory/test_memory_context.py @@ -7,35 +7,20 @@ import pytest -from trpc_agent_sdk.advanced_memory import AutoCompactCallback -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig +from trpc_agent_sdk.advanced_memory import AdvancedMemoryServiceConfig from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime -from trpc_agent_sdk.advanced_memory import HistorySnipCallback from trpc_agent_sdk.advanced_memory import LongTermMemoryContext from trpc_agent_sdk.advanced_memory import LongTermMemoryContextCallback from trpc_agent_sdk.advanced_memory import MemoryIndexEntry -from trpc_agent_sdk.advanced_memory import MicrocompactCallback -from trpc_agent_sdk.advanced_memory import setup_advanced_memory -from trpc_agent_sdk.advanced_memory import setup_context_management -from trpc_agent_sdk.advanced_memory import ToolResultBudgetCallback -from trpc_agent_sdk.advanced_memory import TranscriptSessionService -from trpc_agent_sdk.advanced_memory._callbacks import install_staged_callback +from trpc_agent_sdk.memory import AdvancedMemoryService from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.sessions.compact._callbacks import install_staged_callback from trpc_agent_sdk.sessions import InMemorySessionService -class FakeSummaryGenerator: - """Provide a summary generator that does not call a real model.""" - - async def generate(self, history: str, ctx) -> str: - """Return a fixed test summary.""" - del history, ctx - return "summary" - - def _runtime(tmp_path: Path) -> AdvancedMemoryRuntime: """Create a test runtime with long-term memory injection enabled.""" - return AdvancedMemoryRuntime.create(AdvancedMemoryConfig( + return AdvancedMemoryRuntime.create(AdvancedMemoryServiceConfig( enabled=True, root_dir=tmp_path, )) @@ -104,61 +89,52 @@ async def test_long_term_memory_index_is_injected_once(tmp_path: Path) -> None: assert "secrets, credentials, tokens, and other sensitive data" in instruction -async def test_unified_setup_installs_complete_pipeline_in_order(tmp_path: Path) -> None: - """Ensure unified setup installs the five components in order.""" - runtime = _runtime(tmp_path) - agent = SimpleNamespace(before_model_callback=None) +async def test_custom_memory_focus_is_injected_into_system_instruction(tmp_path: Path) -> None: + """Ensure applications can prioritize a custom long-term memory focus.""" + runtime = AdvancedMemoryRuntime.create( + AdvancedMemoryServiceConfig( + enabled=True, + root_dir=tmp_path, + memory_focus_instruction="重点记住用户长期稳定的兴趣爱好。", + )) + request = LlmRequest(model="test-model") - components = setup_context_management( - agent, - runtime, - FakeSummaryGenerator(), - ) + applied = await LongTermMemoryContext(runtime).apply(request) - assert components.long_term_memory.runtime is runtime - assert isinstance(agent.before_model_callback[0], LongTermMemoryContextCallback) - assert isinstance(agent.before_model_callback[1], ToolResultBudgetCallback) - assert isinstance(agent.before_model_callback[2], HistorySnipCallback) - assert isinstance(agent.before_model_callback[3], MicrocompactCallback) - assert isinstance(agent.before_model_callback[4], AutoCompactCallback) + instruction = str(request.config.system_instruction) + assert applied is True + assert "## Custom memory focus" in instruction + assert "重点记住用户长期稳定的兴趣爱好。" in instruction -async def test_full_setup_wraps_session_service_and_is_idempotent(tmp_path: Path, ) -> None: - """Ensure unified setup assembles transcript, session memory, and callbacks.""" +async def test_memory_service_does_not_install_session_compression(tmp_path: Path, ) -> None: + """Ensure the MemoryService leaves the supplied SessionService unchanged.""" runtime = _runtime(tmp_path) + memory_service = AdvancedMemoryService(runtime=runtime) + session_service = InMemorySessionService() agent = SimpleNamespace(before_model_callback=None, tools=[]) - first = setup_advanced_memory( - agent, - InMemorySessionService(), - runtime, - FakeSummaryGenerator(), - ) - second = setup_advanced_memory( - agent, - first.session_service, - runtime, - FakeSummaryGenerator(), - ) + bound = memory_service.bind(agent, session_service) - assert isinstance(first.session_service, TranscriptSessionService) - assert first.session_memory_extractor.runtime is runtime - assert first.session_service.session_memory_extractor is first.session_memory_extractor - assert second.session_service is first.session_service - assert second.session_memory_extractor is first.session_memory_extractor - assert second.long_term_memory_tools is first.long_term_memory_tools - assert len(agent.before_model_callback) == 5 + assert bound is session_service + assert len(agent.before_model_callback) == 1 + assert isinstance( + agent.before_model_callback[0], + LongTermMemoryContextCallback, + ) tool_names = {tool.name for tool in agent.tools} assert tool_names == { "save_memory", "read_memory", "list_memory_index", } + await session_service.close() + await memory_service.close() async def test_disabled_runtime_does_not_modify_system_instruction(tmp_path: Path) -> None: """Ensure disabled runtime does not inject long-term memory.""" - runtime = AdvancedMemoryRuntime.create(AdvancedMemoryConfig(enabled=False, root_dir=tmp_path)) + runtime = AdvancedMemoryRuntime.create(AdvancedMemoryServiceConfig(enabled=False, root_dir=tmp_path)) request = LlmRequest(model="test-model") applied = await LongTermMemoryContext(runtime).apply(request) diff --git a/tests/advanced_memory/test_microcompact.py b/tests/advanced_memory/test_microcompact.py deleted file mode 100644 index 91887898e..000000000 --- a/tests/advanced_memory/test_microcompact.py +++ /dev/null @@ -1,178 +0,0 @@ -"""Unit tests for mechanically cleaning old tool results.""" - -from __future__ import annotations - -from pathlib import Path -from types import SimpleNamespace - -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig -from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime -from trpc_agent_sdk.advanced_memory import Microcompact -from trpc_agent_sdk.advanced_memory import MicrocompactCallback -from trpc_agent_sdk.advanced_memory import setup_microcompact -from trpc_agent_sdk.advanced_memory import setup_tool_result_budget -from trpc_agent_sdk.advanced_memory import ToolResultBudgetCallback -from trpc_agent_sdk.models import LlmRequest -from trpc_agent_sdk.types import Content -from trpc_agent_sdk.types import FunctionResponse -from trpc_agent_sdk.types import Part - - -def _runtime( - tmp_path: Path, - *, - enabled: bool = True, - microcompact_enabled: bool = True, - trigger_count: int = 4, - keep_recent: int = 2, - gap_seconds: float = 60.0, -) -> AdvancedMemoryRuntime: - """Create an isolated runtime with small mechanical-compaction limits.""" - return AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( - enabled=enabled, - root_dir=tmp_path, - tool_result_max_chars=1_000, - tool_result_preview_chars=100, - microcompact_enabled=microcompact_enabled, - microcompact_trigger_count=trigger_count, - microcompact_keep_recent=keep_recent, - microcompact_gap_seconds=gap_seconds, - )) - - -def _request(count: int, *, tool_name: str = "Read") -> tuple[LlmRequest, list[Part]]: - """Create a model request with a specified number of tool results.""" - parts = [ - Part(function_response=FunctionResponse( - id=f"result-{index}", - name=tool_name, - response={"output": chr(97 + index) * 200}, - )) for index in range(count) - ] - return LlmRequest(model="test-model", contents=[Content(role="user", parts=parts)]), parts - - -def _outputs(request: LlmRequest) -> list[str]: - """Extract the output text for each tool result.""" - return [part.function_response.response["output"] for part in request.contents[0].parts] - - -async def test_count_trigger_clears_old_results_and_keeps_recent(tmp_path: Path) -> None: - """Ensure count pressure cleans only old results.""" - request, original_parts = _request(5) - original_first_response = original_parts[0].function_response.response.copy() - - result = await Microcompact(_runtime(tmp_path)).apply( - request, - session_id="session-a", - last_assistant_timestamp=None, - ) - - outputs = _outputs(request) - assert result.trigger == "count" - assert result.cleared_count == 3 - assert outputs[:3] == ["[Old tool result content cleared]"] * 3 - assert outputs[3:] == ["d" * 200, "e" * 200] - assert original_parts[0].function_response.response == original_first_response - - -async def test_time_trigger_runs_below_count_threshold(tmp_path: Path) -> None: - """Ensure a long time gap cleans old results before count pressure.""" - request, _ = _request(4) - - result = await Microcompact(_runtime(tmp_path)).apply( - request, - session_id="session-a", - last_assistant_timestamp=100.0, - now=161.0, - ) - - assert result.trigger == "time" - assert result.cleared_count == 2 - assert _outputs(request)[:2] == ["[Old tool result content cleared]"] * 2 - - -async def test_time_trigger_does_not_clear_when_only_recent_results_exist(tmp_path: Path) -> None: - """Ensure the configured recent results remain after time pressure.""" - request, _ = _request(2) - - result = await Microcompact(_runtime(tmp_path)).apply( - request, - session_id="session-a", - last_assistant_timestamp=100.0, - now=161.0, - ) - - assert result.trigger is None - assert result.cleared_count == 0 - assert _outputs(request) == ["a" * 200, "b" * 200] - - -async def test_cleared_results_are_reapplied_after_restart(tmp_path: Path) -> None: - """Ensure restart restores and reapplies the same cleanup.""" - first_request, _ = _request(5) - await Microcompact(_runtime(tmp_path)).apply( - first_request, - session_id="session-a", - last_assistant_timestamp=None, - ) - - second_request, _ = _request(3) - result = await Microcompact(_runtime(tmp_path)).apply( - second_request, - session_id="session-a", - last_assistant_timestamp=None, - ) - records = await _runtime(tmp_path).transcripts.read_all("session-a") - - assert result.trigger is None - assert result.reapplied_count == 3 - assert _outputs(second_request) == ["[Old tool result content cleared]"] * 3 - assert len([record for record in records if record["kind"] == "microcompact-clear"]) == 3 - - -async def test_non_compactable_tools_are_ignored(tmp_path: Path) -> None: - """Ensure unconfigured tools do not affect thresholds or cleanup.""" - request, _ = _request(6, tool_name="CustomTool") - - result = await Microcompact(_runtime(tmp_path)).apply( - request, - session_id="session-a", - last_assistant_timestamp=None, - ) - - assert result.cleared_count == 0 - assert _outputs(request)[0] == "a" * 200 - - -async def test_disabled_microcompact_does_not_copy_or_persist(tmp_path: Path) -> None: - """Ensure disabled compaction preserves requests and disk state.""" - request, _ = _request(5) - original_content = request.contents[0] - - result = await Microcompact(_runtime(tmp_path, microcompact_enabled=False)).apply( - request, - session_id="session-a", - last_assistant_timestamp=None, - ) - - assert result.cleared_count == 0 - assert request.contents[0] is original_content - assert not (tmp_path / "SESSION").exists() - - -def test_setup_orders_budget_before_microcompact_in_both_call_orders(tmp_path: Path) -> None: - """Ensure both setup functions keep budgeting before mechanical cleanup.""" - runtime = _runtime(tmp_path) - first_agent = SimpleNamespace(before_model_callback=None) - setup_microcompact(first_agent, runtime) - setup_tool_result_budget(first_agent, runtime) - - second_agent = SimpleNamespace(before_model_callback=None) - setup_tool_result_budget(second_agent, runtime) - setup_microcompact(second_agent, runtime) - - for agent in (first_agent, second_agent): - assert isinstance(agent.before_model_callback[0], ToolResultBudgetCallback) - assert isinstance(agent.before_model_callback[1], MicrocompactCallback) diff --git a/tests/advanced_memory/test_preload_memory.py b/tests/advanced_memory/test_preload_memory.py index 8a854da8f..a2c61deb8 100644 --- a/tests/advanced_memory/test_preload_memory.py +++ b/tests/advanced_memory/test_preload_memory.py @@ -5,7 +5,7 @@ from pathlib import Path from types import SimpleNamespace -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig +from trpc_agent_sdk.advanced_memory import AdvancedMemoryServiceConfig from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime from trpc_agent_sdk.advanced_memory import MemoryDocument from trpc_agent_sdk.advanced_memory import MemoryPreloader @@ -32,14 +32,14 @@ async def select(self, query, candidates, ctx, *, limit): async def test_preloader_injects_selected_topic_with_budget(tmp_path: Path) -> None: """Ensure selected topic content is rendered and bounded.""" runtime = AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( + AdvancedMemoryServiceConfig( enabled=True, root_dir=tmp_path, preload_memory_enabled=True, preload_memory_max_chars=200, )) await runtime.initialize() - await runtime.long_term_memory.write_topic( + await runtime.for_scope("demo-app", "demo-user").long_term_memory.write_topic( "project.md", MemoryDocument( name="Project", @@ -48,10 +48,15 @@ async def test_preloader_injects_selected_topic_with_budget(tmp_path: Path) -> N content="important project details", ), ) + ctx = SimpleNamespace(session=SimpleNamespace( + app_name="demo-app", + user_id="demo-user", + id="session-a", + )) result = await MemoryPreloader(runtime, _FakeSelector()).preload( "What is relevant?", - SimpleNamespace(), + ctx, ) assert result is not None @@ -63,14 +68,14 @@ async def test_preloader_injects_selected_topic_with_budget(tmp_path: Path) -> N async def test_preloader_marks_truncated_content(tmp_path: Path) -> None: """Tell the main model when the configured content budget truncated a topic.""" runtime = AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( + AdvancedMemoryServiceConfig( enabled=True, root_dir=tmp_path, preload_memory_enabled=True, preload_memory_max_chars=12, )) await runtime.initialize() - await runtime.long_term_memory.write_topic( + await runtime.for_scope("demo-app", "demo-user").long_term_memory.write_topic( "project.md", MemoryDocument( name="Project", @@ -79,10 +84,15 @@ async def test_preloader_marks_truncated_content(tmp_path: Path) -> None: content="important project details", ), ) + ctx = SimpleNamespace(session=SimpleNamespace( + app_name="demo-app", + user_id="demo-user", + id="session-a", + )) result = await MemoryPreloader(runtime, _FakeSelector()).preload( "What is relevant?", - SimpleNamespace(), + ctx, ) assert result is not None @@ -93,13 +103,13 @@ async def test_preloader_marks_truncated_content(tmp_path: Path) -> None: async def test_preloader_failure_is_best_effort(tmp_path: Path) -> None: """Return no prompt content when relevance screening fails.""" runtime = AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( + AdvancedMemoryServiceConfig( enabled=True, root_dir=tmp_path, preload_memory_enabled=True, )) await runtime.initialize() - await runtime.long_term_memory.write_topic( + await runtime.for_scope("demo-app", "demo-user").long_term_memory.write_topic( "project.md", MemoryDocument( name="Project", diff --git a/tests/advanced_memory/test_session_memory_extractor.py b/tests/advanced_memory/test_session_memory_extractor.py deleted file mode 100644 index 240e1f2fc..000000000 --- a/tests/advanced_memory/test_session_memory_extractor.py +++ /dev/null @@ -1,516 +0,0 @@ -"""Unit tests for full-context session-memory extraction and isolation.""" - -from __future__ import annotations - -from pathlib import Path -from types import SimpleNamespace - -import pytest - -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig -from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime -from trpc_agent_sdk.advanced_memory import ForkedSessionMemoryGenerator -from trpc_agent_sdk.advanced_memory import SessionMemoryExtractionInput -from trpc_agent_sdk.advanced_memory import SessionMemoryDocument -from trpc_agent_sdk.advanced_memory import SessionMemoryExtractor -from trpc_agent_sdk.advanced_memory import TranscriptSessionService -from trpc_agent_sdk.events import Event -from trpc_agent_sdk.models import LLMModel -from trpc_agent_sdk.models import LlmResponse -from trpc_agent_sdk.sessions import InMemorySessionService -from trpc_agent_sdk.types import Content -from trpc_agent_sdk.types import FunctionCall -from trpc_agent_sdk.types import Part - - -class FakeSessionMemoryGenerator: - """Record extraction input and return a deterministic document.""" - - def __init__(self, *, fail: bool = False) -> None: - """Initialize call tracking and the optional failure switch.""" - self.inputs = [] - self.fail = fail - - async def generate(self, extraction_input, ctx) -> SessionMemoryDocument: - """Generate a fixed test document from the last Event ID.""" - self.inputs.append(extraction_input) - if self.fail: - raise RuntimeError("generator failed") - return SessionMemoryDocument( - session_title="增量抽取测试", - current_state=f"已处理到 {extraction_input.last_event_id}", - task_specification="验证 session memory 增量更新。", - worklog=f"- {extraction_input.first_event_id} -> {extraction_input.last_event_id}", - ) - - -class EmptySessionMemoryGenerator: - """Simulate an invalid extractor returning ten empty sections.""" - - async def generate(self, extraction_input, ctx) -> SessionMemoryDocument: - """Ignore input and return a complete empty template.""" - del extraction_input, ctx - return SessionMemoryDocument() - - -class StructuredMemoryModel(LLMModel): - """Return an isolated Runner model with fixed Markdown memory.""" - - def __init__(self, *, empty: bool = False) -> None: - """Initialize the test model and store received requests.""" - super().__init__(model_name="session-memory-test-model") - self.requests = [] - self.empty = empty - - @classmethod - def supported_models(cls): - """Declare the names supported by the test model.""" - return [r"session-memory-test-model"] - - async def _generate_async_impl(self, request, stream=False, ctx=None): - """Record a request and return parser-compatible Markdown.""" - self.requests.append(request) - payload = ("# Session Title\n\n" if self.empty else SessionMemoryDocument( - session_title="隔离 Runner", - current_state="子 Agent 已完成。", - ).to_markdown()) - yield LlmResponse(content=Content( - role="model", - parts=[Part.from_text(text=payload)], - )) - - def validate_request(self, request): - """Allow all model requests in tests.""" - return None - - -def _runtime( - tmp_path: Path, - *, - initial_chars: int = 1, - update_chars: int = 1, - prompt_max_chars: int = 10_000, - section_max_chars: int = 8_000, -) -> AdvancedMemoryRuntime: - """Create an isolated runtime with small extraction limits.""" - return AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( - enabled=True, - root_dir=tmp_path, - session_memory_initial_chars=initial_chars, - session_memory_update_chars=update_chars, - session_memory_prompt_max_chars=prompt_max_chars, - session_memory_section_max_chars=section_max_chars, - )) - - -def _event(event_id: str, text: str) -> Event: - """Create a non-streaming Event for a transcript.""" - return Event( - id=event_id, - invocation_id="invocation-1", - author="agent", - content=Content(parts=[Part.from_text(text=text)]), - ) - - -async def _service_and_session(runtime: AdvancedMemoryRuntime): - """Create a test SessionService and session with automatic transcripts.""" - service = TranscriptSessionService(InMemorySessionService(), runtime) - session = await service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="demo-session", - ) - return service, session - - -def _ctx(session): - """Create the minimal InvocationContext stand-in for generator tests.""" - return SimpleNamespace(session=session, agent=SimpleNamespace(model="fake-model")) - - -async def test_first_extraction_writes_document_and_checkpoint(tmp_path: Path) -> None: - """Ensure the first threshold hit generates a document and records a boundary.""" - runtime = _runtime(tmp_path) - service, session = await _service_and_session(runtime) - await service.append_event(session, _event("event-1", "分析项目结构")) - await service.append_event(session, _event("event-2", "完成第一阶段")) - generator = FakeSessionMemoryGenerator() - - result = await SessionMemoryExtractor(runtime, generator).extract_if_needed( - session, - _ctx(session), - ) - - memory = await runtime.session_memory.read(session.id) - records = await runtime.transcripts.read_all(session.id) - checkpoints = [record for record in records if record["kind"] == "session-memory-checkpoint"] - assert result.extracted is True - assert result.processed_events == 2 - assert "# Session Title\n_A short and distinctive" in memory - assert "\n\n增量抽取测试" in memory - assert "# Learnings\n_What has worked well?" in memory - assert checkpoints[-1]["last_event_id"] == "event-2" - - -async def test_token_threshold_triggers_extraction_before_character_threshold(tmp_path: Path) -> None: - """Ensure session memory uses token thresholds when configured.""" - runtime = AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( - enabled=True, - root_dir=tmp_path, - session_memory_initial_chars=100_000, - session_memory_update_chars=100_000, - session_memory_initial_tokens=10, - session_memory_update_tokens=10, - session_memory_tool_calls_between_updates=1, - model_context_window_tokens=1_000, - max_output_tokens=100, - session_memory_request_overhead_tokens=50, - )) - service, session = await _service_and_session(runtime) - await service.append_event(session, _event("event-1", "x" * 200)) - - result = await SessionMemoryExtractor( - runtime, - FakeSessionMemoryGenerator(), - ).extract_if_needed(session, _ctx(session)) - - assert result.extracted is True - - -async def test_next_extraction_uses_full_context_after_checkpoint(tmp_path: Path) -> None: - """Ensure each update receives the full visible context and a checkpoint delta.""" - runtime = _runtime(tmp_path) - service, session = await _service_and_session(runtime) - generator = FakeSessionMemoryGenerator() - extractor = SessionMemoryExtractor(runtime, generator) - await service.append_event(session, _event("event-1", "first")) - await extractor.extract_if_needed(session, _ctx(session)) - await service.append_event(session, _event("event-2", "second")) - - result = await extractor.extract_if_needed(session, _ctx(session)) - - assert result.extracted is True - assert result.processed_events == 1 - assert generator.inputs[-1].first_event_id == "event-2" - assert "已处理到 event-1" in generator.inputs[-1].current_memory - assert "first" in generator.inputs[-1].context_messages - assert "second" in generator.inputs[-1].context_messages - assert generator.inputs[-1].new_events == "" - - -async def test_context_messages_keep_latest_content_and_remove_metadata(tmp_path: Path, ) -> None: - """Ensure full visible Content excludes thoughts and Event metadata.""" - runtime = _runtime(tmp_path, prompt_max_chars=5_000) - service, session = await _service_and_session(runtime) - await service.append_event(session, _event("event-old", "old-" + "x" * 300)) - latest = Event( - id="event-latest", - invocation_id="invocation-secret", - author="agent-secret", - content=Content( - role="model", - parts=[ - Part(text="hidden reasoning", thought=True), - Part.from_text(text="latest visible answer"), - ], - ), - ) - await service.append_event(session, latest) - generator = FakeSessionMemoryGenerator() - - result = await SessionMemoryExtractor(runtime, generator).extract_if_needed( - session, - _ctx(session), - force=True, - ) - - context_messages = generator.inputs[0].context_messages - all_context = context_messages + generator.inputs[0].new_events - assert result.extracted is True - assert "latest visible answer" in context_messages - assert "old-" in context_messages - assert "hidden reasoning" not in all_context - assert "event-latest" not in all_context - assert "invocation-secret" not in all_context - assert "agent-secret" not in all_context - - -async def test_full_context_is_sent_without_checkpoint_duplication(tmp_path: Path, ) -> None: - """Ensure the session memory Agent receives the complete visible context.""" - runtime = _runtime(tmp_path, prompt_max_chars=20_000) - service, session = await _service_and_session(runtime) - await service.append_event(session, _event("event-1", "old-context-" + "x" * 1_000)) - await service.append_event(session, _event("event-2", "recent-context-" + "y" * 1_000)) - await service.append_event(session, _event("event-3", "latest-context-" + "z" * 1_000)) - generator = FakeSessionMemoryGenerator() - - result = await SessionMemoryExtractor(runtime, generator).extract_if_needed( - session, - _ctx(session), - force=True, - ) - - extraction_input = generator.inputs[0] - assert result.extracted is True - assert result.processed_events == 3 - assert "latest-context-" in extraction_input.context_messages - assert "old-context-" in extraction_input.context_messages - assert "recent-context-" in extraction_input.context_messages - assert extraction_input.new_events == "" - - -async def test_full_context_over_budget_does_not_advance_checkpoint(tmp_path: Path, ) -> None: - """Process the largest safe event prefix instead of stalling forever.""" - runtime = _runtime(tmp_path, prompt_max_chars=3_000) - service, session = await _service_and_session(runtime) - for index in range(3): - await service.append_event(session, _event(f"event-{index}", "x" * 1_000)) - result = await SessionMemoryExtractor(runtime, FakeSessionMemoryGenerator()).extract_if_needed( - session, - _ctx(session), - force=True, - ) - - assert result.extracted is True - assert result.processed_events == 3 - - -async def test_compacted_context_can_still_process_transcript_delta(tmp_path: Path, ) -> None: - """Ensure events omitted by compaction are supplied from the transcript delta.""" - runtime = _runtime(tmp_path, prompt_max_chars=10_000) - service, session = await _service_and_session(runtime) - generator = FakeSessionMemoryGenerator() - extractor = SessionMemoryExtractor(runtime, generator) - - await service.append_event(session, _event("event-1", "old context")) - await extractor.extract_if_needed(session, _ctx(session)) - await service.append_event(session, _event("event-2", "new context")) - - compacted_ctx = SimpleNamespace( - session=session, - agent=SimpleNamespace(model="fake-model"), - override_messages=[ - Content(parts=[Part.from_text(text="compact summary")]), - ], - ) - result = await extractor.extract_if_needed(session, compacted_ctx) - - assert result.extracted is True - assert result.last_event_id == "event-2" - assert "new context" in generator.inputs[-1].context_messages - assert generator.inputs[-1].new_events == "" - assert "old context" not in generator.inputs[-1].context_messages - assert generator.inputs[-1].context_messages.index("new context") < generator.inputs[-1].context_messages.index( - "compact summary") - - -async def test_missing_checkpoint_recovers_only_newer_timestamped_events(tmp_path: Path, ) -> None: - """Ensure a missing checkpoint Event does not re-extract the transcript.""" - runtime = _runtime(tmp_path) - service, session = await _service_and_session(runtime) - await service.append_event(session, _event("event-old", "旧内容")) - await runtime.transcripts.append( - session.id, - { - "kind": "session-memory-checkpoint", - "checkpoint_id": "session-memory:event-missing", - "first_event_id": "event-missing", - "last_event_id": "event-missing", - }, - ) - await service.append_event(session, _event("event-new", "新内容")) - generator = FakeSessionMemoryGenerator() - - result = await SessionMemoryExtractor( - runtime, - generator, - ).extract_if_needed( - session, - _ctx(session), - force=True, - ) - - assert result.extracted is True - assert result.processed_events == 1 - assert generator.inputs[0].first_event_id == "event-new" - assert generator.inputs[0].last_event_id == "event-new" - - -async def test_no_new_events_does_not_call_generator(tmp_path: Path) -> None: - """Ensure no checkpoint increment means no repeated sub-agent call.""" - runtime = _runtime(tmp_path) - service, session = await _service_and_session(runtime) - generator = FakeSessionMemoryGenerator() - extractor = SessionMemoryExtractor(runtime, generator) - await service.append_event(session, _event("event-1", "first")) - await extractor.extract_if_needed(session, _ctx(session)) - - result = await extractor.extract_if_needed(session, _ctx(session)) - - assert result.reason == "no-new-events" - assert len(generator.inputs) == 1 - - -async def test_failure_does_not_advance_checkpoint_and_can_retry(tmp_path: Path) -> None: - """Ensure extraction failure leaves the increment for the next attempt.""" - runtime = _runtime(tmp_path) - service, session = await _service_and_session(runtime) - await service.append_event(session, _event("event-1", "first")) - failing_generator = FakeSessionMemoryGenerator(fail=True) - - failed = await SessionMemoryExtractor(runtime, failing_generator).extract_if_needed( - session, - _ctx(session), - ) - successful_generator = FakeSessionMemoryGenerator() - succeeded = await SessionMemoryExtractor(runtime, successful_generator).extract_if_needed( - session, - _ctx(session), - ) - - assert failed.reason == "extraction-failed" - assert succeeded.extracted is True - assert successful_generator.inputs[0].first_event_id == "event-1" - - -async def test_empty_document_does_not_overwrite_or_advance_checkpoint(tmp_path: Path, ) -> None: - """Ensure all-empty output fails and preserves old session memory.""" - runtime = _runtime(tmp_path) - service, session = await _service_and_session(runtime) - old_document = SessionMemoryDocument( - session_title="已有记忆", - current_state="等待新事件。", - ) - await runtime.session_memory.write(session.id, old_document) - await service.append_event(session, _event("event-1", "first")) - - result = await SessionMemoryExtractor( - runtime, - EmptySessionMemoryGenerator(), - ).extract_if_needed( - session, - _ctx(session), - force=True, - ) - - records = await runtime.transcripts.read_all(session.id) - assert result.reason == "extraction-failed" - assert await runtime.session_memory.read(session.id) == old_document.to_markdown() - assert not any(record.get("kind") == "session-memory-checkpoint" for record in records) - - -async def test_force_bypasses_initial_threshold(tmp_path: Path) -> None: - """Ensure forced extraction bypasses the initial character threshold.""" - runtime = _runtime(tmp_path, initial_chars=100_000, update_chars=100_000) - service, session = await _service_and_session(runtime) - await service.append_event(session, _event("event-1", "small")) - - result = await SessionMemoryExtractor( - runtime, - FakeSessionMemoryGenerator(), - ).extract_if_needed( - session, - _ctx(session), - force=True, - ) - - assert result.extracted is True - assert result.reason == "forced" - - -async def test_pending_tool_call_does_not_create_a_checkpoint_boundary(tmp_path: Path) -> None: - """Ensure a pending tool call cannot become a compaction boundary.""" - runtime = _runtime(tmp_path) - service, session = await _service_and_session(runtime) - tool_event = Event( - id="event-tool", - invocation_id="invocation-1", - author="agent", - content=Content( - role="model", - parts=[Part(function_call=FunctionCall( - id="call-1", - name="Read", - args={"file_path": "demo.py"}, - ))], - ), - ) - await service.append_event(session, tool_event) - - result = await SessionMemoryExtractor( - runtime, - FakeSessionMemoryGenerator(), - ).extract_if_needed( - session, - _ctx(session), - ) - - assert result.reason == "unsafe-boundary" - - -async def test_session_service_runs_extractor_after_old_summary(tmp_path: Path) -> None: - """Ensure the Runner post-turn extension automatically triggers extraction.""" - runtime = _runtime(tmp_path) - generator = FakeSessionMemoryGenerator() - extractor = SessionMemoryExtractor(runtime, generator) - service = TranscriptSessionService( - InMemorySessionService(), - runtime, - session_memory_extractor=extractor, - ) - session = await service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="demo-session", - ) - await service.append_event(session, _event("event-1", "post turn")) - - await service.create_session_summary(session, ctx=_ctx(session)) - - assert len(generator.inputs) == 1 - assert await runtime.session_memory.read(session.id) is not None - - -async def test_forked_generator_uses_isolated_runner_and_returns_memory() -> None: - """Ensure the default generator makes one isolated Markdown Runner call.""" - model = StructuredMemoryModel() - generator = ForkedSessionMemoryGenerator(model) - extraction_input = SessionMemoryExtractionInput( - current_memory=SessionMemoryDocument().to_markdown(), - first_event_id="event-1", - last_event_id="event-1", - context_messages="surrounding context", - ) - ctx = SimpleNamespace( - app_name="demo-app", - agent=SimpleNamespace(model=model), - ) - - document = await generator.generate(extraction_input, ctx) - - assert document.session_title == "隔离 Runner" - assert document.current_state == "子 Agent 已完成。" - assert len(model.requests) == 1 - assert "surrounding context" in model.requests[0].contents[-1].parts[0].text - - -async def test_forked_generator_rejects_empty_markdown_output() -> None: - """Ensure an empty Markdown response does not create empty session memory.""" - model = StructuredMemoryModel(empty=True) - generator = ForkedSessionMemoryGenerator(model) - extraction_input = SessionMemoryExtractionInput( - current_memory=SessionMemoryDocument().to_markdown(), - first_event_id="event-1", - last_event_id="event-1", - context_messages="new work", - ) - ctx = SimpleNamespace( - app_name="demo-app", - agent=SimpleNamespace(model=model), - ) - - with pytest.raises(ValueError): - await generator.generate(extraction_input, ctx) diff --git a/tests/advanced_memory/test_storage.py b/tests/advanced_memory/test_storage.py deleted file mode 100644 index d529e3e9c..000000000 --- a/tests/advanced_memory/test_storage.py +++ /dev/null @@ -1,334 +0,0 @@ -"""Unit tests for the independent Advanced Memory stores.""" - -from __future__ import annotations - -import asyncio -import json -import threading -from datetime import datetime -from datetime import timezone -from pathlib import Path - -import pytest - -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig -from trpc_agent_sdk.advanced_memory import AdvancedMemoryPaths -from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime -from trpc_agent_sdk.advanced_memory import MemoryDocument -from trpc_agent_sdk.advanced_memory import MemoryIndexEntry -from trpc_agent_sdk.advanced_memory import MemoryType -from trpc_agent_sdk.advanced_memory import SESSION_MEMORY_SECTIONS -from trpc_agent_sdk.advanced_memory import SessionMemoryDocument -from trpc_agent_sdk.advanced_memory import memory_freshness -from trpc_agent_sdk.advanced_memory import parse_memory_updated_at - - -def _enabled_config(tmp_path: Path, **overrides: object) -> AdvancedMemoryConfig: - """Create an enabled configuration rooted at the test directory.""" - return AdvancedMemoryConfig(enabled=True, root_dir=tmp_path, **overrides) - - -def test_config_reads_context_window_from_environment(monkeypatch: pytest.MonkeyPatch, ) -> None: - """Use both model limits from the environment when not provided explicitly.""" - monkeypatch.setenv("TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS", "128000") - monkeypatch.setenv("TRPC_AGENT_MAX_OUTPUT_TOKENS", "8192") - - config = AdvancedMemoryConfig() - - assert config.model_context_window_tokens == 128_000 - assert config.max_output_tokens == 8_192 - - -def test_config_rejects_invalid_context_window_environment(monkeypatch: pytest.MonkeyPatch, ) -> None: - """Reject invalid environment values with a clear configuration error.""" - monkeypatch.setenv("TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS", "not-a-number") - - with pytest.raises(ValueError, match="TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS"): - AdvancedMemoryConfig() - - -def test_config_rejects_invalid_max_output_tokens_environment(monkeypatch: pytest.MonkeyPatch, ) -> None: - """Reject invalid maximum output-token environment values.""" - monkeypatch.setenv("TRPC_AGENT_MAX_OUTPUT_TOKENS", "-1") - - with pytest.raises(ValueError, match="TRPC_AGENT_MAX_OUTPUT_TOKENS"): - AdvancedMemoryConfig() - - -async def test_disabled_runtime_does_not_create_directories(tmp_path: Path) -> None: - """Ensure disabled runtime initialization creates no directories.""" - runtime = AdvancedMemoryRuntime.create(AdvancedMemoryConfig(enabled=False, root_dir=tmp_path)) - - initialized = await runtime.initialize() - - assert initialized is False - assert not (tmp_path / "MEMORY").exists() - assert not (tmp_path / "SESSION").exists() - - -async def test_enabled_runtime_creates_expected_layout(tmp_path: Path) -> None: - """Ensure enabled initialization creates the expected empty layout.""" - runtime = AdvancedMemoryRuntime.create(_enabled_config(tmp_path)) - - initialized = await runtime.initialize() - - assert initialized is True - assert (tmp_path / "MEMORY" / "MEMORY.md").read_text() == "" - assert (tmp_path / "SESSION").is_dir() - - -async def test_long_term_memory_writes_index_and_topics(tmp_path: Path) -> None: - """Ensure the index and detail files share the MEMORY directory.""" - runtime = AdvancedMemoryRuntime.create(_enabled_config(tmp_path)) - await runtime.initialize() - - await runtime.long_term_memory.write_index([ - MemoryIndexEntry( - name="认证方案", - filename="auth.md", - summary="记录项目采用的认证方案", - ), - ]) - topic_path = await runtime.long_term_memory.write_topic( - "auth", - MemoryDocument( - name="认证方案", - description="记录项目采用的认证方案", - memory_type=MemoryType.PROJECT, - content="# Authentication\n\nUse OAuth.", - ), - ) - - assert await runtime.long_term_memory.read_index() == "- [认证方案](auth.md):记录项目采用的认证方案\n" - topic_content = await runtime.long_term_memory.read_topic("auth") - assert topic_content is not None - assert topic_content.startswith("---\n" - "name: 认证方案\n" - "description: 记录项目采用的认证方案\n" - "type: project\n" - "updated_at: ") - assert topic_content.endswith("---\n# Authentication\n\nUse OAuth.\n") - assert parse_memory_updated_at(topic_content) is not None - assert topic_path == tmp_path / "MEMORY" / "auth.md" - assert await runtime.long_term_memory.list_topics() == [topic_path] - - -async def test_memory_index_is_truncated_when_read_over_line_limit(tmp_path: Path) -> None: - """Ensure prompt reads respect the configured line limit without rejecting writes.""" - runtime = AdvancedMemoryRuntime.create(_enabled_config(tmp_path, memory_index_max_lines=2), ) - await runtime.initialize() - - await runtime.long_term_memory.write_index([ - MemoryIndexEntry(name="one", filename="one.md", summary="one"), - MemoryIndexEntry(name="two", filename="two.md", summary="two"), - MemoryIndexEntry(name="three", filename="three.md", summary="three"), - ]) - - assert (await runtime.long_term_memory.read_index()).splitlines() == [ - "- [one](one.md):one", - "- [two](two.md):two", - ] - - -async def test_session_memory_is_isolated_by_session_id(tmp_path: Path) -> None: - """Ensure structured summaries for different sessions do not overlap.""" - runtime = AdvancedMemoryRuntime.create(_enabled_config(tmp_path)) - await runtime.initialize() - - first_document = SessionMemoryDocument(session_title="会话 A", current_state="A") - second_document = SessionMemoryDocument(session_title="会话 B", current_state="B") - first_path = await runtime.session_memory.write("session-a", first_document) - second_path = await runtime.session_memory.write("session-b", second_document) - - assert first_path == tmp_path / "SESSION" / "session-a" / "session_memory.md" - assert second_path == tmp_path / "SESSION" / "session-b" / "session_memory.md" - first_content = await runtime.session_memory.read("session-a") - second_content = await runtime.session_memory.read("session-b") - assert first_content == first_document.to_markdown() - assert second_content == second_document.to_markdown() - assert first_content is not None - assert all(f"# {section}" in first_content for section in SESSION_MEMORY_SECTIONS) - - -async def test_transcript_appends_jsonl_in_order(tmp_path: Path) -> None: - """Ensure transcripts preserve order and payloads as JSONL.""" - runtime = AdvancedMemoryRuntime.create(_enabled_config(tmp_path)) - await runtime.initialize() - - transcript_path = await runtime.transcripts.append( - "session-a", - { - "kind": "user", - "payload": { - "text": "你好" - } - }, - ) - await runtime.transcripts.append( - "session-a", - { - "kind": "assistant", - "payload": { - "text": "你好" - } - }, - ) - - records = await runtime.transcripts.read_all("session-a") - raw_lines = transcript_path.read_text().splitlines() - assert [record["kind"] for record in records] == ["user", "assistant"] - assert records[0]["payload"] == {"text": "你好"} - assert all("recorded_at" in record for record in records) - assert len(raw_lines) == 2 - assert all(isinstance(json.loads(line), dict) for line in raw_lines) - - -async def test_transcript_append_unique_uses_persisted_ids(tmp_path: Path) -> None: - """Ensure transcript de-duplication recognizes persisted event IDs.""" - first_runtime = AdvancedMemoryRuntime.create(_enabled_config(tmp_path)) - await first_runtime.initialize() - await first_runtime.transcripts.append_unique( - "session-a", - { - "kind": "event", - "event_id": "event-1" - }, - unique_key="event_id", - ) - - second_runtime = AdvancedMemoryRuntime.create(_enabled_config(tmp_path)) - _, appended = await second_runtime.transcripts.append_unique( - "session-a", - { - "kind": "event", - "event_id": "event-1" - }, - unique_key="event_id", - ) - - assert appended is False - assert len(await second_runtime.transcripts.read_all("session-a")) == 1 - - -async def test_transcript_read_waits_for_in_progress_append( - tmp_path: Path, - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Ensure reads do not observe a partially written JSONL record.""" - runtime = AdvancedMemoryRuntime.create(_enabled_config(tmp_path)) - started = threading.Event() - release = threading.Event() - - def slow_append(path: Path, serialized: str) -> None: - """Pause under the write lock to simulate a partial write.""" - midpoint = len(serialized) // 2 - with path.open("a", encoding="utf-8") as transcript_file: - transcript_file.write(serialized[:midpoint]) - transcript_file.flush() - started.set() - release.wait(timeout=2) - transcript_file.write(serialized[midpoint:] + "\n") - transcript_file.flush() - - monkeypatch.setattr( - runtime.transcripts, - "_append_serialized_unlocked", - slow_append, - ) - append_task = asyncio.create_task(runtime.transcripts.append("session-a", {"kind": "event"})) - assert await asyncio.to_thread(started.wait, 2) - read_task = asyncio.create_task(runtime.transcripts.read_all("session-a")) - await asyncio.sleep(0.05) - - assert read_task.done() is False - release.set() - await append_task - records = await read_task - assert len(records) == 1 - assert records[0]["kind"] == "event" - assert "recorded_at" in records[0] - - -async def test_memory_index_is_truncated_when_read_over_byte_budget(tmp_path: Path) -> None: - """Ensure prompt reads respect the configured byte limit without rejecting writes.""" - config = AdvancedMemoryConfig( - enabled=True, - root_dir=tmp_path, - memory_index_max_bytes=80, - ) - runtime = AdvancedMemoryRuntime.create(config) - entries = [MemoryIndexEntry( - name="较长中文记忆名称", - filename="memory.md", - summary="这是一段会按 UTF-8 字节计数的较长中文概述", - )] - - await runtime.long_term_memory.write_index(entries) - - assert await runtime.long_term_memory.read_index() == "" - - -def test_paths_sanitize_external_identifiers(tmp_path: Path) -> None: - """Ensure session and topic identifiers cannot escape the root directory.""" - paths = AdvancedMemoryPaths(_enabled_config(tmp_path)) - - session_path = paths.session_dir("../../session") - topic_path = paths.memory_topic_path("../auth notes") - assert session_path.parent == tmp_path / "SESSION" - assert session_path.name.startswith("session-") - assert topic_path.parent == tmp_path / "MEMORY" - assert topic_path.name.startswith("auth_notes-") - assert topic_path.suffix == ".md" - assert paths.session_dir("session") != session_path - assert paths.memory_topic_path("auth_notes") != topic_path - - -def test_config_rejects_nested_path_components(tmp_path: Path) -> None: - """Ensure directory and file settings accept only safe path components.""" - with pytest.raises(ValueError, match="Invalid memory path component"): - AdvancedMemoryConfig(root_dir=tmp_path, memory_dir_name="../MEMORY") - - -def test_memory_freshness_uses_expected_buckets() -> None: - now = datetime(2026, 8, 18, 12, tzinfo=timezone.utc) - - assert memory_freshness(now, now=now) == "today" - assert memory_freshness( - datetime(2026, 8, 17, 13, tzinfo=timezone.utc), - now=now, - ) == "today" - assert memory_freshness( - datetime(2026, 8, 17, 0, tzinfo=timezone.utc), - now=now, - ) == "yesterday" - assert memory_freshness( - datetime(2026, 8, 12, 12, tzinfo=timezone.utc), - now=now, - ) == "within 7 days" - assert memory_freshness( - datetime(2026, 7, 25, 12, tzinfo=timezone.utc), - now=now, - ) == "within 30 days" - assert memory_freshness( - datetime(2026, 7, 1, 12, tzinfo=timezone.utc), - now=now, - ) == "over 30 days" - assert memory_freshness(None, now=now) == "unknown" - - -def test_parse_memory_updated_at_only_reads_frontmatter() -> None: - content = ("---\n" - "name: Example\n" - "description: Example memory\n" - "type: project\n" - "updated_at: 2026-08-18T10:00:00+00:00\n" - "---\n" - "The body mentions updated_at: 1999-01-01T00:00:00+00:00.\n") - - assert parse_memory_updated_at(content) == datetime( - 2026, - 8, - 18, - 10, - tzinfo=timezone.utc, - ) diff --git a/tests/advanced_memory/test_tool_result_budget.py b/tests/advanced_memory/test_tool_result_budget.py deleted file mode 100644 index 4df8d9036..000000000 --- a/tests/advanced_memory/test_tool_result_budget.py +++ /dev/null @@ -1,310 +0,0 @@ -"""Unit tests for tool-result context budgeting.""" - -from __future__ import annotations - -import json -from pathlib import Path -from types import SimpleNamespace - -import pytest - -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig -from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime -from trpc_agent_sdk.advanced_memory import setup_tool_result_budget -from trpc_agent_sdk.advanced_memory import ToolResultBudget -from trpc_agent_sdk.advanced_memory import ToolResultBudgetCallback -from trpc_agent_sdk.models import LlmRequest -from trpc_agent_sdk.types import Content -from trpc_agent_sdk.types import FunctionResponse -from trpc_agent_sdk.types import Part - - -def _runtime( - tmp_path: Path, - *, - enabled: bool = True, - per_result: int = 500, - per_message: int = 2_000, - preview: int = 50, -) -> AdvancedMemoryRuntime: - """Create an isolated runtime with small test limits.""" - return AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( - enabled=enabled, - root_dir=tmp_path, - tool_result_max_chars=per_result, - tool_results_per_message_max_chars=per_message, - tool_result_preview_chars=preview, - )) - - -def _request(*responses: tuple[str, str]) -> tuple[LlmRequest, list[Part]]: - """Create one user Content request from a tool ID and output text.""" - parts = [ - Part(function_response=FunctionResponse( - id=result_id, - name="demo_tool", - response={"output": output}, - )) for result_id, output in responses - ] - return LlmRequest(model="test-model", contents=[Content(role="user", parts=parts)]), parts - - -async def test_single_large_result_is_persisted_and_replaced(tmp_path: Path) -> None: - """Ensure oversized single results are persisted and previewed.""" - runtime = _runtime(tmp_path, per_result=200, per_message=5_000, preview=40) - budget = ToolResultBudget(runtime) - request, original_parts = _request(("result-1", "x" * 500)) - original_response = original_parts[0].function_response.response.copy() - - result = await budget.apply(request, session_id="session-a") - - replacement = request.contents[0].parts[0].function_response.response - assert result.replaced_count == 1 - assert "persisted_output" in replacement - assert replacement["persisted_output"]["truncated"] is True - assert original_parts[0].function_response.response == original_response - persisted = await runtime.tool_results.read("session-a", "result-1") - assert persisted is not None - assert '"output":"' in persisted - assert "x" * 100 in persisted - - -async def test_aggregate_budget_replaces_largest_fresh_results(tmp_path: Path) -> None: - """Ensure aggregate pressure replaces the largest new result first.""" - runtime = _runtime(tmp_path, per_result=5_000, per_message=2_300, preview=50) - budget = ToolResultBudget(runtime) - request, _ = _request( - ("small", "s" * 400), - ("largest", "l" * 1_400), - ("medium", "m" * 900), - ) - - result = await budget.apply(request, session_id="session-a") - - responses = {part.function_response.id: part.function_response.response for part in request.contents[0].parts} - assert result.replaced_count == 1 - assert "persisted_output" in responses["largest"] - assert responses["small"]["output"] == "s" * 400 - assert responses["medium"]["output"] == "m" * 900 - - -async def test_aggregate_budget_groups_consecutive_user_contents(tmp_path: Path, ) -> None: - """Ensure consecutive user Contents share one aggregate budget.""" - runtime = _runtime(tmp_path, per_result=5_000, per_message=1_800, preview=50) - request = LlmRequest( - model="test-model", - contents=[ - Content( - role="user", - parts=[ - Part(function_response=FunctionResponse( - id=f"result-{index}", - name="demo_tool", - response={"output": char * 1_100}, - )) - ], - ) for index, char in enumerate(("a", "b")) - ], - ) - - result = await ToolResultBudget(runtime).apply( - request, - session_id="session-a", - ) - - responses = [content.parts[0].function_response.response for content in request.contents] - assert result.replaced_count == 1 - assert sum("persisted_output" in response for response in responses) == 1 - - -async def test_model_content_starts_a_new_aggregate_budget_group(tmp_path: Path, ) -> None: - """Ensure results after a model boundary are not merged with the prior group.""" - runtime = _runtime(tmp_path, per_result=5_000, per_message=1_800, preview=50) - request = LlmRequest( - model="test-model", - contents=[ - Content( - role="user", - parts=[ - Part(function_response=FunctionResponse( - id="result-1", - name="demo_tool", - response={"output": "a" * 1_100}, - )) - ], - ), - Content( - role="model", - parts=[Part.from_text(text="继续调用工具")], - ), - Content( - role="user", - parts=[ - Part(function_response=FunctionResponse( - id="result-2", - name="demo_tool", - response={"output": "b" * 1_100}, - )) - ], - ), - ], - ) - - result = await ToolResultBudget(runtime).apply( - request, - session_id="session-a", - ) - - assert result.replaced_count == 0 - - -async def test_reapplying_budget_uses_exact_cached_replacement(tmp_path: Path) -> None: - """Ensure repeated requests reuse replacements without duplicate records.""" - runtime = _runtime(tmp_path, per_result=200, per_message=5_000, preview=40) - budget = ToolResultBudget(runtime) - first_request, _ = _request(("result-1", "x" * 500)) - await budget.apply(first_request, session_id="session-a") - first_replacement = first_request.contents[0].parts[0].function_response.response - - second_request, _ = _request(("result-1", "x" * 500)) - second_result = await budget.apply(second_request, session_id="session-a") - second_replacement = second_request.contents[0].parts[0].function_response.response - - records = await runtime.transcripts.read_all("session-a") - replacement_records = [record for record in records if record["kind"] == "content-replacement"] - assert second_result.replaced_count == 0 - assert second_replacement == first_replacement - assert len(replacement_records) == 1 - - -async def test_unreplaced_result_remains_frozen_after_restart(tmp_path: Path) -> None: - """Ensure already-sent results do not change after restart or lower limits.""" - first_runtime = _runtime(tmp_path, per_result=2_000, per_message=5_000, preview=40) - first_budget = ToolResultBudget(first_runtime) - first_request, _ = _request(("result-1", "x" * 500)) - await first_budget.apply(first_request, session_id="session-a") - - second_runtime = _runtime(tmp_path, per_result=200, per_message=1_000, preview=40) - second_budget = ToolResultBudget(second_runtime) - second_request, _ = _request(("result-1", "x" * 500)) - second_result = await second_budget.apply(second_request, session_id="session-a") - - response = second_request.contents[0].parts[0].function_response.response - assert second_result.replaced_count == 0 - assert response["output"] == "x" * 500 - assert await second_runtime.tool_results.read("session-a", "result-1") is None - - -async def test_disabled_budget_does_not_copy_or_persist_request(tmp_path: Path) -> None: - """Ensure disabled mode preserves the request and disk state.""" - runtime = _runtime(tmp_path, enabled=False) - budget = ToolResultBudget(runtime) - request, original_parts = _request(("result-1", "x" * 1_000)) - original_content = request.contents[0] - - result = await budget.apply(request, session_id="session-a") - - assert result.replaced_count == 0 - assert request.contents[0] is original_content - assert request.contents[0].parts[0] is original_parts[0] - assert not (tmp_path / "MEMORY").exists() - assert not (tmp_path / "SESSION").exists() - - -async def test_exact_single_result_limit_is_not_replaced(tmp_path: Path) -> None: - """Ensure a result exactly at the per-item limit is not replaced.""" - probe_request, _ = _request(("result-1", "x" * 100)) - probe_response = probe_request.contents[0].parts[0].function_response.response - serialized_size = len(json.dumps( - probe_response, - ensure_ascii=False, - sort_keys=True, - separators=(",", ":"), - )) - runtime = _runtime( - tmp_path, - per_result=serialized_size, - per_message=5_000, - preview=20, - ) - request, _ = _request(("result-1", "x" * 100)) - - result = await ToolResultBudget(runtime).apply(request, session_id="session-a") - - assert result.replaced_count == 0 - assert request.contents[0].parts[0].function_response.response["output"] == "x" * 100 - - -async def test_aggregate_budget_is_independent_across_model_boundaries(tmp_path: Path, ) -> None: - """Ensure model-separated result groups budget independently.""" - runtime = _runtime(tmp_path, per_result=5_000, per_message=1_500, preview=40) - first_request, _ = _request(("first", "a" * 1_000)) - second_request, _ = _request(("second", "b" * 1_000)) - request = LlmRequest( - model="test-model", - contents=[ - first_request.contents[0], - Content(role="model", parts=[Part.from_text(text="next")]), - second_request.contents[0], - ], - ) - - result = await ToolResultBudget(runtime).apply(request, session_id="session-a") - - assert result.replaced_count == 0 - assert request.contents[0].parts[0].function_response.response["output"] == "a" * 1_000 - assert request.contents[2].parts[0].function_response.response["output"] == "b" * 1_000 - - -def test_setup_preserves_existing_callback_and_is_idempotent(tmp_path: Path) -> None: - """Ensure setup preserves callbacks and is idempotent.""" - - async def existing_callback(ctx, request): - """Simulate an existing model pre-callback.""" - return None - - agent = SimpleNamespace(before_model_callback=existing_callback) - runtime = _runtime(tmp_path) - - first_budget = setup_tool_result_budget(agent, runtime) - second_budget = setup_tool_result_budget(agent, runtime) - - assert first_budget is second_budget - assert agent.before_model_callback[0] is existing_callback - assert isinstance(agent.before_model_callback[1], ToolResultBudgetCallback) - assert len(agent.before_model_callback) == 2 - - -async def test_reused_result_id_with_different_content_is_rejected(tmp_path: Path) -> None: - """Ensure conflicting duplicate tool IDs fail instead of reusing replacements.""" - runtime = _runtime(tmp_path) - request, _ = _request( - ("duplicate", "first"), - ("duplicate", "second"), - ) - - with pytest.raises(ValueError, match="reused with different content"): - await ToolResultBudget(runtime).apply(request, session_id="session-a") - - -async def test_reused_result_id_after_restart_is_rejected(tmp_path: Path) -> None: - """Ensure transcript state rejects tool-ID conflicts after restart.""" - first_runtime = _runtime(tmp_path, per_result=200, per_message=5_000, preview=40) - first_request, _ = _request(("result-1", "first" * 100)) - await ToolResultBudget(first_runtime).apply(first_request, session_id="session-a") - - second_runtime = _runtime(tmp_path, per_result=200, per_message=5_000, preview=40) - second_request, _ = _request(("result-1", "second" * 100)) - - with pytest.raises(ValueError, match="reused with different content"): - await ToolResultBudget(second_runtime).apply(second_request, session_id="session-a") - - -def test_setup_rejects_another_runtime_for_same_agent(tmp_path: Path) -> None: - """Ensure one Agent cannot silently bind two budget runtimes.""" - agent = SimpleNamespace(before_model_callback=None) - setup_tool_result_budget(agent, _runtime(tmp_path / "first")) - - with pytest.raises(ValueError, match="another runtime"): - setup_tool_result_budget(agent, _runtime(tmp_path / "second")) diff --git a/tests/advanced_memory/test_transcript_session_service.py b/tests/advanced_memory/test_transcript_session_service.py deleted file mode 100644 index a23217011..000000000 --- a/tests/advanced_memory/test_transcript_session_service.py +++ /dev/null @@ -1,127 +0,0 @@ -"""Unit tests for TranscriptSessionService automatic recording.""" - -from __future__ import annotations - -from pathlib import Path - -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig -from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime -from trpc_agent_sdk.advanced_memory import TranscriptSessionService -from trpc_agent_sdk.events import Event -from trpc_agent_sdk.sessions import InMemorySessionService -from trpc_agent_sdk.types import Content -from trpc_agent_sdk.types import Part - - -def _event(event_id: str, text: str, *, partial: bool = False) -> Event: - """Create a fixed Event for transcript tests.""" - return Event( - id=event_id, - invocation_id="invocation-1", - author="agent", - content=Content(parts=[Part.from_text(text=text)]), - partial=partial, - ) - - -async def _session(service: TranscriptSessionService): - """Create a test session through the decorated service.""" - return await service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="demo-session", - ) - - -async def test_append_event_writes_versioned_parent_chain(tmp_path: Path) -> None: - """Ensure persisted Events produce an ordered parent-linked transcript.""" - runtime = AdvancedMemoryRuntime.create(AdvancedMemoryConfig(enabled=True, root_dir=tmp_path)) - service = TranscriptSessionService(InMemorySessionService(), runtime) - session = await _session(service) - - await service.append_event(session, _event("event-1", "hello")) - await service.append_event(session, _event("event-2", "world")) - - records = await runtime.transcripts.read_all(session.id) - assert [record["event_id"] for record in records] == ["event-1", "event-2"] - assert records[0]["parent_event_id"] is None - assert records[1]["parent_event_id"] == "event-1" - assert records[0]["schema_version"] == 1 - assert records[0]["session"] == { - "id": "demo-session", - "app_name": "demo-app", - "user_id": "demo-user", - } - assert records[0]["event"]["invocationId"] == "invocation-1" - - -async def test_duplicate_event_id_is_not_written_twice(tmp_path: Path) -> None: - """Ensure duplicate Event IDs are not written twice.""" - runtime = AdvancedMemoryRuntime.create(AdvancedMemoryConfig(enabled=True, root_dir=tmp_path)) - service = TranscriptSessionService(InMemorySessionService(), runtime) - session = await _session(service) - duplicate = _event("event-1", "hello") - - await service.append_event(session, duplicate) - await service.append_event(session, duplicate.model_copy(deep=True)) - - records = await runtime.transcripts.read_all(session.id) - assert [record["event_id"] for record in records] == ["event-1"] - - -async def test_old_duplicate_does_not_rewind_parent_chain(tmp_path: Path) -> None: - """Ensure replaying an old Event does not rewind the parent chain.""" - runtime = AdvancedMemoryRuntime.create(AdvancedMemoryConfig(enabled=True, root_dir=tmp_path)) - service = TranscriptSessionService(InMemorySessionService(), runtime) - session = await _session(service) - await service.append_event(session, _event("event-1", "first")) - await service.append_event(session, _event("event-2", "second")) - await service.append_event(session, _event("event-1", "first")) - await service.append_event(session, _event("event-3", "third")) - - records = await runtime.transcripts.read_all(session.id) - - assert [record["event_id"] for record in records] == ["event-1", "event-2", "event-3"] - assert records[-1]["parent_event_id"] == "event-2" - - -async def test_new_wrapper_restores_parent_from_existing_transcript(tmp_path: Path) -> None: - """Ensure a rebuilt wrapper restores the parent-chain tail from disk.""" - runtime = AdvancedMemoryRuntime.create(AdvancedMemoryConfig(enabled=True, root_dir=tmp_path)) - delegate = InMemorySessionService() - first_service = TranscriptSessionService(delegate, runtime) - session = await _session(first_service) - await first_service.append_event(session, _event("event-1", "first")) - - second_runtime = AdvancedMemoryRuntime.create(AdvancedMemoryConfig(enabled=True, root_dir=tmp_path)) - second_service = TranscriptSessionService(delegate, second_runtime) - await second_service.append_event(session, _event("event-2", "second")) - - records = await second_runtime.transcripts.read_all(session.id) - assert records[-1]["parent_event_id"] == "event-1" - - -async def test_disabled_runtime_preserves_old_service_without_disk_writes(tmp_path: Path) -> None: - """Ensure disabled mode preserves the legacy service without disk writes.""" - runtime = AdvancedMemoryRuntime.create(AdvancedMemoryConfig(enabled=False, root_dir=tmp_path)) - service = TranscriptSessionService(InMemorySessionService(), runtime) - session = await _session(service) - - persisted_event = await service.append_event(session, _event("event-1", "hello")) - - assert persisted_event.id == "event-1" - assert [event.id for event in session.events] == ["event-1"] - assert not (tmp_path / "MEMORY").exists() - assert not (tmp_path / "SESSION").exists() - - -async def test_partial_event_is_not_written_to_transcript(tmp_path: Path) -> None: - """Ensure streaming partial Events enter neither session nor transcript.""" - runtime = AdvancedMemoryRuntime.create(AdvancedMemoryConfig(enabled=True, root_dir=tmp_path)) - service = TranscriptSessionService(InMemorySessionService(), runtime) - session = await _session(service) - - await service.append_event(session, _event("partial-1", "chunk", partial=True)) - - assert session.events == [] - assert await runtime.transcripts.read_all(session.id) == [] diff --git a/tests/advanced_memory/test_coordination.py b/tests/sessions/compact/test_coordination.py similarity index 92% rename from tests/advanced_memory/test_coordination.py rename to tests/sessions/compact/test_coordination.py index 030d01bc5..3dba41f86 100644 --- a/tests/advanced_memory/test_coordination.py +++ b/tests/sessions/compact/test_coordination.py @@ -6,7 +6,7 @@ import pytest -from trpc_agent_sdk.advanced_memory._coordination import CrossLoopLock +from trpc_agent_sdk.sessions.compact.advanced._coordination import CrossLoopLock @pytest.mark.asyncio diff --git a/tests/sessions/compact/test_session_compact.py b/tests/sessions/compact/test_session_compact.py new file mode 100644 index 000000000..f9ef740ef --- /dev/null +++ b/tests/sessions/compact/test_session_compact.py @@ -0,0 +1,268 @@ +"""Tests for SessionService-owned Session Compact.""" + +from __future__ import annotations + +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest + +from trpc_agent_sdk.abc import CompactSummarizerABC +from trpc_agent_sdk.abc import CompactTrigger +from trpc_agent_sdk.context import new_agent_context +from trpc_agent_sdk.events import Event +from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.models import LlmResponse +from trpc_agent_sdk.sessions import InMemorySessionService +from trpc_agent_sdk.sessions import SessionServiceConfig +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizer +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerConfig +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerManager +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerRuntime +from trpc_agent_sdk.sessions.compact import AutoCompactSummarizerConfig +from trpc_agent_sdk.sessions.compact import SessionMemoryExtractor +from trpc_agent_sdk.sessions.compact.advanced import SessionMemoryExtractorConfig +from trpc_agent_sdk.sessions.compact.advanced import ToolResultBudgetConfig +from trpc_agent_sdk.sessions.compact.advanced._tool_result_budget import ToolResultBudget +from trpc_agent_sdk.types import Content +from trpc_agent_sdk.types import FunctionResponse +from trpc_agent_sdk.types import Part + + +class _MemoryModel: + name = "memory-model" + + async def generate_async(self, request, *, stream, ctx): + del request, stream, ctx + yield LlmResponse( + content=Content( + role="model", + parts=[Part.from_text(text="# Session Title\nTest session\n\n# Current State\nUpdated")], + ), + ) + + +class _SummaryModel: + name = "summary-model" + + async def generate_async(self, request, *, stream, ctx): + del request, stream, ctx + yield LlmResponse(content=Content( + role="model", + parts=[ + Part.from_text( + text="covered" + "# Session Title\nCompact summary") + ], + )) + + +def _event(event_id: str, content: Content) -> Event: + return Event( + id=event_id, + invocation_id="invocation", + author="agent", + content=content, + ) + + +def test_compact_runtime_has_no_external_storage() -> None: + runtime = AdvancedAutoCompactSummarizerRuntime(AdvancedAutoCompactSummarizerConfig()) + + assert not hasattr(runtime, "transcripts") + assert not hasattr(runtime, "tool_results") + assert not hasattr(runtime, "paths") + + +@pytest.mark.asyncio +async def test_disabled_auto_compact_returns_without_recursion() -> None: + config = AdvancedAutoCompactSummarizerConfig( + auto_compact=AutoCompactSummarizerConfig(enabled=False), + ) + summarizer = AdvancedAutoCompactSummarizer(config) + session = SimpleNamespace(id="session", app_name="app", user_id="user") + ctx = SimpleNamespace(session=session, session_id=session.id) + + result = await summarizer.apply(LlmRequest(), ctx=ctx) + + assert result.compacted is False + assert result.blocked is False + + +@pytest.mark.asyncio +async def test_session_service_accepts_a_configured_compact_manager() -> None: + summarizer = AdvancedAutoCompactSummarizer(AdvancedAutoCompactSummarizerConfig()) + manager = AdvancedAutoCompactSummarizerManager(summarizer) + service = InMemorySessionService( + session_config=SessionServiceConfig(store_historical_events=True), + summarizer_manager=manager, + ) + + assert service.summarizer_manager is manager + await service.close() + + +@pytest.mark.asyncio +async def test_advanced_summarizer_implements_compact_abc_and_timing() -> None: + summarizer = AdvancedAutoCompactSummarizer(AdvancedAutoCompactSummarizerConfig( + session_memory=SessionMemoryExtractorConfig(enabled=False), + )) + before_model_manager = AdvancedAutoCompactSummarizerManager(summarizer) + after_turn_manager = AdvancedAutoCompactSummarizerManager( + summarizer, + compact_trigger=CompactTrigger.AFTER_TURN, + ) + + assert isinstance(summarizer, CompactSummarizerABC) + assert before_model_manager.compact_trigger == CompactTrigger.BEFORE_MODEL + assert after_turn_manager.compact_trigger == CompactTrigger.AFTER_TURN + + session = SimpleNamespace(events=[]) + ctx = SimpleNamespace(session=session) + summarizer.should_summarize = AsyncMock(return_value=True) + summarizer.create_session_summary = AsyncMock(return_value="summary") + await before_model_manager.create_session_summary(session, ctx=ctx) + summarizer.should_summarize.assert_not_awaited() + + await after_turn_manager.create_session_summary(session, ctx=ctx) + summarizer.should_summarize.assert_awaited_once_with(session) + summarizer.create_session_summary.assert_awaited_once_with( + session, + ctx=ctx, + store_historical_events=True, + ) + + +@pytest.mark.asyncio +async def test_advanced_end_of_turn_compaction_moves_old_events_to_history() -> None: + service = InMemorySessionService( + session_config=SessionServiceConfig(store_historical_events=True), + ) + session = await service.create_session( + app_name="app", + user_id="user", + session_id="session", + ) + for index in range(3): + await service.append_event( + session, + _event( + f"event-{index}", + Content(role="user", parts=[Part.from_text(text=f"{index}:" + "x" * 2_000)]), + ), + ) + original_ids = [event.id for event in session.events] + config = AdvancedAutoCompactSummarizerConfig( + session_memory=SessionMemoryExtractorConfig(enabled=False), + auto_compact=AutoCompactSummarizerConfig( + trigger_chars=500, + target_chars=250, + blocking_chars=10_000, + keep_recent_contents=1, + ), + ) + summarizer = AdvancedAutoCompactSummarizer(config, model=_SummaryModel()) + ctx = SimpleNamespace( + session=session, + session_id=session.id, + session_service=service, + agent_context=new_agent_context(), + agent=SimpleNamespace(model=_SummaryModel()), + ) + + assert await summarizer.should_summarize(session) is True + summary = await summarizer.create_session_summary( + session, + ctx=ctx, + store_historical_events=True, + ) + + assert summary is not None and "Compact summary" in summary + assert len(session.events) == 2 + assert session.events[0].is_summary_event() + assert session.events[1].id == original_ids[-1] + assert [event.id for event in session.historical_events] == original_ids[:-1] + await service.close() + + +@pytest.mark.asyncio +async def test_session_memory_is_written_to_session_state() -> None: + service = InMemorySessionService(session_config=SessionServiceConfig(store_historical_events=True), ) + session = await service.create_session( + app_name="app", + user_id="user", + session_id="session", + ) + await service.append_event( + session, + _event("event-1", Content(parts=[Part.from_text(text="hello")])), + ) + config = AdvancedAutoCompactSummarizerConfig( + session_memory=SessionMemoryExtractorConfig( + initial_chars=1, + update_chars=1, + ), + ) + extractor = SessionMemoryExtractor( + AdvancedAutoCompactSummarizerRuntime(config), + model=_MemoryModel(), + ) + + result = await extractor.extract_if_needed( + SimpleNamespace( + session=session, + session_service=service, + agent=SimpleNamespace(model=_MemoryModel(), generate_content_config=None), + override_messages=None, + ), + force=True, + ) + + assert result.extracted is True + assert "_trpc_agent:summary" in session.state + stored = await service.get_session( + app_name="app", + user_id="user", + session_id="session", + ) + assert stored is not None + assert "_trpc_agent:summary" in stored.state + await service.close() + + +@pytest.mark.asyncio +async def test_tool_result_budget_keeps_the_session_event_id() -> None: + service = InMemorySessionService() + session = await service.create_session( + app_name="app", + user_id="user", + session_id="session", + ) + content = Content(parts=[ + Part(function_response=FunctionResponse( + id="tool-call-1", + name="demo", + response={"output": "x" * 500}, + )) + ]) + await service.append_event(session, _event("event-tool", content)) + request = LlmRequest(model="test", contents=[content.model_copy(deep=True)]) + config = AdvancedAutoCompactSummarizerConfig( + tool_result_budget=ToolResultBudgetConfig( + max_chars=100, + preview_chars=20, + ), + ) + budget = ToolResultBudget( + AdvancedAutoCompactSummarizerRuntime(config), + ) + + await budget.apply( + request, + ctx=SimpleNamespace(session=session), + ) + + replacement = request.contents[0].parts[0].function_response.response + assert replacement["session_event_id"] == "event-tool" + assert "path" not in replacement + await service.close() diff --git a/tests/advanced_memory/test_token_budget.py b/tests/sessions/compact/test_token_budget.py similarity index 79% rename from tests/advanced_memory/test_token_budget.py rename to tests/sessions/compact/test_token_budget.py index 0431f9515..ead1cd6cc 100644 --- a/tests/advanced_memory/test_token_budget.py +++ b/tests/sessions/compact/test_token_budget.py @@ -4,9 +4,9 @@ from types import SimpleNamespace -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig -from trpc_agent_sdk.advanced_memory import TokenContextTracker from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.sessions.compact.advanced import TokenContextTrackerConfig +from trpc_agent_sdk.sessions.compact.advanced._token_budget import TokenContextTracker from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import Part @@ -37,9 +37,8 @@ def test_usage_baseline_adds_only_contents_after_matching_event(tmp_path) -> Non agent=SimpleNamespace(model="test-model"), ) tracker = TokenContextTracker( - AdvancedMemoryConfig( + TokenContextTrackerConfig( enabled=True, - root_dir=tmp_path, model_context_window_tokens=1_000, max_output_tokens=100, )) @@ -63,7 +62,7 @@ def test_usage_boundary_mismatch_falls_back_to_full_request_estimate(tmp_path) - session=SimpleNamespace(events=[event]), agent=SimpleNamespace(model="test-model"), ) - tracker = TokenContextTracker(AdvancedMemoryConfig(enabled=True, root_dir=tmp_path)) + tracker = TokenContextTracker(TokenContextTrackerConfig(enabled=True)) estimate = tracker.estimate(request, ctx) @@ -85,7 +84,7 @@ def test_changed_recorded_system_or_tool_fingerprint_falls_back(tmp_path) -> Non agent=SimpleNamespace(model="test-model"), ) - estimate = TokenContextTracker(AdvancedMemoryConfig(enabled=True, root_dir=tmp_path)).estimate(request, ctx) + estimate = TokenContextTracker(TokenContextTrackerConfig(enabled=True)).estimate(request, ctx) assert estimate.source == "estimated" assert estimate.tokens < 999_999 @@ -94,25 +93,34 @@ def test_changed_recorded_system_or_tool_fingerprint_falls_back(tmp_path) -> Non def test_budget_reserves_max_output_and_calculates_three_thresholds(tmp_path) -> None: """Ensure thresholds use the window after reserving max output.""" tracker = TokenContextTracker( - AdvancedMemoryConfig( + TokenContextTrackerConfig( enabled=True, - root_dir=tmp_path, model_context_window_tokens=10_000, max_output_tokens=2_000, )) - budget = tracker.budget(_request("测试请求")) + ctx = SimpleNamespace( + session=SimpleNamespace(events=[]), + agent=SimpleNamespace(model="test-model"), + ) + budget = tracker.budget(_request("测试请求"), ctx) assert budget.effective_window_tokens == 8_000 assert budget.warning_threshold_tokens == 6_800 - assert budget.autocompact_threshold_tokens == 7_200 + assert budget.auto_compact_threshold_tokens == 7_200 assert budget.blocking_threshold_tokens == 7_600 def test_no_window_keeps_compatibility_mode(tmp_path) -> None: """Ensure token decisions remain disabled without a model window.""" - budget = TokenContextTracker(AdvancedMemoryConfig(enabled=True, - root_dir=tmp_path)).budget(_request("compatibility request")) + ctx = SimpleNamespace( + session=SimpleNamespace(events=[]), + agent=SimpleNamespace(model="test-model"), + ) + budget = TokenContextTracker(TokenContextTrackerConfig(enabled=True)).budget( + _request("compatibility request"), + ctx, + ) assert not budget.token_mode_enabled assert budget.estimate.source == "estimated" diff --git a/tests/sessions/replay/backends.py b/tests/sessions/replay/backends.py index 5c638a0fe..c5b5a7462 100644 --- a/tests/sessions/replay/backends.py +++ b/tests/sessions/replay/backends.py @@ -25,7 +25,9 @@ from trpc_agent_sdk.sessions import SessionServiceConfig from trpc_agent_sdk.sessions import SessionSummarizer from trpc_agent_sdk.sessions import SqlSessionService -from trpc_agent_sdk.sessions._summarizer_manager import SummarizerSessionManager +from trpc_agent_sdk.sessions.compact.default._summarizer_manager import ( + DefaultSessionSummarizerManager as SummarizerSessionManager, +) from .harness import ReplayBackend from .report import BackendStatus diff --git a/tests/sessions/session_memory_summary_diff_report.json b/tests/sessions/session_memory_summary_diff_report.json index 8d3240a0d..daa8dd7ef 100644 --- a/tests/sessions/session_memory_summary_diff_report.json +++ b/tests/sessions/session_memory_summary_diff_report.json @@ -203,7 +203,8 @@ "tool_call": null, "tool_response": null, "part_metadata": null, - "audio_transcription": null + "audio_transcription": null, + "media_processing": null } ], "role": "model" @@ -269,7 +270,8 @@ "tool_call": null, "tool_response": null, "part_metadata": null, - "audio_transcription": null + "audio_transcription": null, + "media_processing": null } ], "role": "model" @@ -409,7 +411,8 @@ "tool_call": null, "tool_response": null, "part_metadata": null, - "audio_transcription": null + "audio_transcription": null, + "media_processing": null } ], "role": "model" @@ -475,7 +478,8 @@ "tool_call": null, "tool_response": null, "part_metadata": null, - "audio_transcription": null + "audio_transcription": null, + "media_processing": null } ], "role": "model" diff --git a/tests/sessions/test_base_session_service.py b/tests/sessions/test_base_session_service.py index bcb63037a..bc6fb733f 100644 --- a/tests/sessions/test_base_session_service.py +++ b/tests/sessions/test_base_session_service.py @@ -13,15 +13,15 @@ from __future__ import annotations import time -from unittest.mock import AsyncMock, MagicMock, patch - -import pytest +from unittest.mock import AsyncMock, MagicMock from trpc_agent_sdk.abc import ListSessionsResponse from trpc_agent_sdk.events import Event from trpc_agent_sdk.sessions._base_session_service import BaseSessionService from trpc_agent_sdk.sessions._session import Session -from trpc_agent_sdk.sessions._summarizer_manager import SummarizerSessionManager +from trpc_agent_sdk.sessions.compact.default._summarizer_manager import ( + DefaultSessionSummarizerManager as SummarizerSessionManager, +) from trpc_agent_sdk.sessions._types import SessionServiceConfig from trpc_agent_sdk.types import Content, EventActions, Part, State @@ -344,6 +344,16 @@ async def test_update_session_default_noop(self): session = _make_session() await svc.update_session(session) + async def test_update_session_state_falls_back_to_full_update(self): + svc = ConcreteSessionService() + svc.update_session = AsyncMock() + session = _make_session() + + await svc.update_session_state(session, {"summary": {"version": 1}}) + + assert session.state["summary"] == {"version": 1} + svc.update_session.assert_awaited_once_with(session) + async def test_close(self): svc = ConcreteSessionService() await svc.close() diff --git a/tests/sessions/test_in_memory_session_service.py b/tests/sessions/test_in_memory_session_service.py index 174daa311..08e93953b 100644 --- a/tests/sessions/test_in_memory_session_service.py +++ b/tests/sessions/test_in_memory_session_service.py @@ -394,6 +394,52 @@ async def test_update_existing(self): await svc.update_session(session) await svc.close() + async def test_update_persists_compacted_active_and_historical_events(self): + svc = InMemorySessionService( + session_config=_make_session_config(store_historical_events=True), + ) + session = await svc.create_session(app_name="app", user_id="user", session_id="s1") + original = [_make_event(text=f"msg{i}") for i in range(4)] + for event in original: + await svc.append_event(session, event) + summary = _make_event(author="system", text="summary") + + assert session.compact_events( + summary, + original[1].id, + compaction_id="compact-1", + ) + await svc.update_session(session) + + stored = await svc.get_session(app_name="app", user_id="user", session_id="s1") + assert stored is not None + assert stored.events[0].is_summary_event() + assert [event.id for event in stored.events[1:]] == [event.id for event in original[2:]] + assert [event.id for event in stored.historical_events] == [event.id for event in original[:2]] + await svc.close() + + async def test_patch_state_preserves_stored_events(self): + svc = InMemorySessionService(session_config=_make_session_config()) + session = await svc.create_session( + app_name="app", + user_id="user", + session_id="s1", + ) + await svc.append_event(session, _make_event(text="keep me")) + stale = session.model_copy(deep=True) + stale.events = [] + + await svc.update_session_state(stale, {"_trpc_agent:summary": {"v": 1}}) + + stored = await svc.get_session( + app_name="app", + user_id="user", + session_id="s1", + ) + assert [event.content.parts[0].text for event in stored.events] == ["keep me"] + assert stored.state["_trpc_agent:summary"] == {"v": 1} + await svc.close() + async def test_update_nonexistent_app(self): svc = InMemorySessionService(session_config=_make_session_config()) session = Session(id="s1", app_name="nonexistent", user_id="user", save_key="k") diff --git a/tests/sessions/test_redis_session_service.py b/tests/sessions/test_redis_session_service.py index 8269b1862..a2e1f49d2 100644 --- a/tests/sessions/test_redis_session_service.py +++ b/tests/sessions/test_redis_session_service.py @@ -92,6 +92,20 @@ async def execute_command(self, session, command): elif method == 'hgetall': key = args[0] return self._hash_store.get(key, {}) + elif method == 'eval': + key = args[2] + raw = self._store.get(key) + if raw is None: + return None + value = json.loads(raw) + value.setdefault("state", {}).update(json.loads(args[3])) + if "last_update_time" in value: + value["last_update_time"] = args[4] + if "lastUpdateTime" in value: + value["lastUpdateTime"] = args[4] + encoded = json.dumps(value) + self._store[key] = encoded + return encoded return None async def delete(self, session, key): @@ -329,6 +343,76 @@ async def test_update_existing(self): assert stored.state.get("new_key") == "new_val" await svc.close() + async def test_update_persists_compacted_active_and_historical_events(self): + config = _make_config(store_historical_events=True) + svc = _create_service(config=config) + session = await svc.create_session(app_name="app", user_id="user", session_id="s1") + original = [_make_event(text=f"msg{i}") for i in range(4)] + for event in original: + await svc.append_event(session, event) + summary = _make_event(author="system", text="summary") + + assert session.compact_events( + summary, + original[1].id, + compaction_id="compact-1", + ) + await svc.update_session(session) + + stored = await svc.get_session(app_name="app", user_id="user", session_id="s1") + assert stored is not None + assert stored.events[0].is_summary_event() + assert [event.id for event in stored.events[1:]] == [event.id for event in original[2:]] + assert [event.id for event in stored.historical_events] == [event.id for event in original[:2]] + await svc.close() + + async def test_patch_state_preserves_stored_events(self): + svc = _create_service() + session = await svc.create_session( + app_name="app", + user_id="user", + session_id="s1", + ) + await svc.append_event(session, _make_event(text="keep me")) + + stale = session.model_copy(deep=True) + stale.events = [] + await svc.update_session_state(stale, {"_trpc_agent:summary": {"v": 1}}) + + stored = await svc.get_session( + app_name="app", + user_id="user", + session_id="s1", + ) + assert [event.content.parts[0].text for event in stored.events] == ["keep me"] + assert stored.state["_trpc_agent:summary"] == {"v": 1} + await svc.close() + + async def test_patch_state_repairs_lua_empty_array_encoding(self): + config = _make_config(store_historical_events=True) + svc = _create_service(config=config) + session = await svc.create_session( + app_name="app", + user_id="user", + session_id="s1", + ) + key = "session:app:user:s1" + payload = json.loads(svc._redis_storage._store[key]) + payload["historical_events"] = {} + svc._redis_storage._store[key] = json.dumps(payload) + + loaded = await svc.get_session( + app_name="app", + user_id="user", + session_id="s1", + ) + assert loaded is not None + assert loaded.historical_events == [] + + await svc.update_session_state(loaded, {"_trpc_agent:summary": {"v": 1}}) + assert loaded.state["_trpc_agent:summary"] == {"v": 1} + await svc.close() + async def test_update_nonexistent(self): svc = _create_service() session = _make_session_obj(id="nonexistent") diff --git a/tests/sessions/test_session_summarizer.py b/tests/sessions/test_session_summarizer.py index ef46db4bb..472d31676 100644 --- a/tests/sessions/test_session_summarizer.py +++ b/tests/sessions/test_session_summarizer.py @@ -21,10 +21,10 @@ from trpc_agent_sdk.events import Event from trpc_agent_sdk.sessions._session import Session -from trpc_agent_sdk.sessions._session_summarizer import ( +from trpc_agent_sdk.sessions.compact.default._summarizer import ( DEFAULT_SUMMARIZER_PROMPT, - SessionSummarizer, - SessionSummary, + DefaultSessionSummarizer as SessionSummarizer, + DefaultSessionSummary as SessionSummary, ) from trpc_agent_sdk.types import Content, EventActions, FunctionCall, FunctionResponse, Part diff --git a/tests/sessions/test_sql_session_service.py b/tests/sessions/test_sql_session_service.py index e1730ec6c..8a0d7c136 100644 --- a/tests/sessions/test_sql_session_service.py +++ b/tests/sessions/test_sql_session_service.py @@ -428,6 +428,50 @@ async def test_update_existing(self): assert len(stored.events) == 0 await svc.close() + async def test_update_persists_compacted_active_and_historical_events(self): + svc = await _create_service(_make_config(store_historical_events=True)) + session = await svc.create_session(app_name="app", user_id="user", session_id="s1") + original = [_make_event(text=f"msg{i}") for i in range(4)] + for event in original: + await svc.append_event(session, event) + summary = _make_event(author="system", text="summary") + + assert session.compact_events( + summary, + original[1].id, + compaction_id="compact-1", + ) + await svc.update_session(session) + + stored = await svc.get_session(app_name="app", user_id="user", session_id="s1") + assert stored is not None + assert stored.events[0].is_summary_event() + assert [event.id for event in stored.events[1:]] == [event.id for event in original[2:]] + assert [event.id for event in stored.historical_events] == [event.id for event in original[:2]] + await svc.close() + + async def test_patch_state_preserves_stored_events(self): + svc = await _create_service() + session = await svc.create_session( + app_name="app", + user_id="user", + session_id="s1", + ) + await svc.append_event(session, _make_event(text="keep me")) + + stale = session.model_copy(deep=True) + stale.events = [] + await svc.update_session_state(stale, {"_trpc_agent:summary": {"v": 1}}) + + stored = await svc.get_session( + app_name="app", + user_id="user", + session_id="s1", + ) + assert [event.content.parts[0].text for event in stored.events] == ["keep me"] + assert stored.state["_trpc_agent:summary"] == {"v": 1} + await svc.close() + async def test_update_nonexistent(self): svc = await _create_service() session = Session(id="nonexistent", app_name="app", user_id="user", save_key="k") diff --git a/tests/sessions/test_summarizer_checker.py b/tests/sessions/test_summarizer_checker.py index 62613ed40..8d12360ef 100644 --- a/tests/sessions/test_summarizer_checker.py +++ b/tests/sessions/test_summarizer_checker.py @@ -24,7 +24,7 @@ from trpc_agent_sdk.events import Event from trpc_agent_sdk.sessions._session import Session -from trpc_agent_sdk.sessions._summarizer_checker import ( +from trpc_agent_sdk.sessions.compact.default._checker import ( set_summarizer_check_functions_by_and, set_summarizer_check_functions_by_or, set_summarizer_conversation_threshold, @@ -284,7 +284,7 @@ def test_above_threshold(self): checker = set_summarizer_conversation_threshold(10) session = _make_session(conversation_count=15) assert checker(session) is True - assert session.conversation_count == 0 + assert session.conversation_count == 15 def test_below_threshold(self): checker = set_summarizer_conversation_threshold(10) @@ -301,12 +301,12 @@ def test_default_threshold(self): session = _make_session(conversation_count=101) assert checker(session) is True - def test_resets_count_on_true(self): + def test_does_not_mutate_count(self): checker = set_summarizer_conversation_threshold(5) session = _make_session(conversation_count=10) result = checker(session) assert result is True - assert session.conversation_count == 0 + assert session.conversation_count == 10 class TestCheckFunctionsByAnd: diff --git a/tests/sessions/test_summarizer_manager.py b/tests/sessions/test_summarizer_manager.py index fa3292983..0086f5aed 100644 --- a/tests/sessions/test_summarizer_manager.py +++ b/tests/sessions/test_summarizer_manager.py @@ -20,8 +20,11 @@ from trpc_agent_sdk.events import Event from trpc_agent_sdk.sessions._session import Session -from trpc_agent_sdk.sessions._session_summarizer import SessionSummarizer, SessionSummary -from trpc_agent_sdk.sessions._summarizer_manager import SummarizerSessionManager +from trpc_agent_sdk.sessions.compact.default._summarizer import DefaultSessionSummarizer as SessionSummarizer +from trpc_agent_sdk.sessions.compact.default._summarizer import DefaultSessionSummary as SessionSummary +from trpc_agent_sdk.sessions.compact.default._summarizer_manager import ( + DefaultSessionSummarizerManager as SummarizerSessionManager, +) from trpc_agent_sdk.types import Content, Part @@ -139,11 +142,12 @@ async def test_summary_when_should_summarize(self): mock_service = AsyncMock() manager.set_session_service(mock_service) - session = _make_session(events=[_make_event()]) + session = _make_session(events=[_make_event()], conversation_count=15) await manager.create_session_summary(session) manager._summarizer.create_session_summary.assert_called_once() mock_service.update_session.assert_called_once() + assert session.conversation_count == 0 async def test_no_summary_when_should_not_summarize(self): model = _make_model() diff --git a/trpc_agent_sdk/abc/__init__.py b/trpc_agent_sdk/abc/__init__.py index 5a3f75125..f9e3cf650 100644 --- a/trpc_agent_sdk/abc/__init__.py +++ b/trpc_agent_sdk/abc/__init__.py @@ -16,6 +16,9 @@ from ._artifact_service import ArtifactId from ._artifact_service import ArtifactServiceABC from ._artifact_service import ArtifactVersion +from ._compact import CompactSummarizerABC +from ._compact import CompactSummarizerManagerABC +from ._compact import CompactTrigger from ._filter import FilterABC from ._filter import FilterAsyncGenHandleType from ._filter import FilterAsyncGenReturnType @@ -43,6 +46,9 @@ "ArtifactId", "ArtifactServiceABC", "ArtifactVersion", + "CompactSummarizerABC", + "CompactSummarizerManagerABC", + "CompactTrigger", "FilterABC", "FilterAsyncGenHandleType", "FilterAsyncGenReturnType", diff --git a/trpc_agent_sdk/abc/_compact.py b/trpc_agent_sdk/abc/_compact.py new file mode 100644 index 000000000..b96db4808 --- /dev/null +++ b/trpc_agent_sdk/abc/_compact.py @@ -0,0 +1,187 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""The base class for compact summarizers.""" + +from abc import ABC +from abc import abstractmethod +from enum import Enum +from typing import List +from typing import Optional +from typing import Dict +from typing import Any +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from trpc_agent_sdk.context import InvocationContext + +from ._session import SessionABC +from ._response import ResponseABC +from ._request import RequestABC +from ._session_service import SessionServiceABC + + +class CompactTrigger(str, Enum): + """Select when session compaction is evaluated.""" + + AFTER_TURN = "after_turn" + BEFORE_MODEL = "before_model" + + +class CompactSummarizerABC(ABC): + """The base class for compact summarizers.""" + + @abstractmethod + async def should_summarize(self, session: SessionABC) -> bool: + """Check if the session should be summarized. + + Args: + session: The session to check. + + Returns: + True if the session should be summarized, False otherwise. + """ + + @abstractmethod + async def create_session_summary_by_events( + self, + events: List[ResponseABC], + session_id: str, + keep_recent_count: int = 10, + ctx: Optional["InvocationContext"] = None, + historical_events: Optional[List[ResponseABC]] = None, + store_historical_events: bool = False) -> tuple[Optional[str], List[ResponseABC]]: + """Create a session summary by events. + + Args: + events: The events to summarize. + session_id: The session ID. + keep_recent_count: The number of recent events to keep. + ctx: The invocation context. + historical_events: The historical events. + store_historical_events: Whether to store the historical events. + + Returns: + A tuple containing the session summary and the historical events. + """ + + @abstractmethod + async def create_session_summary(self, + session: SessionABC, + ctx: Optional["InvocationContext"] = None, + store_historical_events: bool = False) -> Optional[str]: + """Create a session summary. + + Args: + session: The session to summarize. + ctx: The invocation context. + store_historical_events: Whether to store the historical events. + + Returns: + The session summary. + """ + + def get_summary_metadata(self) -> Dict[str, Any]: + """Get the summary metadata. + + Returns: + The summary metadata. + """ + return {} + + async def create_session_summary_by_request( + self, + request: RequestABC, + ctx: Optional["InvocationContext"] = None, + force: bool = False, + ) -> Optional[ResponseABC]: + """Compact one model request before generation. + + The default implementation is intentionally a no-op so existing + summarizers only implementing end-of-turn compaction remain + compatible. + """ + del request, ctx, force + return None + + +class CompactSummarizerManagerABC(ABC): + """Coordinate one CompactSummarizer implementation with a SessionService.""" + + def __init__( + self, + summarizer: CompactSummarizerABC, + compact_trigger: CompactTrigger = CompactTrigger.AFTER_TURN, + ): + self._summarizer = summarizer + self._base_service = None + self._compact_trigger = compact_trigger + + @property + def summarizer(self) -> CompactSummarizerABC: + """Get the CompactSummarizer implementation.""" + return self._summarizer + + @property + def session_service(self) -> SessionServiceABC: + """Get the base session service.""" + return self._base_service + + @property + def compact_trigger(self) -> CompactTrigger: + """Return when this manager evaluates compaction.""" + return self._compact_trigger + + def set_session_service(self, session_service: SessionServiceABC, force: bool = False) -> None: + """Set the session service to use. + + Args: + session_service: The session service to use. + force: Whether to force update even if already set. + """ + if not self._base_service or force: + self._base_service = session_service + + def set_summarizer(self, summarizer: CompactSummarizerABC, force: bool = False) -> None: + """Set the summarizer to use. + + Args: + summarizer: The summarizer to use + force: Whether to force update even if already set + """ + if not self._summarizer or force: + self._summarizer = summarizer + + @abstractmethod + async def create_session_summary( + self, + session: SessionABC, + force: bool = False, + ctx: Optional["InvocationContext"] = None, + ) -> None: + """Update compact state through the SessionService post-turn hook.""" + + @abstractmethod + async def get_session_summary(self, session: SessionABC) -> Optional[str]: + """Return the compact representation exposed as a session summary.""" + + async def create_session_summary_before_model( + self, + request: RequestABC, + ctx: "InvocationContext", + force: bool = False, + ) -> Optional[ResponseABC]: + """Run request compaction when configured for the before-model phase.""" + if self._compact_trigger != CompactTrigger.BEFORE_MODEL: + return None + return await self._summarizer.create_session_summary_by_request( + request, + ctx=ctx, + force=force, + ) + + async def close(self) -> None: + """Release resources owned by this manager.""" + return None diff --git a/trpc_agent_sdk/abc/_session_service.py b/trpc_agent_sdk/abc/_session_service.py index 419fac067..f1cbd1923 100644 --- a/trpc_agent_sdk/abc/_session_service.py +++ b/trpc_agent_sdk/abc/_session_service.py @@ -125,6 +125,20 @@ async def update_session(self, session: SessionABC) -> None: session: The session to update """ + async def update_session_state( + self, + session: SessionABC, + state_delta: dict[str, Any], + ) -> None: + """Persist a session-scoped state delta. + + Backends may override this method with an efficient partial update. + The default implementation preserves compatibility with existing + SessionService implementations by falling back to ``update_session``. + """ + session.state.update(state_delta) + await self.update_session(session) + @abstractmethod async def create_session_summary(self, session: SessionABC, ctx: "InvocationContext" = None) -> None: """Summarize a session.""" diff --git a/trpc_agent_sdk/advanced_memory/__init__.py b/trpc_agent_sdk/advanced_memory/__init__.py deleted file mode 100644 index c346cd2d3..000000000 --- a/trpc_agent_sdk/advanced_memory/__init__.py +++ /dev/null @@ -1,132 +0,0 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. -# -# Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Optional Advanced Memory module that leaves the legacy mechanism unchanged.""" - -from ._autocompact import AutoCompact -from ._autocompact import AutoCompactCallback -from ._autocompact import AutoCompactResult -from ._autocompact import content_signature -from ._autocompact import ForkedLegacySummaryGenerator -from ._autocompact import setup_autocompact -from ._config import AdvancedMemoryConfig -from ._formats import MemoryDocument -from ._formats import MemoryIndexEntry -from ._formats import MemoryType -from ._formats import memory_freshness -from ._formats import parse_memory_updated_at -from ._formats import SESSION_MEMORY_SECTION_DESCRIPTIONS -from ._formats import SESSION_MEMORY_SECTIONS -from ._formats import SessionMemoryDocument -from ._history_snip import estimate_request_chars -from ._history_snip import HistorySnip -from ._history_snip import HistorySnipCallback -from ._history_snip import HistorySnipResult -from ._history_snip import setup_history_snip -from ._memory_context import LongTermMemoryContext -from ._memory_context import LongTermMemoryContextCallback -from ._memory_context import setup_long_term_memory_context -from ._microcompact import Microcompact -from ._microcompact import MicrocompactCallback -from ._microcompact import MicrocompactResult -from ._microcompact import setup_microcompact -from ._paths import AdvancedMemoryPaths -from ._preload_memory import MemoryCandidate -from ._preload_memory import MemoryPreloader -from ._preload_memory import MemoryRelevanceSelector -from ._preload_memory import ModelMemoryRelevanceSelector -from ._preload_memory import select_relevant_memory_filenames -from ._runtime import AdvancedMemoryRuntime -from ._session_memory import build_session_memory_prompt -from ._session_memory import ForkedSessionMemoryGenerator -from ._session_memory import has_session_memory_content -from ._session_memory import limit_session_memory_document -from ._session_memory import SessionMemoryExtractionInput -from ._session_memory import SessionMemoryExtractionResult -from ._session_memory import SessionMemoryExtractor -from ._session_service import TranscriptSessionService -from ._integration import AdvancedContextManagement -from ._integration import AdvancedMemoryIntegration -from ._integration import setup_advanced_memory -from ._integration import setup_context_management -from ._storage import LongTermMemoryStore -from ._storage import SessionMemoryStore -from ._storage import ToolResultStore -from ._storage import TranscriptStore -from ._tool_result_budget import setup_tool_result_budget -from ._tool_result_budget import ToolResultBudget -from ._tool_result_budget import ToolResultBudgetCallback -from ._tool_result_budget import ToolResultBudgetResult -from ._transcript import TRANSCRIPT_SCHEMA_VERSION -from ._token_budget import ContextBudget -from ._token_budget import ContextTokenEstimate -from ._token_budget import HeuristicTokenEstimator -from ._token_budget import ModelContextWindowResolver -from ._token_budget import TokenContextTracker -from ._token_budget import TokenEstimator - -__all__ = [ - "AutoCompact", - "AutoCompactCallback", - "AutoCompactResult", - "AdvancedMemoryConfig", - "AdvancedContextManagement", - "AdvancedMemoryIntegration", - "AdvancedMemoryPaths", - "AdvancedMemoryRuntime", - "ContextBudget", - "ContextTokenEstimate", - "build_session_memory_prompt", - "content_signature", - "estimate_request_chars", - "ForkedLegacySummaryGenerator", - "ForkedSessionMemoryGenerator", - "has_session_memory_content", - "HistorySnip", - "HistorySnipCallback", - "HistorySnipResult", - "HeuristicTokenEstimator", - "LongTermMemoryStore", - "LongTermMemoryContext", - "LongTermMemoryContextCallback", - "MemoryDocument", - "MemoryIndexEntry", - "MemoryType", - "MemoryCandidate", - "MemoryPreloader", - "MemoryRelevanceSelector", - "ModelMemoryRelevanceSelector", - "select_relevant_memory_filenames", - "memory_freshness", - "parse_memory_updated_at", - "Microcompact", - "MicrocompactCallback", - "MicrocompactResult", - "ModelContextWindowResolver", - "SESSION_MEMORY_SECTION_DESCRIPTIONS", - "SESSION_MEMORY_SECTIONS", - "SessionMemoryDocument", - "SessionMemoryExtractionInput", - "SessionMemoryExtractionResult", - "SessionMemoryExtractor", - "SessionMemoryStore", - "TRANSCRIPT_SCHEMA_VERSION", - "ToolResultBudget", - "ToolResultBudgetCallback", - "ToolResultBudgetResult", - "ToolResultStore", - "TokenContextTracker", - "TokenEstimator", - "TranscriptSessionService", - "TranscriptStore", - "setup_autocompact", - "setup_advanced_memory", - "limit_session_memory_document", - "setup_history_snip", - "setup_context_management", - "setup_long_term_memory_context", - "setup_microcompact", - "setup_tool_result_budget", -] diff --git a/trpc_agent_sdk/advanced_memory/_autocompact.py b/trpc_agent_sdk/advanced_memory/_autocompact.py deleted file mode 100644 index e7efb436a..000000000 --- a/trpc_agent_sdk/advanced_memory/_autocompact.py +++ /dev/null @@ -1,766 +0,0 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. -# -# Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Automatically compact history before model requests and circuit-break failures.""" - -from __future__ import annotations - -import asyncio -import hashlib -import json -import re -import uuid -from dataclasses import dataclass -from typing import Any -from typing import Protocol -from typing import TYPE_CHECKING - -from trpc_agent_sdk.agents import LlmAgent -from trpc_agent_sdk.models import LlmResponse -from trpc_agent_sdk.runners import Runner -from trpc_agent_sdk.sessions import InMemorySessionService -from trpc_agent_sdk.types import Content -from trpc_agent_sdk.types import Part - -from ._callbacks import install_staged_callback -from ._formats import SESSION_MEMORY_SECTIONS -from ._formats import SessionMemoryDocument -from ._history_snip import estimate_request_chars -from ._runtime import AdvancedMemoryRuntime -from ._token_budget import TokenContextTracker - -if TYPE_CHECKING: - from trpc_agent_sdk.agents import LlmAgent as ParentLlmAgent - from trpc_agent_sdk.context import InvocationContext - from trpc_agent_sdk.models import LlmRequest - -AUTOCOMPACT_SCHEMA_VERSION = 1 -AUTOCOMPACT_BLOCKED_MESSAGE = ( - "Automatic context compaction has failed repeatedly and the request is near the hard context limit. " - "To avoid sending a request that will certainly fail, reduce the input, start a new session, " - "or manually organize session memory before retrying.") -AUTOCOMPACT_SUMMARY_PREFIX = """This session is being continued from a compacted context. -The following summary contains the important information from earlier messages. -The complete original events remain available in the session transcript. - -""" -_LEGACY_SESSION_MEMORY_SECTION_LIST = "\n".join(f"- # {section}" for section in SESSION_MEMORY_SECTIONS) - -LEGACY_SUMMARY_INSTRUCTION = """You are an isolated context-compaction Agent. -Compress the provided old conversation into a dense Markdown summary that another Agent can continue seamlessly. -Preserve the user's goals, explicit requirements, key technical decisions, files and functions, commands, -errors and fixes, verified results, current state, and next steps. -Do not answer questions from the old conversation, mention this compaction prompt, or invent information. -Return exactly two XML blocks: first use ... to check coverage, then -... for the final Markdown summary. The summary must contain these ten Markdown sections -in this order: -""" + _LEGACY_SESSION_MEMORY_SECTION_LIST + """ -The analysis is only for organization; keep only the summary.""" - - -@dataclass(frozen=True) -class AutoCompactRecord: - """Store stable replay information for the latest successful compaction.""" - - boundary_signature: str - boundary_occurrence: int - summary: str - source: str - - -@dataclass -class AutoCompactState: - """Store the latest compaction record and consecutive failure count.""" - - latest_compaction: AutoCompactRecord | None - consecutive_failures: int - - -@dataclass(frozen=True) -class AutoCompactResult: - """Summarize one compaction, replay, or hard-block result.""" - - compacted: bool - reapplied: bool - blocked: bool - source: str | None - request_chars_before: int - request_chars_after: int - consecutive_failures: int - error: str | None = None - request_tokens_before: int | None = None - request_tokens_after: int | None = None - token_source: str | None = None - - -class LegacySummaryGenerator(Protocol): - """Define the replaceable legacy compaction summary interface.""" - - async def generate(self, history: str, ctx: "InvocationContext") -> str: - """Return a workable Markdown summary for bounded old history.""" - - -def content_signature(content: Content) -> str: - """Generate a stable signature that preserves message identity.""" - parts: list[dict[str, Any]] = [] - for part in content.parts or []: - if part.text is not None: - parts.append({ - "type": "text", - "sha256": hashlib.sha256(part.text.encode("utf-8")).hexdigest(), - }) - elif part.function_call is not None: - parts.append({ - "type": "function_call", - "id": getattr(part.function_call, "id", None), - "name": part.function_call.name, - }) - elif part.function_response is not None: - parts.append({ - "type": "function_response", - "id": getattr(part.function_response, "id", None), - "name": part.function_response.name, - }) - elif part.executable_code is not None: - parts.append({"type": "executable_code"}) - elif part.code_execution_result is not None: - parts.append({"type": "code_execution_result"}) - else: - parts.append({"type": "other"}) - serialized = json.dumps( - { - "role": content.role, - "parts": parts - }, - ensure_ascii=False, - sort_keys=True, - separators=(",", ":"), - ) - return hashlib.sha256(serialized.encode("utf-8")).hexdigest() - - -def _content_text(content: Content) -> str: - """Render one model content item as legacy summary input.""" - return json.dumps( - content.model_dump(mode="json", by_alias=True, exclude_none=True), - ensure_ascii=False, - sort_keys=True, - separators=(",", ":"), - default=str, - ) - - -class ForkedLegacySummaryGenerator: - """Call a tool-free legacy summary Agent through an isolated Runner.""" - - def __init__(self, model: Any | None = None) -> None: - """Store an optional dedicated model, falling back to the parent model.""" - self._model = model - - def _resolve_model(self, ctx: "InvocationContext") -> Any: - """Resolve the model used for legacy compaction.""" - model = self._model or getattr(ctx.agent, "model", None) - if not model: - raise ValueError("Autocompact summary generator cannot resolve an LLM model") - return model - - async def generate(self, history: str, ctx: "InvocationContext") -> str: - """Generate a summary in a temporary session without parent callbacks.""" - config = ctx.agent.generate_content_config if isinstance(ctx.agent, LlmAgent) else None - agent = LlmAgent( - name="advanced_autocompact_summarizer", - description="Generate an isolated context-compaction summary.", - instruction=LEGACY_SUMMARY_INSTRUCTION, - model=self._resolve_model(ctx), - tools=[], - generate_content_config=config, - add_name_to_instruction=False, - ) - app_name = f"{ctx.app_name}_advanced_autocompact" - runner = Runner( - app_name=app_name, - agent=agent, - session_service=InMemorySessionService(), - enable_post_turn_processing=False, - ) - last_event = None - try: - session = await runner.session_service.create_session( - app_name=app_name, - user_id="advanced-autocompact", - state={}, - ) - prompt = ("Compress the following old conversation. The input may contain JSON representations " - "of tool calls and results:\n\n" - f"\n{history}\n") - async for event in runner.run_async( - user_id=session.user_id, - session_id=session.id, - new_message=Content(role="user", parts=[Part.from_text(text=prompt)]), - ): - if not event.partial: - last_event = event - finally: - await runner.close() - if not last_event or not last_event.content or not last_event.content.parts: - raise ValueError("Autocompact summary generator returned no final content") - output = "\n".join(part.text for part in last_event.content.parts if part.text).strip() - summary_match = re.search( - r"\s*(.*?)\s*", - output, - flags=re.DOTALL | re.IGNORECASE, - ) - if summary_match is None or not summary_match.group(1).strip(): - raise ValueError("Autocompact summary generator returned no block") - return summary_match.group(1).strip() - - -class AutoCompact: - """Compact with session memory first, then fall back to a legacy summary.""" - - def __init__( - self, - memory_runtime: AdvancedMemoryRuntime, - summary_generator: LegacySummaryGenerator | None = None, - *, - model: Any | None = None, - ) -> None: - """Initialize the compressor, summary generator, and session locks.""" - if summary_generator is not None and model is not None: - raise ValueError("Provide either summary_generator or model, not both") - self._runtime = memory_runtime - self._summary_generator = summary_generator or ForkedLegacySummaryGenerator(model) - self._states: dict[str, AutoCompactState] = {} - self._session_locks: dict[str, asyncio.Lock] = {} - - @property - def runtime(self) -> AdvancedMemoryRuntime: - """Return the runtime bound to this compressor.""" - return self._runtime - - def _session_lock(self, session_id: str) -> asyncio.Lock: - """Return the unique compaction lock for a session.""" - lock = self._session_locks.get(session_id) - if lock is None: - lock = asyncio.Lock() - self._session_locks[session_id] = lock - return lock - - async def _load_state(self, session_id: str) -> AutoCompactState: - """Restore the latest compaction and failure count from the transcript.""" - state = self._states.get(session_id) - if state is not None: - return state - records = await self._runtime.transcripts.read_all(session_id) - latest: AutoCompactRecord | None = None - failures = 0 - for record in records: - if record.get("kind") == "autocompact-success": - signature = record.get("boundary_signature") - occurrence = record.get("boundary_occurrence") - summary = record.get("summary") - source = record.get("source") - if (all(isinstance(value, str) for value in (signature, summary, source)) - and isinstance(occurrence, int) and occurrence > 0): - latest = AutoCompactRecord( - signature, - occurrence, - summary, - source, - ) - failures = 0 - elif record.get("kind") == "autocompact-failure": - failures += 1 - state = AutoCompactState(latest_compaction=latest, consecutive_failures=failures) - self._states[session_id] = state - return state - - def _summary_content(self, summary: str) -> Content: - """Wrap a compaction summary in stable model-visible user content.""" - return Content( - role="user", - parts=[Part.from_text(text=AUTOCOMPACT_SUMMARY_PREFIX + summary)], - ) - - def _summary_with_recovery_path(self, summary: str, session_id: str) -> str: - """Append recovery paths for the full transcript and session memory.""" - return (f"{summary.rstrip()}\n\n" - "For exact content from before compaction, read the complete transcript: " - f"{self._runtime.paths.transcript_path(session_id)}\n" - "Current session memory: " - f"{self._runtime.paths.session_memory_path(session_id)}") - - def _find_signature_index( - self, - contents: list[Content], - signature: str, - occurrence: int, - ) -> int | None: - """Locate a persisted compaction boundary by signature occurrence.""" - seen = 0 - for index, content in enumerate(contents): - if content_signature(contents[index]) == signature: - seen += 1 - if seen == occurrence: - return index - return None - - def _signature_occurrence( - self, - contents: list[Content], - signature: str, - boundary_index: int, - ) -> int: - """Count a boundary signature's occurrences from the request start.""" - return sum(1 for content in contents[:boundary_index + 1] if content_signature(content) == signature) - - def _adjust_start_for_tool_pairing(self, contents: list[Content], start: int) -> int: - """Extend the retained range to keep calls paired with responses.""" - if start <= 0 or start >= len(contents): - return max(0, start) - response_ids = { - getattr(part.function_response, "id", None) - for content in contents[start:] - for part in content.parts or [] if part.function_response is not None - } - response_ids.discard(None) - if not response_ids: - return start - for index in range(start - 1, -1, -1): - call_ids = { - getattr(part.function_call, "id", None) - for part in contents[index].parts or [] if part.function_call is not None - } - if call_ids & response_ids: - start = index - response_ids -= call_ids - if not response_ids: - break - return start - - def _compaction_start(self, contents: list[Content], boundary_index: int) -> int: - """Return the retained-content start for a legacy compaction.""" - start = min( - boundary_index + 1, - len(contents) - self._runtime.config.autocompact_keep_recent_contents, - ) - return self._adjust_start_for_tool_pairing(contents, start) - - def _session_memory_compaction_start(self, boundary_index: int) -> int: - """Drop everything through the session-memory checkpoint boundary.""" - return boundary_index + 1 - - def _apply_record(self, request: "LlmRequest", record: AutoCompactRecord) -> bool: - """Replay a persisted compaction record into a rebuilt request.""" - boundary_index = self._find_signature_index( - request.contents, - record.boundary_signature, - record.boundary_occurrence, - ) - if boundary_index is None: - return False - start = (self._session_memory_compaction_start(boundary_index) - if record.source == "session-memory" else self._compaction_start(request.contents, boundary_index)) - request.contents = [ - self._summary_content(record.summary), - *request.contents[start:], - ] - return True - - async def _latest_session_memory_record( - self, - session_id: str, - ) -> tuple[str, str] | None: - """Read session memory and its checkpoint Event for model-free compaction.""" - async with self._runtime.coordination.guard( - session_id, - timeout=self._runtime.config.session_memory_wait_timeout_seconds, - ) as acquired: - if not acquired: - return None - memory = await self._runtime.session_memory.read(session_id) - if memory is None or memory.strip() == SessionMemoryDocument().to_markdown().strip(): - return None - records = await self._runtime.transcripts.read_all(session_id) - for record in reversed(records): - if record.get("kind") == "session-memory-checkpoint" and isinstance(record.get("last_event_id"), str): - return memory, record["last_event_id"] - return None - - def _event_content_signature( - self, - records: list[dict[str, Any]], - event_id: str, - ) -> tuple[str, int] | None: - """Recover a boundary signature and occurrence from transcript Events.""" - signatures: list[str] = [] - for record in records: - if record.get("kind") != "event": - continue - raw_content = record.get("event", {}).get("content") - if not isinstance(raw_content, dict): - continue - try: - signature = content_signature(Content.model_validate(raw_content)) - except Exception: # noqa: BLE001 - return None - signatures.append(signature) - if record.get("event_id") == event_id: - return signature, signatures.count(signature) - return None - - def _compact_with_summary( - self, - request: "LlmRequest", - *, - summary: str, - boundary_index: int, - source: str, - strict_boundary: bool = False, - ) -> AutoCompactRecord: - """Replace the old prefix with a summary and return a replay record.""" - boundary_signature = content_signature(request.contents[boundary_index]) - boundary_occurrence = self._signature_occurrence( - request.contents, - boundary_signature, - boundary_index, - ) - start = (self._session_memory_compaction_start(boundary_index) if strict_boundary else self._compaction_start( - request.contents, boundary_index)) - request.contents = [self._summary_content(summary), *request.contents[start:]] - return AutoCompactRecord( - boundary_signature, - boundary_occurrence, - summary, - source, - ) - - def _bounded_history(self, contents: list[Content]) -> str: - """Bound old history to the configured summary-input character limit.""" - rendered = "\n".join(f"\n{_content_text(content)}\n" for content in contents) - limit = self._runtime.config.autocompact_summary_input_max_chars - if len(rendered) <= limit: - return rendered - marker = "\n...[middle of old history omitted due to the summary input limit]...\n" - first_size = max(1, (limit - len(marker)) // 3) - last_size = max(1, limit - len(marker) - first_size) - return rendered[:first_size] + marker + rendered[-last_size:] - - async def _legacy_summary( - self, - contents: list[Content], - ctx: "InvocationContext", - ) -> str: - """Shrink old history across retries and generate a legacy summary.""" - retries = self._runtime.config.autocompact_summary_retries - working = list(contents) - last_error: Exception | None = None - for attempt in range(retries): - try: - return await self._summary_generator.generate( - self._bounded_history(working), - ctx, - ) - except Exception as exc: # noqa: BLE001 - last_error = exc - if len(working) <= 1: - break - drop_count = max(1, len(working) // (retries - attempt + 1)) - working = working[drop_count:] - raise RuntimeError("Legacy autocompact summary failed after retries") from last_error - - async def _persist_success( - self, - session_id: str, - record: AutoCompactRecord, - before_chars: int, - after_chars: int, - before_tokens: int | None = None, - after_tokens: int | None = None, - token_source: str | None = None, - ) -> None: - """Persist a successful compaction and reset the circuit-breaker count.""" - await self._runtime.transcripts.append( - session_id, - { - "schema_version": AUTOCOMPACT_SCHEMA_VERSION, - "kind": "autocompact-success", - "compaction_id": f"autocompact:{uuid.uuid4().hex}", - "boundary_signature": record.boundary_signature, - "boundary_occurrence": record.boundary_occurrence, - "summary": record.summary, - "source": record.source, - "request_chars_before": before_chars, - "request_chars_after": after_chars, - "request_tokens_before": before_tokens, - "request_tokens_after": after_tokens, - "token_source": token_source, - }, - ) - - async def _persist_failure( - self, - session_id: str, - error: Exception, - failures: int, - token_budget: Any | None = None, - ) -> None: - """Persist failures so the circuit breaker survives a restart.""" - await self._runtime.transcripts.append( - session_id, - { - "schema_version": - AUTOCOMPACT_SCHEMA_VERSION, - "kind": - "autocompact-failure", - "attempt_id": - f"autocompact:{uuid.uuid4().hex}", - "consecutive_failures": - failures, - "error": - str(error), - "request_tokens": (token_budget.estimate.tokens - if token_budget is not None and token_budget.token_mode_enabled else None), - "context_window_tokens": (token_budget.context_window_tokens - if token_budget is not None and token_budget.token_mode_enabled else None), - "token_source": (token_budget.estimate.source - if token_budget is not None and token_budget.token_mode_enabled else None), - }, - ) - - async def apply( - self, - request: "LlmRequest", - *, - session_id: str, - ctx: "InvocationContext", - force: bool = False, - ) -> AutoCompactResult: - """Replay old compaction and compact again when pressure is high.""" - config = self._runtime.config - tracker = TokenContextTracker(config) - if not config.enabled or not config.autocompact_enabled: - request_chars = estimate_request_chars(request) - return AutoCompactResult(False, False, False, None, request_chars, request_chars, 0) - await self._runtime.initialize() - async with self._session_lock(session_id): - request.contents = [content.model_copy(deep=True) for content in request.contents] - state = await self._load_state(session_id) - reapplied = False - if state.latest_compaction is not None: - reapplied = self._apply_record(request, state.latest_compaction) - - request_chars_before = estimate_request_chars(request) - token_budget_before = tracker.budget(request, ctx) - token_mode = token_budget_before.token_mode_enabled - request_tokens_before = token_budget_before.estimate.tokens - blocking_reached = (request_tokens_before >= token_budget_before.blocking_threshold_tokens - if token_mode else request_chars_before >= config.autocompact_blocking_chars) - if state.consecutive_failures >= config.autocompact_max_failures and blocking_reached: - return AutoCompactResult( - False, - reapplied, - True, - None, - request_chars_before, - request_chars_before, - state.consecutive_failures, - ) - if state.consecutive_failures >= config.autocompact_max_failures: - return AutoCompactResult( - False, - reapplied, - False, - None, - request_chars_before, - request_chars_before, - state.consecutive_failures, - ) - autocompact_reached = (request_tokens_before >= token_budget_before.autocompact_threshold_tokens - if token_mode else request_chars_before >= config.autocompact_trigger_chars) - if not force and not autocompact_reached: - return AutoCompactResult( - False, - reapplied, - False, - state.latest_compaction.source if reapplied and state.latest_compaction else None, - request_chars_before, - request_chars_before, - state.consecutive_failures, - ) - - original_contents = [content.model_copy(deep=True) for content in request.contents] - try: - compact_record: AutoCompactRecord | None = None - session_memory = await self._latest_session_memory_record(session_id) - if session_memory is not None: - memory, checkpoint_event_id = session_memory - transcript_records = await self._runtime.transcripts.read_all(session_id) - boundary = self._event_content_signature( - transcript_records, - checkpoint_event_id, - ) - if boundary is not None: - boundary_signature, boundary_occurrence = boundary - boundary_index = self._find_signature_index( - request.contents, - boundary_signature, - boundary_occurrence, - ) - if boundary_index is not None: - compact_record = self._compact_with_summary( - request, - summary=self._summary_with_recovery_path( - memory, - session_id, - ), - boundary_index=boundary_index, - source="session-memory", - strict_boundary=True, - ) - target_reached = (tracker.budget(request, ctx).estimate.tokens - <= token_budget_before.warning_threshold_tokens if token_mode else - estimate_request_chars(request) <= config.autocompact_target_chars) - if not target_reached: - request.contents = [content.model_copy(deep=True) for content in original_contents] - compact_record = None - - if compact_record is None: - keep_count = min( - config.autocompact_keep_recent_contents, - max(1, - len(request.contents) - 1), - ) - boundary_index = len(request.contents) - keep_count - 1 - if boundary_index < 0: - raise ValueError("Not enough model contents to compact") - summary = await self._legacy_summary( - request.contents[:boundary_index + 1], - ctx, - ) - compact_record = self._compact_with_summary( - request, - summary=self._summary_with_recovery_path( - summary, - session_id, - ), - boundary_index=boundary_index, - source="legacy", - ) - - request_chars_after = estimate_request_chars(request) - if request_chars_after >= request_chars_before: - raise ValueError("Autocompact did not reduce request size") - token_budget_after = tracker.budget(request, ctx) - if token_mode and token_budget_after.estimate.tokens >= request_tokens_before: - raise ValueError("Autocompact did not reduce request token estimate") - await self._persist_success( - session_id, - compact_record, - request_chars_before, - request_chars_after, - request_tokens_before if token_mode else None, - token_budget_after.estimate.tokens if token_mode else None, - token_budget_after.estimate.source if token_mode else None, - ) - state.latest_compaction = compact_record - state.consecutive_failures = 0 - return AutoCompactResult( - True, - reapplied, - False, - compact_record.source, - request_chars_before, - request_chars_after, - 0, - request_tokens_before=request_tokens_before if token_mode else None, - request_tokens_after=(token_budget_after.estimate.tokens if token_mode else None), - token_source=token_budget_after.estimate.source if token_mode else None, - ) - except Exception as exc: # noqa: BLE001 - request.contents = original_contents - state.consecutive_failures += 1 - await self._persist_failure( - session_id, - exc, - state.consecutive_failures, - token_budget_before, - ) - blocked = state.consecutive_failures >= config.autocompact_max_failures and blocking_reached - return AutoCompactResult( - False, - reapplied, - blocked, - None, - request_chars_before, - request_chars_before, - state.consecutive_failures, - error=str(exc), - request_tokens_before=request_tokens_before if token_mode else None, - request_tokens_after=request_tokens_before if token_mode else None, - token_source=token_budget_before.estimate.source if token_mode else None, - ) - - -class AutoCompactCallback: - """Adapt the automatic compressor to before_model_callback.""" - - advanced_memory_stage = 40 - - def __init__(self, autocompact: AutoCompact) -> None: - """Store the compressor executed before model requests.""" - self._autocompact = autocompact - - @property - def autocompact(self) -> AutoCompact: - """Return the compressor used by this callback.""" - return self._autocompact - - async def __call__( - self, - ctx: "InvocationContext", - request: "LlmRequest", - ) -> LlmResponse | None: - """Compact before each request and return a local block after failures.""" - result = await self._autocompact.apply( - request, - session_id=ctx.session_id, - ctx=ctx, - ) - if not result.blocked: - TokenContextTracker(self._autocompact.runtime.config).record_request_context( - request, - ctx, - ) - return None - return LlmResponse(content=Content( - role="model", - parts=[Part.from_text(text=AUTOCOMPACT_BLOCKED_MESSAGE)], - )) - - -def setup_autocompact( - agent: "ParentLlmAgent", - memory_runtime: AdvancedMemoryRuntime, - summary_generator: LegacySummaryGenerator | None = None, - *, - model: Any | None = None, -) -> AutoCompact: - """Install the automatic compaction callback in pipeline stage order.""" - autocompact = AutoCompact( - memory_runtime, - summary_generator, - model=model, - ) - callback = AutoCompactCallback(autocompact) - existing_autocompact = install_staged_callback( - agent, - callback, - callback_type=AutoCompactCallback, - component_attribute="autocompact", - memory_runtime=memory_runtime, - conflict_message="Autocompact is already configured with another runtime", - ) - return existing_autocompact or autocompact diff --git a/trpc_agent_sdk/advanced_memory/_config.py b/trpc_agent_sdk/advanced_memory/_config.py deleted file mode 100644 index 975875d3e..000000000 --- a/trpc_agent_sdk/advanced_memory/_config.py +++ /dev/null @@ -1,257 +0,0 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. -# -# Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Configuration for the independent Advanced Memory mechanism.""" - -from __future__ import annotations - -import os -from dataclasses import dataclass -from dataclasses import field -from pathlib import Path -from typing import Any - -DEFAULT_COMPACTABLE_TOOL_NAMES = ( - "Read", - "Bash", - "Grep", - "Glob", - "WebSearch", - "WebFetch", - "Edit", - "Write", -) - - -def _integer_from_environment( - name: str, - *, - default: int | None, - minimum: int, -) -> int | None: - """Read and validate an optional integer setting from the environment.""" - raw_value = os.environ.get(name, "").strip() - if not raw_value: - return default - try: - value = int(raw_value) - except ValueError as exc: - description = "positive integer" if minimum > 0 else "non-negative integer" - raise ValueError(f"{name} must be a {description}") from exc - if value < minimum: - description = "positive integer" if minimum > 0 else "non-negative integer" - raise ValueError(f"{name} must be a {description}") - return value - - -def _require_positive(**values: int | float) -> None: - """Require each named numeric setting to be greater than zero.""" - for name, value in values.items(): - if value <= 0: - raise ValueError(f"{name} must be greater than zero") - - -def _require_non_negative(**values: int | float) -> None: - """Require each named numeric setting to be non-negative.""" - for name, value in values.items(): - if value < 0: - raise ValueError(f"{name} must not be negative") - - -def _require_less_than( - name: str, - value: int | float, - upper_name: str, - upper_value: int | float, -) -> None: - """Require one named numeric setting to be smaller than another.""" - if value >= upper_value: - raise ValueError(f"{name} must be smaller than {upper_name}") - - -def _require_greater_than( - name: str, - value: int | float, - lower_name: str, - lower_value: int | float, -) -> None: - """Require one named numeric setting to be greater than another.""" - if value <= lower_value: - raise ValueError(f"{name} must be greater than {lower_name}") - - -def _require_non_empty_names(name: str, values: tuple[str, ...]) -> None: - """Require a non-empty sequence containing only non-empty names.""" - if not values or any(not value.strip() for value in values): - raise ValueError(f"{name} must contain non-empty names") - - -def _validate_path_components(values: tuple[str, ...]) -> None: - """Require safe, single-component names for memory storage paths.""" - for value in values: - if not value or Path(value).name != value: - raise ValueError(f"Invalid memory path component: {value!r}") - - -@dataclass(frozen=True) -class AdvancedMemoryConfig: - """Configure the independent memory directory and storage limits.""" - - enabled: bool = True - root_dir: Path = field(default_factory=Path.cwd) - memory_dir_name: str = "MEMORY" - session_dir_name: str = "SESSION" - memory_index_name: str = "MEMORY.md" - transcript_name: str = "transcript.jsonl" - session_memory_name: str = "session_memory.md" - memory_index_max_lines: int = 200 - memory_index_max_bytes: int = 25_000 - long_term_memory_injection_enabled: bool = True - tool_result_max_chars: int = 50_000 - tool_results_per_message_max_chars: int = 200_000 - tool_result_preview_chars: int = 2_000 - history_snip_enabled: bool = True - history_snip_trigger_chars: int = 600_000 - history_snip_target_chars: int = 400_000 - history_snip_keep_recent: int = 5 - history_snip_tool_names: tuple[str, ...] = DEFAULT_COMPACTABLE_TOOL_NAMES - model_context_window_tokens: int | None = field(default_factory=lambda: _integer_from_environment( - "TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS", - default=None, - minimum=1, - )) - max_output_tokens: int = field(default_factory=lambda: _integer_from_environment( - "TRPC_AGENT_MAX_OUTPUT_TOKENS", - default=0, - minimum=0, - )) - token_warning_ratio: float = 0.85 - token_autocompact_ratio: float = 0.90 - token_blocking_ratio: float = 0.95 - token_estimator: Any | None = field(default=None, repr=False, compare=False) - context_window_resolver: Any | None = field(default=None, repr=False, compare=False) - session_memory_enabled: bool = True - session_memory_initial_chars: int = 40_000 - session_memory_update_chars: int = 20_000 - session_memory_initial_tokens: int = 10_000 - session_memory_update_tokens: int = 5_000 - session_memory_tool_calls_between_updates: int = 3 - session_memory_prompt_max_chars: int = 200_000 - session_memory_request_overhead_tokens: int = 2_048 - session_memory_section_max_chars: int = 8_000 - session_memory_total_max_chars: int = 54_000 - session_memory_wait_timeout_seconds: float = 15.0 - autocompact_enabled: bool = True - autocompact_trigger_chars: int = 700_000 - autocompact_target_chars: int = 350_000 - autocompact_blocking_chars: int = 780_000 - autocompact_keep_recent_contents: int = 8 - autocompact_max_failures: int = 3 - autocompact_summary_input_max_chars: int = 600_000 - autocompact_summary_retries: int = 3 - microcompact_enabled: bool = True - microcompact_gap_seconds: float = 3_600.0 - microcompact_trigger_count: int = 20 - microcompact_keep_recent: int = 5 - microcompact_tool_names: tuple[str, ...] = DEFAULT_COMPACTABLE_TOOL_NAMES - encoding: str = "utf-8" - transcript_fsync: bool = False - preload_memory_enabled: bool = False - preload_memory_max_topics: int = 5 - preload_memory_max_chars: int = 50_000 - preload_memory_candidate_limit: int = 200 - - def __post_init__(self) -> None: - """Validate the configuration and normalize the root directory.""" - _require_positive( - memory_index_max_lines=self.memory_index_max_lines, - memory_index_max_bytes=self.memory_index_max_bytes, - preload_memory_max_topics=self.preload_memory_max_topics, - preload_memory_max_chars=self.preload_memory_max_chars, - preload_memory_candidate_limit=self.preload_memory_candidate_limit, - ) - _validate_path_components(( - self.memory_dir_name, - self.session_dir_name, - self.memory_index_name, - self.transcript_name, - self.session_memory_name, - )) - _require_positive( - tool_result_max_chars=self.tool_result_max_chars, - tool_results_per_message_max_chars=self.tool_results_per_message_max_chars, - tool_result_preview_chars=self.tool_result_preview_chars, - ) - _require_less_than( - "tool_result_preview_chars", - self.tool_result_preview_chars, - "tool_result_max_chars", - self.tool_result_max_chars, - ) - _require_positive( - history_snip_trigger_chars=self.history_snip_trigger_chars, - history_snip_target_chars=self.history_snip_target_chars, - ) - _require_less_than( - "history_snip_target_chars", - self.history_snip_target_chars, - "history_snip_trigger_chars", - self.history_snip_trigger_chars, - ) - _require_positive(history_snip_keep_recent=self.history_snip_keep_recent) - _require_non_empty_names("history_snip_tool_names", self.history_snip_tool_names) - if self.model_context_window_tokens is not None and self.model_context_window_tokens <= 0: - raise ValueError("model_context_window_tokens must be greater than zero when provided") - _require_non_negative(max_output_tokens=self.max_output_tokens) - if self.model_context_window_tokens is not None and self.max_output_tokens >= self.model_context_window_tokens: - raise ValueError("max_output_tokens must be smaller than model_context_window_tokens") - if not (0 < self.token_warning_ratio < self.token_autocompact_ratio < self.token_blocking_ratio < 1): - raise ValueError("token ratios must satisfy 0 < warning < autocompact < blocking < 1") - _require_positive( - session_memory_initial_chars=self.session_memory_initial_chars, - session_memory_update_chars=self.session_memory_update_chars, - session_memory_initial_tokens=self.session_memory_initial_tokens, - session_memory_update_tokens=self.session_memory_update_tokens, - session_memory_tool_calls_between_updates=self.session_memory_tool_calls_between_updates, - session_memory_prompt_max_chars=self.session_memory_prompt_max_chars, - ) - _require_non_negative(session_memory_request_overhead_tokens=self.session_memory_request_overhead_tokens) - _require_positive( - session_memory_section_max_chars=self.session_memory_section_max_chars, - session_memory_total_max_chars=self.session_memory_total_max_chars, - session_memory_wait_timeout_seconds=self.session_memory_wait_timeout_seconds, - ) - _require_positive(autocompact_target_chars=self.autocompact_target_chars) - _require_greater_than( - "autocompact_trigger_chars", - self.autocompact_trigger_chars, - "autocompact_target_chars", - self.autocompact_target_chars, - ) - _require_greater_than( - "autocompact_blocking_chars", - self.autocompact_blocking_chars, - "autocompact_trigger_chars", - self.autocompact_trigger_chars, - ) - _require_positive( - autocompact_keep_recent_contents=self.autocompact_keep_recent_contents, - autocompact_max_failures=self.autocompact_max_failures, - autocompact_summary_input_max_chars=self.autocompact_summary_input_max_chars, - autocompact_summary_retries=self.autocompact_summary_retries, - ) - _require_positive( - microcompact_gap_seconds=self.microcompact_gap_seconds, - microcompact_trigger_count=self.microcompact_trigger_count, - microcompact_keep_recent=self.microcompact_keep_recent, - ) - _require_less_than( - "microcompact_keep_recent", - self.microcompact_keep_recent, - "microcompact_trigger_count", - self.microcompact_trigger_count, - ) - _require_non_empty_names("microcompact_tool_names", self.microcompact_tool_names) - object.__setattr__(self, "root_dir", self.root_dir.expanduser().resolve()) diff --git a/trpc_agent_sdk/advanced_memory/_formats.py b/trpc_agent_sdk/advanced_memory/_formats.py deleted file mode 100644 index f0fa25ad1..000000000 --- a/trpc_agent_sdk/advanced_memory/_formats.py +++ /dev/null @@ -1,191 +0,0 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. -# -# Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Define shared formats for long-term and session memory.""" - -from __future__ import annotations - -import re -from dataclasses import dataclass -from datetime import datetime -from datetime import timezone -from enum import Enum - -_FRONTMATTER_PATTERN = re.compile(r"\A---\n(?P.*?)\n---(?:\n|\Z)", re.DOTALL) -_UPDATED_AT_PATTERN = re.compile(r"^updated_at:\s*(?P\S+)\s*$", re.MULTILINE) - - -def _as_utc(value: datetime) -> datetime: - """Normalize an aware or naive datetime to UTC.""" - if value.tzinfo is None: - value = value.replace(tzinfo=timezone.utc) - return value.astimezone(timezone.utc) - - -class MemoryType(str, Enum): - """Semantic types allowed for long-term memory documents.""" - - USER = "user" - FEEDBACK = "feedback" - PROJECT = "project" - REFERENCE = "reference" - - -@dataclass(frozen=True) -class MemoryIndexEntry: - """Represent one standard entry in MEMORY.md.""" - - name: str - filename: str - summary: str - - def __post_init__(self) -> None: - """Validate that index fields are non-empty single-line strings.""" - for field_name, value in ( - ("name", self.name), - ("filename", self.filename), - ("summary", self.summary), - ): - if not value.strip() or "\n" in value or "\r" in value: - raise ValueError(f"{field_name} must be non-empty single-line text") - - def to_markdown(self) -> str: - """Render one long-term memory index entry.""" - return f"- [{self.name.strip()}]({self.filename.strip()}):{self.summary.strip()}" - - -@dataclass(frozen=True) -class MemoryDocument: - """Represent a long-term memory document with frontmatter.""" - - name: str - description: str - memory_type: MemoryType - content: str - updated_at: datetime | None = None - - def __post_init__(self) -> None: - """Validate frontmatter and reject unsafe multiline values.""" - for field_name, value in ( - ("name", self.name), - ("description", self.description), - ): - if not value.strip() or "\n" in value or "\r" in value: - raise ValueError(f"{field_name} must be non-empty single-line text") - - def to_markdown(self) -> str: - """Render standard frontmatter and document content.""" - body = self.content.strip() - updated_at = (_as_utc(self.updated_at).isoformat() if self.updated_at is not None else None) - updated_at_line = f"updated_at: {updated_at}\n" if updated_at else "" - return ("---\n" - f"name: {self.name.strip()}\n" - f"description: {self.description.strip()}\n" - f"type: {self.memory_type.value}\n" - f"{updated_at_line}" - "---\n" - f"{body}\n") - - -def parse_memory_updated_at(content: str) -> datetime | None: - """Extract the UTC update timestamp from a memory document.""" - frontmatter_match = _FRONTMATTER_PATTERN.match(content) - if frontmatter_match is None: - return None - match = _UPDATED_AT_PATTERN.search(frontmatter_match.group("frontmatter")) - if match is None: - return None - try: - value = match.group("value").replace("Z", "+00:00") - parsed = datetime.fromisoformat(value) - except ValueError: - return None - if parsed.tzinfo is None: - parsed = parsed.replace(tzinfo=timezone.utc) - return _as_utc(parsed) - - -def memory_freshness(updated_at: datetime | None, *, now: datetime | None = None) -> str: - """Return a compact freshness bucket suitable for model-facing output.""" - if updated_at is None: - return "unknown" - current = _as_utc(now or datetime.now(timezone.utc)) - timestamp = _as_utc(updated_at) - age_days = max(0, int((current - timestamp).total_seconds()) // 86_400) - if age_days == 0: - return "today" - if age_days == 1: - return "yesterday" - if age_days <= 7: - return "within 7 days" - if age_days <= 30: - return "within 30 days" - return "over 30 days" - - -SESSION_MEMORY_SECTIONS = ( - "Session Title", - "Current State", - "Task specification", - "Files and Functions", - "Workflow", - "Errors & Corrections", - "Codebase and System Documentation", - "Learnings", - "Key results", - "Worklog", -) - -SESSION_MEMORY_SECTION_DESCRIPTIONS = ( - "A short and distinctive 5-10 word descriptive title for the session", - "What is actively being worked on right now? Pending tasks not yet completed.", - "What did the user ask to build? Any design decisions or other explanatory context", - "What are the important files? In short, what do they contain?", - "What bash commands are usually run and in what order?", - "Errors encountered and how they were fixed. What approaches failed?", - "What are the important system components? How do they work/fit together?", - "What has worked well? What has not? What to avoid?", - "If the user asked a specific output, repeat the exact result here", - "Step by step, what was attempted, done? Very terse summary", -) - - -@dataclass(frozen=True) -class SessionMemoryDocument: - """Represent structured session memory with ten fixed sections.""" - - session_title: str = "" - current_state: str = "" - task_specification: str = "" - files_and_functions: str = "" - workflow: str = "" - errors_and_corrections: str = "" - codebase_and_system_documentation: str = "" - learnings: str = "" - key_results: str = "" - worklog: str = "" - - def to_markdown(self) -> str: - """Render all sections in fixed order, including empty sections.""" - values = ( - self.session_title, - self.current_state, - self.task_specification, - self.files_and_functions, - self.workflow, - self.errors_and_corrections, - self.codebase_and_system_documentation, - self.learnings, - self.key_results, - self.worklog, - ) - sections = [ - f"# {section}\n_{description}_\n\n{value.strip()}" for section, description, value in zip( - SESSION_MEMORY_SECTIONS, - SESSION_MEMORY_SECTION_DESCRIPTIONS, - values, - ) - ] - return "\n\n".join(sections).rstrip() + "\n" diff --git a/trpc_agent_sdk/advanced_memory/_integration.py b/trpc_agent_sdk/advanced_memory/_integration.py deleted file mode 100644 index e4f2ca3e9..000000000 --- a/trpc_agent_sdk/advanced_memory/_integration.py +++ /dev/null @@ -1,188 +0,0 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. -# -# Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Provide the one-shot entry point for the context pipeline.""" - -from __future__ import annotations - -from dataclasses import dataclass -from typing import Any -from typing import TYPE_CHECKING - -from ._autocompact import AutoCompact -from ._autocompact import LegacySummaryGenerator -from ._autocompact import setup_autocompact -from ._history_snip import HistorySnip -from ._history_snip import setup_history_snip -from ._memory_context import LongTermMemoryContext -from ._memory_context import setup_long_term_memory_context -from ._microcompact import Microcompact -from ._microcompact import setup_microcompact -from ._runtime import AdvancedMemoryRuntime -from ._session_memory import SessionMemoryExtractor -from ._session_memory import SessionMemoryGenerator -from ._session_service import TranscriptSessionService -from ._tool_result_budget import setup_tool_result_budget -from ._tool_result_budget import ToolResultBudget - -if TYPE_CHECKING: - from trpc_agent_sdk.agents import LlmAgent - from trpc_agent_sdk.sessions import SessionServiceABC - from trpc_agent_sdk.tools._advanced_memory_tool import AdvancedMemoryTools - - -@dataclass(frozen=True) -class AdvancedContextManagement: - """Aggregate the five components installed by one setup call.""" - - long_term_memory: LongTermMemoryContext - tool_result_budget: ToolResultBudget - history_snip: HistorySnip - microcompact: Microcompact - autocompact: AutoCompact - - -@dataclass(frozen=True) -class AdvancedMemoryIntegration: - """Aggregate Agent callbacks, the session memory extractor, and service.""" - - context_management: AdvancedContextManagement - session_memory_extractor: SessionMemoryExtractor - session_service: TranscriptSessionService - long_term_memory_tools: "AdvancedMemoryTools | None" - - -def _setup_long_term_memory_tools( - agent: "LlmAgent", - memory_runtime: AdvancedMemoryRuntime, -) -> "AdvancedMemoryTools": - """Install the three official memory tools idempotently.""" - from trpc_agent_sdk.tools._advanced_memory_tool import ( - ADVANCED_MEMORY_TOOL_NAMES, ) - from trpc_agent_sdk.tools._advanced_memory_tool import AdvancedMemoryTools - - matching_tools = [tool for tool in agent.tools if getattr(tool, "name", None) in ADVANCED_MEMORY_TOOL_NAMES] - if matching_tools: - owners = {getattr(getattr(tool, "func", None), "__self__", None) for tool in matching_tools} - if len(owners) != 1: - raise ValueError("Advanced Memory tool names are already used by different tools") - owner = owners.pop() - if not isinstance(owner, AdvancedMemoryTools): - raise ValueError("Advanced Memory tool names are already used by non-SDK tools") - if owner.runtime is not memory_runtime: - raise ValueError("Advanced Memory tools use another runtime") - installed_names = {getattr(tool, "name", None) for tool in matching_tools} - if installed_names != ADVANCED_MEMORY_TOOL_NAMES: - raise ValueError("Advanced Memory tools are only partially installed") - return owner - tools = AdvancedMemoryTools(memory_runtime) - agent.tools.extend(tools.as_tools()) - return tools - - -def _setup_preload_memory_tool( - agent: "LlmAgent", - memory_runtime: AdvancedMemoryRuntime, - model: Any | None = None, -) -> None: - """Install the automatic topic-memory preprocessor when enabled.""" - if not memory_runtime.config.enabled or not memory_runtime.config.preload_memory_enabled: - return - from trpc_agent_sdk.advanced_memory._preload_memory import MemoryPreloader - from trpc_agent_sdk.advanced_memory._preload_memory import ( - ModelMemoryRelevanceSelector, ) - from trpc_agent_sdk.tools import PreloadMemoryTool - - existing = [tool for tool in agent.tools if getattr(tool, "name", None) == "preload_memory"] - use_legacy_memory = False - if existing: - if len(existing) != 1 or not isinstance(existing[0], PreloadMemoryTool): - raise ValueError("Advanced Memory preload tool name is already used by another tool") - use_legacy_memory = existing[0].uses_legacy_memory - agent.tools.remove(existing[0]) - preloader = MemoryPreloader(memory_runtime, ModelMemoryRelevanceSelector(model)) - agent.tools.append(PreloadMemoryTool( - memory_preloader=preloader.preload, - use_legacy_memory=use_legacy_memory, - )) - - -def setup_context_management( - agent: "LlmAgent", - memory_runtime: AdvancedMemoryRuntime, - summary_generator: LegacySummaryGenerator | None = None, - *, - compact_model: Any | None = None, -) -> AdvancedContextManagement: - """Install the complete Advanced Memory pipeline in fixed stages.""" - return AdvancedContextManagement( - long_term_memory=setup_long_term_memory_context(agent, memory_runtime), - tool_result_budget=setup_tool_result_budget(agent, memory_runtime), - history_snip=setup_history_snip(agent, memory_runtime), - microcompact=setup_microcompact(agent, memory_runtime), - autocompact=setup_autocompact( - agent, - memory_runtime, - summary_generator, - model=compact_model, - ), - ) - - -def setup_advanced_memory( - agent: "LlmAgent", - session_service: "SessionServiceABC", - memory_runtime: AdvancedMemoryRuntime, - summary_generator: LegacySummaryGenerator | None = None, - session_memory_generator: SessionMemoryGenerator | None = None, - *, - compact_model: Any | None = None, - session_memory_model: Any | None = None, - preload_memory_model: Any | None = None, - install_long_term_memory_tools: bool = True, -) -> AdvancedMemoryIntegration: - """Assemble callbacks, the transcript decorator, and session memory.""" - context_management = setup_context_management( - agent, - memory_runtime, - summary_generator, - compact_model=compact_model, - ) - long_term_memory_tools = (_setup_long_term_memory_tools(agent, memory_runtime) - if install_long_term_memory_tools and memory_runtime.config.enabled else None) - _setup_preload_memory_tool(agent, memory_runtime, model=preload_memory_model) - if isinstance(session_service, TranscriptSessionService): - if session_service.memory_runtime is not memory_runtime: - raise ValueError("Transcript session service uses another runtime") - extractor = session_service.session_memory_extractor - if extractor is not None: - if session_memory_generator is not None or session_memory_model is not None: - raise ValueError("Session memory extractor is already configured; " - "do not provide another generator or model") - else: - extractor = SessionMemoryExtractor( - memory_runtime, - session_memory_generator, - model=session_memory_model, - ) - session_service.attach_session_memory_extractor(extractor) - wrapped_service = session_service - else: - extractor = SessionMemoryExtractor( - memory_runtime, - session_memory_generator, - model=session_memory_model, - ) - wrapped_service = TranscriptSessionService( - session_service, - memory_runtime, - extractor, - ) - return AdvancedMemoryIntegration( - context_management=context_management, - session_memory_extractor=extractor, - session_service=wrapped_service, - long_term_memory_tools=long_term_memory_tools, - ) diff --git a/trpc_agent_sdk/advanced_memory/_paths.py b/trpc_agent_sdk/advanced_memory/_paths.py deleted file mode 100644 index da1a41e07..000000000 --- a/trpc_agent_sdk/advanced_memory/_paths.py +++ /dev/null @@ -1,101 +0,0 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. -# -# Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Safe path resolution for the independent memory mechanism.""" - -from __future__ import annotations - -import hashlib -import re -from dataclasses import dataclass -from pathlib import Path - -from ._config import AdvancedMemoryConfig - -_SAFE_COMPONENT_PATTERN = re.compile(r"[^A-Za-z0-9._-]+") - - -def _safe_component(value: str, *, field_name: str) -> str: - """Convert an external identifier into a safe path component.""" - normalized = _SAFE_COMPONENT_PATTERN.sub("_", value.strip()).strip("._") - if not normalized: - raise ValueError(f"{field_name} must contain at least one safe character") - return normalized - - -def _collision_safe_component(value: str, *, field_name: str) -> str: - """Add a digest when sanitization could cause path collisions.""" - stripped = value.strip() - normalized = _safe_component(stripped, field_name=field_name) - if normalized == stripped: - return normalized - digest = hashlib.sha256(stripped.encode("utf-8")).hexdigest()[:12] - return f"{normalized}-{digest}" - - -@dataclass(frozen=True) -class AdvancedMemoryPaths: - """Build all disk paths for long-term and session memory.""" - - config: AdvancedMemoryConfig - - @property - def memory_dir(self) -> Path: - """Return the long-term memory directory.""" - return self.config.root_dir / self.config.memory_dir_name - - @property - def session_root_dir(self) -> Path: - """Return the root directory for session memory.""" - return self.config.root_dir / self.config.session_dir_name - - @property - def memory_index_path(self) -> Path: - """Return the long-term memory index path.""" - return self.memory_dir / self.config.memory_index_name - - def memory_topic_path(self, topic_name: str) -> Path: - """Return a safe path for a long-term memory topic.""" - safe_name = _collision_safe_component(topic_name, field_name="topic_name") - if not safe_name.lower().endswith(".md"): - safe_name = f"{safe_name}.md" - if safe_name == self.config.memory_index_name: - raise ValueError("Topic file cannot overwrite the memory index") - return self.memory_dir / safe_name - - def session_dir(self, session_id: str) -> Path: - """Return the isolated storage directory for a session.""" - return self.session_root_dir / _collision_safe_component( - session_id, - field_name="session_id", - ) - - def transcript_path(self, session_id: str) -> Path: - """Return the transcript path for a session.""" - return self.session_dir(session_id) / self.config.transcript_name - - def session_memory_path(self, session_id: str) -> Path: - """Return the session memory path for a session.""" - return self.session_dir(session_id) / self.config.session_memory_name - - def tool_results_dir(self, session_id: str) -> Path: - """Return the large tool-result directory for a session.""" - return self.session_dir(session_id) / "tool-results" - - def tool_result_path(self, session_id: str, result_id: str) -> Path: - """Return a safe JSON path for a large tool result.""" - safe_result_id = _collision_safe_component(result_id, field_name="result_id") - return self.tool_results_dir(session_id) / f"{safe_result_id}.json" - - def ensure_base_directories(self) -> None: - """Create the long-term and session memory directories.""" - self.memory_dir.mkdir(parents=True, exist_ok=True) - self.session_root_dir.mkdir(parents=True, exist_ok=True) - - def ensure_session_directory(self, session_id: str) -> Path: - """Create and return a session's storage directory.""" - path = self.session_dir(session_id) - path.mkdir(parents=True, exist_ok=True) - return path diff --git a/trpc_agent_sdk/advanced_memory/_runtime.py b/trpc_agent_sdk/advanced_memory/_runtime.py deleted file mode 100644 index c26def35f..000000000 --- a/trpc_agent_sdk/advanced_memory/_runtime.py +++ /dev/null @@ -1,53 +0,0 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. -# -# Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Unified runtime entry point for the independent memory mechanism.""" - -from __future__ import annotations - -from dataclasses import dataclass - -from ._config import AdvancedMemoryConfig -from ._coordination import SessionOperationCoordinator -from ._paths import AdvancedMemoryPaths -from ._storage import LongTermMemoryStore -from ._storage import SessionMemoryStore -from ._storage import ToolResultStore -from ._storage import TranscriptStore - - -@dataclass(frozen=True) -class AdvancedMemoryRuntime: - """Aggregate configuration, paths, and the three storage objects.""" - - config: AdvancedMemoryConfig - paths: AdvancedMemoryPaths - coordination: SessionOperationCoordinator - long_term_memory: LongTermMemoryStore - session_memory: SessionMemoryStore - tool_results: ToolResultStore - transcripts: TranscriptStore - - @classmethod - def create(cls, config: AdvancedMemoryConfig | None = None) -> "AdvancedMemoryRuntime": - """Create a runtime isolated from the legacy mechanism.""" - resolved_config = config or AdvancedMemoryConfig() - paths = AdvancedMemoryPaths(resolved_config) - return cls( - config=resolved_config, - paths=paths, - coordination=SessionOperationCoordinator(), - long_term_memory=LongTermMemoryStore(resolved_config, paths), - session_memory=SessionMemoryStore(resolved_config, paths), - tool_results=ToolResultStore(resolved_config, paths), - transcripts=TranscriptStore(resolved_config, paths), - ) - - async def initialize(self) -> bool: - """Create memory directories only when the mechanism is enabled.""" - if not self.config.enabled: - return False - await self.long_term_memory.initialize() - return True diff --git a/trpc_agent_sdk/advanced_memory/_session_service.py b/trpc_agent_sdk/advanced_memory/_session_service.py deleted file mode 100644 index a8cccd52d..000000000 --- a/trpc_agent_sdk/advanced_memory/_session_service.py +++ /dev/null @@ -1,207 +0,0 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. -# -# Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Decorate a SessionService to record a complete transcript.""" - -from __future__ import annotations - -from typing import Any -from typing import TYPE_CHECKING - -from trpc_agent_sdk.abc import ListSessionsResponse -from trpc_agent_sdk.abc import ResponseABC -from trpc_agent_sdk.abc import SessionABC -from trpc_agent_sdk.abc import SessionServiceABC - -if TYPE_CHECKING: - from trpc_agent_sdk.context import AgentContext - from trpc_agent_sdk.context import InvocationContext - -from ._runtime import AdvancedMemoryRuntime -from ._coordination import CrossLoopLock -from ._session_memory import SessionMemoryExtractor -from ._transcript import build_event_transcript_record -from ._transcript import find_last_event_id - - -class TranscriptSessionService(SessionServiceABC): - """Decorate a legacy SessionService and append persisted Events.""" - - def __init__( - self, - delegate: SessionServiceABC, - memory_runtime: AdvancedMemoryRuntime, - session_memory_extractor: SessionMemoryExtractor | None = None, - ) -> None: - """Store the legacy service and optional Advanced Memory runtime.""" - self._delegate = delegate - self._memory_runtime = memory_runtime - self._session_memory_extractor = session_memory_extractor - self._initialize_lock = CrossLoopLock() - self._initialized = False - self._session_locks: dict[str, CrossLoopLock] = {} - self._loaded_parent_sessions: set[str] = set() - self._last_event_ids: dict[str, str | None] = {} - - @property - def delegate(self) -> SessionServiceABC: - """Return the unchanged underlying SessionService.""" - return self._delegate - - @property - def memory_runtime(self) -> AdvancedMemoryRuntime: - """Return the Advanced Memory runtime used by the decorator.""" - return self._memory_runtime - - @property - def session_memory_extractor(self) -> SessionMemoryExtractor | None: - """Return the session memory extractor used after each turn.""" - return self._session_memory_extractor - - def attach_session_memory_extractor( - self, - extractor: SessionMemoryExtractor, - ) -> None: - """Attach a session memory extractor when one is not configured.""" - if self._session_memory_extractor is not None: - if self._session_memory_extractor is not extractor: - raise ValueError("Session memory extractor is already configured") - return - if extractor.runtime is not self._memory_runtime: - raise ValueError("Session memory extractor uses another runtime") - self._session_memory_extractor = extractor - - async def _ensure_initialized(self) -> None: - """Initialize memory directories before the first transcript write.""" - if self._initialized or not self._memory_runtime.config.enabled: - return - async with self._initialize_lock: - if self._initialized: - return - self._initialized = await self._memory_runtime.initialize() - - def _session_lock(self, session_id: str) -> CrossLoopLock: - """Return an independent asynchronous write lock per session.""" - lock = self._session_locks.get(session_id) - if lock is None: - lock = CrossLoopLock() - self._session_locks[session_id] = lock - return lock - - async def _load_parent_if_needed(self, session_id: str) -> None: - """Restore the parent-chain tail before the first session write.""" - if session_id in self._loaded_parent_sessions: - return - records = await self._memory_runtime.transcripts.read_all(session_id) - self._last_event_ids[session_id] = find_last_event_id(records) - self._loaded_parent_sessions.add(session_id) - - async def create_session( - self, - *, - app_name: str, - user_id: str, - state: dict[str, Any] | None = None, - session_id: str | None = None, - agent_context: AgentContext | None = None, - ) -> SessionABC: - """Delegate session creation to the underlying service.""" - return await self._delegate.create_session( - app_name=app_name, - user_id=user_id, - state=state, - session_id=session_id, - agent_context=agent_context, - ) - - async def get_session( - self, - *, - app_name: str, - user_id: str, - session_id: str, - agent_context: AgentContext | None = None, - ) -> SessionABC | None: - """Delegate session reads to the underlying service.""" - return await self._delegate.get_session( - app_name=app_name, - user_id=user_id, - session_id=session_id, - agent_context=agent_context, - ) - - async def list_sessions( - self, - *, - app_name: str, - user_id: str | None = None, - ) -> ListSessionsResponse: - """Delegate session listing to the underlying service.""" - return await self._delegate.list_sessions(app_name=app_name, user_id=user_id) - - async def delete_session(self, *, app_name: str, user_id: str, session_id: str) -> None: - """Delete only the legacy session and retain transcript records.""" - async with self._session_lock(session_id): - await self._delegate.delete_session( - app_name=app_name, - user_id=user_id, - session_id=session_id, - ) - self._session_locks.pop(session_id, None) - self._loaded_parent_sessions.discard(session_id) - self._last_event_ids.pop(session_id, None) - - async def append_event(self, session: SessionABC, event: ResponseABC) -> ResponseABC: - """Append each persisted non-streaming Event in order.""" - usage_metadata = getattr(event, "usage_metadata", None) - state = getattr(session, "state", None) - context_fingerprint = (state.get("advanced_memory_pending_request_context_fingerprint") if isinstance( - state, dict) else None) - if usage_metadata is not None and isinstance(context_fingerprint, str): - metadata = dict(getattr(event, "custom_metadata", None) or {}) - metadata["advanced_memory_request_context_fingerprint"] = context_fingerprint - event.custom_metadata = metadata - persisted_event = await self._delegate.append_event(session=session, event=event) - if not self._memory_runtime.config.enabled or getattr(persisted_event, "partial", False): - return persisted_event - - await self._ensure_initialized() - async with self._session_lock(session.id): - await self._load_parent_if_needed(session.id) - record = build_event_transcript_record( - session, - persisted_event, - parent_event_id=self._last_event_ids.get(session.id), - ) - _, appended = await self._memory_runtime.transcripts.append_unique( - session.id, - record, - unique_key="event_id", - ) - if appended: - self._last_event_ids[session.id] = record["event_id"] - return persisted_event - - async def update_session(self, session: SessionABC) -> None: - """Delegate session updates to the underlying service.""" - await self._delegate.update_session(session) - - async def create_session_summary( - self, - session: SessionABC, - ctx: InvocationContext | None = None, - ) -> None: - """Preserve legacy summaries, then update session memory as needed.""" - await self._delegate.create_session_summary(session, ctx=ctx) - if self._session_memory_extractor is not None and ctx is not None: - await self._session_memory_extractor.extract_if_needed(session, ctx) - - async def get_session_summary(self, session: SessionABC) -> str | None: - """Delegate session summary reads to the legacy service.""" - return await self._delegate.get_session_summary(session) - - async def close(self) -> None: - """Close the legacy service while preserving its lifecycle semantics.""" - await self._delegate.close() diff --git a/trpc_agent_sdk/advanced_memory/_storage.py b/trpc_agent_sdk/advanced_memory/_storage.py deleted file mode 100644 index 1fb43591c..000000000 --- a/trpc_agent_sdk/advanced_memory/_storage.py +++ /dev/null @@ -1,331 +0,0 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. -# -# Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Basic disk stores for long-term memory, session memory, and transcripts.""" - -from __future__ import annotations - -import asyncio -import json -import os -import tempfile -import threading -from collections.abc import Mapping -from dataclasses import replace -from datetime import datetime -from datetime import timezone -from pathlib import Path -from typing import Any - -from ._config import AdvancedMemoryConfig -from ._formats import MemoryDocument -from ._formats import MemoryIndexEntry -from ._formats import SessionMemoryDocument -from ._paths import AdvancedMemoryPaths - - -def _atomic_write_text(path: Path, content: str, *, encoding: str) -> None: - """Atomically replace a text file using a temporary sibling file.""" - path.parent.mkdir(parents=True, exist_ok=True) - file_descriptor, temporary_name = tempfile.mkstemp(prefix=f".{path.name}.", dir=path.parent) - try: - with os.fdopen(file_descriptor, "w", encoding=encoding) as temporary_file: - temporary_file.write(content) - temporary_file.flush() - os.fsync(temporary_file.fileno()) - os.replace(temporary_name, path) - except BaseException: - try: - os.unlink(temporary_name) - except FileNotFoundError: - pass - raise - - -class LongTermMemoryStore: - """Manage MEMORY.md and its detail files in the same directory.""" - - def __init__(self, config: AdvancedMemoryConfig, paths: AdvancedMemoryPaths | None = None) -> None: - """Initialize long-term storage without changing legacy memory.""" - self._config = config - self._paths = paths or AdvancedMemoryPaths(config) - - @property - def index_path(self) -> Path: - """Return the disk path for MEMORY.md.""" - return self._paths.memory_index_path - - async def initialize(self) -> None: - """Create the memory directory and an empty index.""" - await asyncio.to_thread(self._initialize_sync) - - def _initialize_sync(self) -> None: - """Synchronously create the memory directory and empty index.""" - self._paths.ensure_base_directories() - if not self.index_path.exists(): - _atomic_write_text(self.index_path, "", encoding=self._config.encoding) - - async def read_index(self) -> str: - """Read only the configured prefix of MEMORY.md.""" - return await asyncio.to_thread(self._read_index_sync) - - def _read_index_sync(self) -> str: - """Synchronously read MEMORY.md within configured limits.""" - if not self.index_path.exists(): - return "" - with self.index_path.open("r", encoding=self._config.encoding) as index_file: - lines: list[str] = [] - used_bytes = 0 - for _ in range(self._config.memory_index_max_lines): - line = index_file.readline() - if not line: - break - line_bytes = len(line.encode(self._config.encoding)) - if used_bytes + line_bytes > self._config.memory_index_max_bytes: - break - lines.append(line) - used_bytes += line_bytes - return "".join(lines) - - async def write_index(self, entries: list[MemoryIndexEntry]) -> None: - """Atomically write MEMORY.md in the standard index format.""" - content = "\n".join(entry.to_markdown() for entry in entries) - if content: - content += "\n" - await asyncio.to_thread(self._write_index_sync, content) - - def _write_index_sync(self, content: str) -> None: - """Synchronously write MEMORY.md; read_index applies prompt-size limits.""" - _atomic_write_text(self.index_path, content, encoding=self._config.encoding) - - async def read_topic(self, topic_name: str) -> str | None: - """Read a detail memory topic, returning None if absent.""" - path = self._paths.memory_topic_path(topic_name) - return await asyncio.to_thread(self._read_optional_text, path) - - async def read_topic_frontmatter(self, topic_name: str) -> str | None: - """Read only the frontmatter of a detail memory topic.""" - path = self._paths.memory_topic_path(topic_name) - return await asyncio.to_thread(self._read_frontmatter, path) - - def _read_optional_text(self, path: Path) -> str | None: - """Synchronously read an optional text file.""" - if not path.exists(): - return None - return path.read_text(encoding=self._config.encoding) - - def _read_frontmatter(self, path: Path) -> str | None: - """Synchronously read a topic's bounded frontmatter block.""" - if not path.exists(): - return None - lines: list[str] = [] - with path.open(encoding=self._config.encoding) as file: - for line in file: - lines.append(line) - if len(lines) > 1 and line.rstrip("\r\n") == "---": - break - return "".join(lines) - - async def write_topic(self, topic_name: str, document: MemoryDocument) -> Path: - """Atomically write a detail memory file with frontmatter.""" - path = self._paths.memory_topic_path(topic_name) - document = replace(document, updated_at=datetime.now(timezone.utc)) - await asyncio.to_thread( - _atomic_write_text, - path, - document.to_markdown(), - encoding=self._config.encoding, - ) - return path - - async def list_topics(self) -> list[Path]: - """List detail memory files by name, excluding MEMORY.md.""" - return await asyncio.to_thread(self._list_topics_sync) - - def _list_topics_sync(self) -> list[Path]: - """Synchronously list all detail memory files.""" - if not self._paths.memory_dir.exists(): - return [] - return sorted( - (path for path in self._paths.memory_dir.glob("*.md") if path.name != self._config.memory_index_name), - key=lambda path: path.name, - ) - - -class SessionMemoryStore: - """Manage an isolated structured Markdown summary per session.""" - - def __init__(self, config: AdvancedMemoryConfig, paths: AdvancedMemoryPaths | None = None) -> None: - """Initialize session memory storage.""" - self._config = config - self._paths = paths or AdvancedMemoryPaths(config) - - async def read(self, session_id: str) -> str | None: - """Read session memory, returning None if absent.""" - path = self._paths.session_memory_path(session_id) - return await asyncio.to_thread(self._read_sync, path) - - def _read_sync(self, path: Path) -> str | None: - """Synchronously read session memory.""" - if not path.exists(): - return None - return path.read_text(encoding=self._config.encoding) - - async def write(self, session_id: str, document: SessionMemoryDocument) -> Path: - """Atomically write session memory using the fixed section template.""" - path = self._paths.session_memory_path(session_id) - await asyncio.to_thread( - _atomic_write_text, - path, - document.to_markdown(), - encoding=self._config.encoding, - ) - return path - - -class ToolResultStore: - """Persist complete tool results that exceed the context budget.""" - - def __init__(self, config: AdvancedMemoryConfig, paths: AdvancedMemoryPaths | None = None) -> None: - """Initialize large tool-result storage.""" - self._config = config - self._paths = paths or AdvancedMemoryPaths(config) - - async def write(self, session_id: str, result_id: str, serialized_result: str) -> Path: - """Atomically write a complete tool result and return its disk path.""" - path = self._paths.tool_result_path(session_id, result_id) - await asyncio.to_thread( - _atomic_write_text, - path, - serialized_result, - encoding=self._config.encoding, - ) - return path - - async def read(self, session_id: str, result_id: str) -> str | None: - """Read a persisted complete tool result.""" - path = self._paths.tool_result_path(session_id, result_id) - return await asyncio.to_thread(self._read_sync, path) - - def _read_sync(self, path: Path) -> str | None: - """Synchronously read an optional complete tool-result file.""" - if not path.exists(): - return None - return path.read_text(encoding=self._config.encoding) - - -class TranscriptStore: - """Store complete per-session records as append-only JSONL.""" - - def __init__(self, config: AdvancedMemoryConfig, paths: AdvancedMemoryPaths | None = None) -> None: - """Initialize transcript storage and its process-local write lock.""" - self._config = config - self._paths = paths or AdvancedMemoryPaths(config) - self._write_lock = threading.Lock() - self._seen_unique_values: dict[tuple[Path, str], set[str]] = {} - - async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: - """Append one JSON-serializable record to a session transcript.""" - path = self._paths.transcript_path(session_id) - payload = dict(record) - payload.setdefault("recorded_at", datetime.now(timezone.utc).isoformat()) - serialized = json.dumps(payload, ensure_ascii=False, separators=(",", ":"), default=str) - await asyncio.to_thread(self._append_sync, path, serialized) - return path - - def _append_sync(self, path: Path, serialized: str) -> None: - """Synchronously append one transcript line under the write lock.""" - path.parent.mkdir(parents=True, exist_ok=True) - with self._write_lock: - self._append_serialized_unlocked(path, serialized) - - def _append_serialized_unlocked(self, path: Path, serialized: str) -> None: - """Append one serialized line while the caller holds the lock.""" - with path.open("a", encoding=self._config.encoding) as transcript_file: - transcript_file.write(serialized) - transcript_file.write("\n") - transcript_file.flush() - if self._config.transcript_fsync: - os.fsync(transcript_file.fileno()) - - async def append_unique( - self, - session_id: str, - record: Mapping[str, Any], - *, - unique_key: str, - ) -> tuple[Path, bool]: - """Append a transcript record after de-duplicating by a field.""" - path = self._paths.transcript_path(session_id) - payload = dict(record) - unique_value = payload.get(unique_key) - if not isinstance(unique_value, str) or not unique_value: - raise ValueError(f"Transcript unique key {unique_key!r} must be a non-empty string") - payload.setdefault("recorded_at", datetime.now(timezone.utc).isoformat()) - serialized = json.dumps(payload, ensure_ascii=False, separators=(",", ":"), default=str) - appended = await asyncio.to_thread( - self._append_unique_sync, - path, - serialized, - unique_key, - unique_value, - ) - return path, appended - - def _append_unique_sync( - self, - path: Path, - serialized: str, - unique_key: str, - unique_value: str, - ) -> bool: - """Load de-duplication state and append only new records.""" - path.parent.mkdir(parents=True, exist_ok=True) - cache_key = (path, unique_key) - with self._write_lock: - seen_values = self._seen_unique_values.get(cache_key) - if seen_values is None: - seen_values = self._load_unique_values_unlocked(path, unique_key) - self._seen_unique_values[cache_key] = seen_values - if unique_value in seen_values: - return False - self._append_serialized_unlocked(path, serialized) - seen_values.add(unique_value) - return True - - def _load_unique_values_unlocked(self, path: Path, unique_key: str) -> set[str]: - """Load existing de-duplication values while holding the lock.""" - if not path.exists(): - return set() - values: set[str] = set() - with path.open("r", encoding=self._config.encoding) as transcript_file: - for line in transcript_file: - if not line.strip(): - continue - parsed = json.loads(line) - if isinstance(parsed, dict) and isinstance(parsed.get(unique_key), str): - values.add(parsed[unique_key]) - return values - - async def read_all(self, session_id: str) -> list[dict[str, Any]]: - """Read all transcript records for a session in write order.""" - path = self._paths.transcript_path(session_id) - return await asyncio.to_thread(self._read_all_sync, path) - - def _read_all_sync(self, path: Path) -> list[dict[str, Any]]: - """Parse a consistent transcript snapshot under the file lock.""" - with self._write_lock: - if not path.exists(): - return [] - records: list[dict[str, Any]] = [] - with path.open("r", encoding=self._config.encoding) as transcript_file: - for line_number, line in enumerate(transcript_file, start=1): - if not line.strip(): - continue - parsed = json.loads(line) - if not isinstance(parsed, dict): - raise ValueError(f"Transcript line {line_number} is not a JSON object") - records.append(parsed) - return records diff --git a/trpc_agent_sdk/advanced_memory/_transcript.py b/trpc_agent_sdk/advanced_memory/_transcript.py deleted file mode 100644 index 6cfe2192b..000000000 --- a/trpc_agent_sdk/advanced_memory/_transcript.py +++ /dev/null @@ -1,49 +0,0 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. -# -# Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Convert tRPC Events into recoverable transcript records.""" - -from __future__ import annotations - -from typing import Any - -from trpc_agent_sdk.abc import ResponseABC -from trpc_agent_sdk.abc import SessionABC - -TRANSCRIPT_SCHEMA_VERSION = 1 - - -def build_event_transcript_record( - session: SessionABC, - event: ResponseABC, - *, - parent_event_id: str | None, -) -> dict[str, Any]: - """Convert a persisted Event into a versioned transcript record.""" - event_id = getattr(event, "id", "") - if not isinstance(event_id, str) or not event_id: - raise ValueError("Persisted event must have a non-empty id") - event_timestamp = getattr(event, "timestamp", None) - return { - "schema_version": TRANSCRIPT_SCHEMA_VERSION, - "kind": "event", - "event_id": event_id, - "parent_event_id": parent_event_id, - "event_timestamp": event_timestamp, - "session": { - "id": session.id, - "app_name": session.app_name, - "user_id": session.user_id, - }, - "event": event.model_dump(mode="json", by_alias=True, exclude_none=True), - } - - -def find_last_event_id(records: list[dict[str, Any]]) -> str | None: - """Find the last valid Event record identifier in a transcript.""" - for record in reversed(records): - if record.get("kind") == "event" and isinstance(record.get("event_id"), str): - return record["event_id"] - return None diff --git a/trpc_agent_sdk/evaluation/_eval_session_service.py b/trpc_agent_sdk/evaluation/_eval_session_service.py index d9e231dbc..a63712a0d 100644 --- a/trpc_agent_sdk/evaluation/_eval_session_service.py +++ b/trpc_agent_sdk/evaluation/_eval_session_service.py @@ -25,6 +25,11 @@ def __init__(self, inner: BaseSessionService, context_messages: Optional[list] = self._inner = inner self._context_messages = context_messages + @property + def session_config(self): + """Expose the storage service's Session configuration.""" + return self._inner.session_config + @override async def create_session( self, diff --git a/trpc_agent_sdk/memory/__init__.py b/trpc_agent_sdk/memory/__init__.py index 78e525456..e0a69ed5e 100644 --- a/trpc_agent_sdk/memory/__init__.py +++ b/trpc_agent_sdk/memory/__init__.py @@ -7,7 +7,7 @@ This module provides memory/RAG functionality including: - Abstract memory service interfaces -- In-memory memory service implementation +- In-memory, Redis, SQL, and Advanced Memory implementations """ from trpc_agent_sdk.abc import MemoryServiceABC as BaseMemoryService @@ -27,7 +27,6 @@ __all__ = [ "BaseMemoryService", "MemoryServiceConfig", - "AdvancedMemoryConfig", "AdvancedMemoryService", "EventTtl", "InMemoryMemoryService", @@ -39,12 +38,3 @@ "extract_words_lower", "format_timestamp", ] - - -def __getattr__(name: str): - """Lazily expose Advanced Memory configuration without import cycles.""" - if name == "AdvancedMemoryConfig": - from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig - - return AdvancedMemoryConfig - raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/trpc_agent_sdk/memory/_advanced_memory_service.py b/trpc_agent_sdk/memory/_advanced_memory_service.py index 8cc2c97f8..a95406f5e 100644 --- a/trpc_agent_sdk/memory/_advanced_memory_service.py +++ b/trpc_agent_sdk/memory/_advanced_memory_service.py @@ -11,59 +11,55 @@ from typing import Optional from typing import TYPE_CHECKING -from trpc_agent_sdk.abc import MemoryServiceABC as BaseMemoryService +from typing_extensions import override + +from trpc_agent_sdk.abc import MemoryServiceABC from trpc_agent_sdk.abc import MemoryServiceConfig from trpc_agent_sdk.abc import SearchMemoryResponse +from trpc_agent_sdk.abc import SessionABC from trpc_agent_sdk.abc import SessionServiceABC from trpc_agent_sdk.context import AgentContext -from trpc_agent_sdk.sessions import Session if TYPE_CHECKING: - from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig - from trpc_agent_sdk.advanced_memory import AdvancedMemoryIntegration - from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime + from trpc_agent_sdk.memory.advanced_memory import AdvancedMemoryServiceConfig + from trpc_agent_sdk.memory.advanced_memory import AdvancedMemoryRuntime + from trpc_agent_sdk.memory.advanced_memory import LongTermMemoryIntegration -class AdvancedMemoryService(BaseMemoryService): - """Expose Advanced Memory through the standard Runner memory API. +class AdvancedMemoryService(MemoryServiceABC): + """Expose tool-driven long-term Memory through the Runner memory API. - Advanced Memory is more than a traditional ``MemoryServiceABC``: it also - installs agent callbacks and decorates the session service. ``Runner`` - calls :meth:`bind` automatically when this service is supplied as its - ``memory_service``. + ``Runner`` calls :meth:`bind` automatically. The standard + :class:`MemoryServiceABC` methods are implemented for lifecycle + compatibility; long-term memory is intentionally still written and read + by the Agent through the Advanced Memory tools. Session compression is + configured independently through ``SessionService.session_compact_manager``. """ def __init__( self, - config: AdvancedMemoryConfig | None = None, + config: AdvancedMemoryServiceConfig | None = None, *, runtime: AdvancedMemoryRuntime | None = None, - summary_generator: Any | None = None, - session_memory_generator: Any | None = None, - compact_model: Any | None = None, - session_memory_model: Any | None = None, + preload_memory_model: Any | None = None, install_long_term_memory_tools: bool = True, ) -> None: """Create an Advanced Memory service without binding it to an agent.""" - from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig - from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime + from trpc_agent_sdk.memory.advanced_memory import AdvancedMemoryServiceConfig + from trpc_agent_sdk.memory.advanced_memory import AdvancedMemoryRuntime if config is not None and runtime is not None and config != runtime.config: raise ValueError("config and runtime must describe the same Advanced Memory configuration") - resolved_config = runtime.config if runtime is not None else (config or AdvancedMemoryConfig()) + resolved_config = runtime.config if runtime is not None else (config or AdvancedMemoryServiceConfig()) super().__init__(MemoryServiceConfig(enabled=resolved_config.enabled)) self._runtime = runtime or AdvancedMemoryRuntime.create(resolved_config) - self._summary_generator = summary_generator - self._session_memory_generator = session_memory_generator - self._compact_model = compact_model - self._session_memory_model = session_memory_model + self._preload_memory_model = preload_memory_model self._install_long_term_memory_tools = install_long_term_memory_tools - self._integration: AdvancedMemoryIntegration | None = None + self._integration: LongTermMemoryIntegration | None = None self._bound_agent: Any | None = None - self._bound_session_service: SessionServiceABC | None = None @property - def config(self) -> AdvancedMemoryConfig: + def config(self) -> AdvancedMemoryServiceConfig: """Return the Advanced Memory configuration.""" return self._runtime.config @@ -73,47 +69,43 @@ def runtime(self) -> AdvancedMemoryRuntime: return self._runtime @property - def integration(self) -> AdvancedMemoryIntegration | None: + def integration(self) -> LongTermMemoryIntegration | None: """Return the binding result after the service is attached to a Runner.""" return self._integration def bind(self, agent: Any, session_service: SessionServiceABC) -> SessionServiceABC: - """Bind callbacks and tools, returning the wrapped session service.""" - from trpc_agent_sdk.advanced_memory import setup_advanced_memory + """Bind long-term Memory and return the unchanged SessionService.""" + from trpc_agent_sdk.memory.advanced_memory import setup_long_term_memory if self._integration is not None: if agent is not self._bound_agent: raise ValueError("AdvancedMemoryService is already bound to another agent") - if session_service is not self._bound_session_service: - raise ValueError("AdvancedMemoryService is already bound to another session service") - return self._integration.session_service + return session_service - self._integration = setup_advanced_memory( + self._integration = setup_long_term_memory( agent, - session_service, self._runtime, - self._summary_generator, - self._session_memory_generator, - compact_model=self._compact_model, - session_memory_model=self._session_memory_model, - install_long_term_memory_tools=self._install_long_term_memory_tools, + preload_memory_model=self._preload_memory_model, + install_tools=self._install_long_term_memory_tools, ) self._bound_agent = agent - self._bound_session_service = session_service - return self._integration.session_service + return session_service + @override async def store_session( self, - session: Session, + session: SessionABC, agent_context: Optional[AgentContext] = None, ) -> None: - """Keep the standard Runner post-turn contract without duplicating work. + """Keep the standard hook side-effect free. - The wrapped session service performs session-memory extraction from - ``create_session_summary`` before Runner reaches this method. + Advanced Memory is model-directed: the Agent decides what is durable + and calls ``save_memory``. Automatically storing every Session here + would mix transient conversation history with long-term memory. """ return None + @override async def search_memory( self, key: str, @@ -121,17 +113,14 @@ async def search_memory( limit: int = 10, agent_context: Optional[AgentContext] = None, ) -> SearchMemoryResponse: - """Return an empty legacy-style response. + """Return the standard empty response for compatibility. Advanced long-term memory is intentionally accessed through its ``save_memory``, ``read_memory``, and ``list_memory_index`` tools. """ return SearchMemoryResponse() + @override async def close(self) -> None: - """Release service-owned resources. - - Advanced Memory stores are file-backed and do not own an external - connection. The wrapped session service is closed by Runner. - """ - return None + """Release service-owned local or external storage resources.""" + await self._runtime.close() diff --git a/trpc_agent_sdk/memory/advanced_memory/__init__.py b/trpc_agent_sdk/memory/advanced_memory/__init__.py new file mode 100644 index 000000000..04625f61a --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/__init__.py @@ -0,0 +1,53 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Optional long-term memory APIs.""" + +from ._config import AdvancedMemoryServiceConfig +from ._formats import MemoryDocument +from ._formats import MemoryIndexEntry +from ._formats import MemoryType +from ._formats import memory_freshness +from ._formats import parse_memory_updated_at +from ._paths import AdvancedMemoryPaths +from ._paths import MemoryScope +from ._runtime import AdvancedMemoryRuntime +from ._runtime import ScopedAdvancedMemoryRuntime +from ._storage import LongTermMemoryStore + +from ._integration import LongTermMemoryIntegration +from ._integration import setup_long_term_memory +from ._memory_context import LongTermMemoryContext +from ._memory_context import LongTermMemoryContextCallback +from ._memory_context import setup_long_term_memory_context +from ._preload_memory import MemoryCandidate +from ._preload_memory import MemoryPreloader +from ._preload_memory import MemoryRelevanceSelector +from ._preload_memory import ModelMemoryRelevanceSelector +from ._preload_memory import select_relevant_memory_filenames + +__all__ = [ + "AdvancedMemoryServiceConfig", + "LongTermMemoryIntegration", + "AdvancedMemoryPaths", + "AdvancedMemoryRuntime", + "ScopedAdvancedMemoryRuntime", + "LongTermMemoryStore", + "LongTermMemoryContext", + "LongTermMemoryContextCallback", + "MemoryDocument", + "MemoryScope", + "MemoryIndexEntry", + "MemoryType", + "MemoryCandidate", + "MemoryPreloader", + "MemoryRelevanceSelector", + "ModelMemoryRelevanceSelector", + "select_relevant_memory_filenames", + "memory_freshness", + "parse_memory_updated_at", + "setup_long_term_memory_context", + "setup_long_term_memory", +] diff --git a/trpc_agent_sdk/memory/advanced_memory/_config.py b/trpc_agent_sdk/memory/advanced_memory/_config.py new file mode 100644 index 000000000..9bb9fcd2f --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_config.py @@ -0,0 +1,84 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Configuration for the independent Advanced Memory mechanism.""" + +from __future__ import annotations + +from dataclasses import dataclass +from dataclasses import field +from pathlib import Path +from typing import Literal + + +def _require_positive(**values: int | float) -> None: + """Require each named numeric setting to be greater than zero.""" + for name, value in values.items(): + if value <= 0: + raise ValueError(f"{name} must be greater than zero") + + +def _validate_path_components(values: tuple[str, ...]) -> None: + """Require safe, single-component names for memory storage paths.""" + for value in values: + if not value or Path(value).name != value: + raise ValueError(f"Invalid memory path component: {value!r}") + + +@dataclass(frozen=True) +class AdvancedMemoryServiceConfig: + """Configure the independent long-term Advanced Memory service.""" + + enabled: bool = True + root_dir: Path = field(default_factory=Path.cwd) + storage_backend: Literal["local", "redis", "sql"] = "local" + redis_url: str | None = None + redis_key_prefix: str = "advanced-memory:v1" + redis_is_async: bool = True + sql_url: str | None = None + sql_is_async: bool = True + sql_cleanup_interval_seconds: float = 60.0 + memory_ttl_seconds: int | None = None + memory_lock_ttl_seconds: int = 30 + memory_lock_acquire_timeout_seconds: float = 10.0 + memory_dir_name: str = "MEMORY" + memory_index_name: str = "MEMORY.md" + memory_index_max_lines: int = 200 + memory_index_max_bytes: int = 25_000 + long_term_memory_injection_enabled: bool = True + memory_focus_instruction: str | None = None + encoding: str = "utf-8" + preload_memory_enabled: bool = False + preload_memory_max_topics: int = 5 + preload_memory_max_chars: int = 50_000 + preload_memory_candidate_limit: int = 200 + + def __post_init__(self) -> None: + """Validate the configuration and normalize the root directory.""" + if self.storage_backend not in {"local", "redis", "sql"}: + raise ValueError("storage_backend must be one of: local, redis, sql") + if self.storage_backend == "redis" and not self.redis_url: + raise ValueError("redis_url is required when storage_backend='redis'") + if self.storage_backend == "sql" and not self.sql_url: + raise ValueError("sql_url is required when storage_backend='sql'") + if not self.redis_key_prefix.strip() or self.redis_key_prefix != self.redis_key_prefix.strip(): + raise ValueError("redis_key_prefix must be a non-empty Redis key prefix") + if self.memory_ttl_seconds is not None and self.memory_ttl_seconds <= 0: + raise ValueError("memory_ttl_seconds must be greater than zero when provided") + if self.memory_lock_ttl_seconds <= 0: + raise ValueError("memory_lock_ttl_seconds must be greater than zero") + if self.memory_lock_acquire_timeout_seconds <= 0: + raise ValueError("memory_lock_acquire_timeout_seconds must be greater than zero") + if self.sql_cleanup_interval_seconds <= 0: + raise ValueError("sql_cleanup_interval_seconds must be greater than zero") + _require_positive( + memory_index_max_lines=self.memory_index_max_lines, + memory_index_max_bytes=self.memory_index_max_bytes, + preload_memory_max_topics=self.preload_memory_max_topics, + preload_memory_max_chars=self.preload_memory_max_chars, + preload_memory_candidate_limit=self.preload_memory_candidate_limit, + ) + _validate_path_components((self.memory_dir_name, self.memory_index_name)) + object.__setattr__(self, "root_dir", self.root_dir.expanduser().resolve()) diff --git a/trpc_agent_sdk/memory/advanced_memory/_formats.py b/trpc_agent_sdk/memory/advanced_memory/_formats.py new file mode 100644 index 000000000..bea5c3d35 --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_formats.py @@ -0,0 +1,136 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Data formats used by Advanced Memory.""" + +from __future__ import annotations + +import re +from dataclasses import dataclass +from datetime import datetime +from datetime import timezone +from enum import Enum + +_FRONTMATTER_PATTERN = re.compile(r"\A---\n(?P.*?)\n---(?:\n|\Z)", re.DOTALL) +_UPDATED_AT_PATTERN = re.compile(r"^updated_at:\s*(?P\S+)\s*$", re.MULTILINE) + + +def _as_utc(value: datetime) -> datetime: + """Normalize an aware or naive datetime to UTC.""" + if value.tzinfo is None: + value = value.replace(tzinfo=timezone.utc) + return value.astimezone(timezone.utc) + + +class MemoryType(str, Enum): + """Semantic types allowed for long-term memory documents.""" + + USER = "user" + FEEDBACK = "feedback" + PROJECT = "project" + REFERENCE = "reference" + + +@dataclass(frozen=True) +class MemoryIndexEntry: + """Represent one entry in MEMORY.md.""" + + name: str + filename: str + summary: str + + def __post_init__(self) -> None: + """Validate that index fields are non-empty single-line strings.""" + for field_name, value in ( + ("name", self.name), + ("filename", self.filename), + ("summary", self.summary), + ): + if not value.strip() or "\n" in value or "\r" in value: + raise ValueError(f"{field_name} must be non-empty single-line text") + + def to_markdown(self) -> str: + """Render one standard index entry.""" + return f"- [{self.name.strip()}]({self.filename.strip()}):{self.summary.strip()}" + + +@dataclass(frozen=True) +class MemoryDocument: + """Represent one long-term memory topic.""" + + name: str + description: str + memory_type: MemoryType + content: str + updated_at: datetime | None = None + + def __post_init__(self) -> None: + """Validate frontmatter fields.""" + for field_name, value in ( + ("name", self.name), + ("description", self.description), + ): + if not value.strip() or "\n" in value or "\r" in value: + raise ValueError(f"{field_name} must be non-empty single-line text") + + def to_markdown(self) -> str: + """Render the topic as Markdown with frontmatter.""" + body = self.content.strip() + updated_at = _as_utc(self.updated_at).isoformat() if self.updated_at is not None else None + updated_at_line = f"updated_at: {updated_at}\n" if updated_at else "" + return ( + "---\n" + f"name: {self.name.strip()}\n" + f"description: {self.description.strip()}\n" + f"type: {self.memory_type.value}\n" + f"{updated_at_line}" + "---\n" + f"{body}\n" + ) + + +def parse_memory_updated_at(content: str) -> datetime | None: + """Extract the UTC update timestamp from a memory document.""" + frontmatter_match = _FRONTMATTER_PATTERN.match(content) + if frontmatter_match is None: + return None + match = _UPDATED_AT_PATTERN.search(frontmatter_match.group("frontmatter")) + if match is None: + return None + try: + parsed = datetime.fromisoformat(match.group("value").replace("Z", "+00:00")) + except ValueError: + return None + return _as_utc(parsed) + + +def memory_freshness(updated_at: datetime | None, *, now: datetime | None = None) -> str: + """Return a compact freshness bucket for model-facing output.""" + if updated_at is None: + return "unknown" + age_days = max(0, int((_as_utc(now or datetime.now(timezone.utc)) - _as_utc(updated_at)).total_seconds()) + // 86_400) + if age_days == 0: + return "today" + if age_days == 1: + return "yesterday" + if age_days <= 7: + return "within 7 days" + if age_days <= 30: + return "within 30 days" + return "over 30 days" + + +def limit_memory_index(index: str, *, max_lines: int, max_bytes: int, encoding: str) -> str: + """Return a bounded view of an index without modifying the stored index.""" + lines: list[str] = [] + used_bytes = 0 + for line in index.splitlines(keepends=True)[:max_lines]: + size = len(line.encode(encoding)) + if used_bytes + size > max_bytes: + break + lines.append(line) + used_bytes += size + return "".join(lines) diff --git a/trpc_agent_sdk/memory/advanced_memory/_integration.py b/trpc_agent_sdk/memory/advanced_memory/_integration.py new file mode 100644 index 000000000..3c01e02f6 --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_integration.py @@ -0,0 +1,106 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Provide setup entry points for long-term memory.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any +from typing import TYPE_CHECKING + +from ._runtime import AdvancedMemoryRuntime + +from ._memory_context import LongTermMemoryContext +from ._memory_context import setup_long_term_memory_context + +if TYPE_CHECKING: + from trpc_agent_sdk.agents import LlmAgent + from trpc_agent_sdk.tools._advanced_memory_tool import AdvancedMemoryTools + + +@dataclass(frozen=True) +class LongTermMemoryIntegration: + """Aggregate the long-term memory callback and tools.""" + + context: LongTermMemoryContext + tools: "AdvancedMemoryTools | None" + + +def _setup_long_term_memory_tools( + agent: "LlmAgent", + memory_runtime: AdvancedMemoryRuntime, +) -> "AdvancedMemoryTools": + """Install the three official memory tools idempotently.""" + from trpc_agent_sdk.tools._advanced_memory_tool import ( + ADVANCED_MEMORY_TOOL_NAMES, ) + from trpc_agent_sdk.tools._advanced_memory_tool import AdvancedMemoryTools + + matching_tools = [tool for tool in agent.tools if getattr(tool, "name", None) in ADVANCED_MEMORY_TOOL_NAMES] + if matching_tools: + owners = {getattr(getattr(tool, "func", None), "__self__", None) for tool in matching_tools} + if len(owners) != 1: + raise ValueError("Advanced Memory tool names are already used by different tools") + owner = owners.pop() + if not isinstance(owner, AdvancedMemoryTools): + raise ValueError("Advanced Memory tool names are already used by non-SDK tools") + if owner.runtime is not memory_runtime: + raise ValueError("Advanced Memory tools use another runtime") + installed_names = {getattr(tool, "name", None) for tool in matching_tools} + if installed_names != ADVANCED_MEMORY_TOOL_NAMES: + raise ValueError("Advanced Memory tools are only partially installed") + return owner + tools = AdvancedMemoryTools(memory_runtime) + agent.tools.extend(tools.as_tools()) + return tools + + +def _setup_preload_memory_tool( + agent: "LlmAgent", + memory_runtime: AdvancedMemoryRuntime, + model: Any | None = None, +) -> None: + """Install the automatic topic-memory preprocessor when enabled.""" + if (not memory_runtime.config.enabled or not memory_runtime.config.preload_memory_enabled): + return + from trpc_agent_sdk.tools import PreloadMemoryTool + + from ._preload_memory import MemoryPreloader + from ._preload_memory import ModelMemoryRelevanceSelector + + existing = [tool for tool in agent.tools if getattr(tool, "name", None) == "preload_memory"] + use_legacy_memory = False + if existing: + if len(existing) != 1 or not isinstance(existing[0], PreloadMemoryTool): + raise ValueError("Advanced Memory preload tool name is already used by another tool") + use_legacy_memory = existing[0].uses_legacy_memory + agent.tools.remove(existing[0]) + preloader = MemoryPreloader( + memory_runtime, + ModelMemoryRelevanceSelector(model), + ) + agent.tools.append(PreloadMemoryTool( + memory_preloader=preloader.preload, + use_legacy_memory=use_legacy_memory, + )) + + +def setup_long_term_memory( + agent: "LlmAgent", + memory_runtime: AdvancedMemoryRuntime, + *, + preload_memory_model: Any | None = None, + install_tools: bool = True, +) -> LongTermMemoryIntegration: + """Install only user-scoped long-term memory behavior.""" + context = setup_long_term_memory_context(agent, memory_runtime) + tools = (_setup_long_term_memory_tools(agent, memory_runtime) + if install_tools and memory_runtime.config.enabled else None) + _setup_preload_memory_tool( + agent, + memory_runtime, + model=preload_memory_model, + ) + return LongTermMemoryIntegration(context=context, tools=tools) diff --git a/trpc_agent_sdk/advanced_memory/_memory_context.py b/trpc_agent_sdk/memory/advanced_memory/_memory_context.py similarity index 81% rename from trpc_agent_sdk/advanced_memory/_memory_context.py rename to trpc_agent_sdk/memory/advanced_memory/_memory_context.py index 947db2330..b3e62a9df 100644 --- a/trpc_agent_sdk/advanced_memory/_memory_context.py +++ b/trpc_agent_sdk/memory/advanced_memory/_memory_context.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING -from ._callbacks import install_staged_callback +from trpc_agent_sdk.sessions.compact._callbacks import install_staged_callback from ._runtime import AdvancedMemoryRuntime if TYPE_CHECKING: @@ -32,17 +32,24 @@ def runtime(self) -> AdvancedMemoryRuntime: """Return the runtime bound to this long-term memory context.""" return self._runtime - async def apply(self, request: "LlmRequest") -> bool: + async def apply(self, request: "LlmRequest", ctx: "InvocationContext | None" = None) -> bool: """Append the MEMORY.md index and on-demand read guidance.""" - config = self._runtime.config + runtime = self._runtime.for_session(ctx.session) if ctx is not None else self._runtime + config = runtime.config if not config.enabled or not config.long_term_memory_injection_enabled: return False - await self._runtime.initialize() + await runtime.initialize() existing_instruction = (str(request.config.system_instruction) if request.config is not None and request.config.system_instruction else "") if LONG_TERM_MEMORY_MARKER in existing_instruction: return False - index = await self._runtime.long_term_memory.read_index() + index = await runtime.long_term_memory.read_index() + focus_instruction = (config.memory_focus_instruction or "").strip() + custom_focus = ("\n\n## Custom memory focus\n" + "The following is an additional application-level memory preference. " + "Give it extra attention when deciding whether stable, explicit information " + "is worth saving, while still following the safety and quality rules above:\n" + f"{focus_instruction}\n" if focus_instruction else "") instruction = ( f"{LONG_TERM_MEMORY_MARKER}\n" "The following is a bounded index of this project's long-term memory. It is a trusted cross-session " @@ -66,13 +73,15 @@ async def apply(self, request: "LlmRequest") -> bool: "Do not save temporary task details, information reconstructable from current code, unverified guesses, " "duplicates, the model's own reasoning, or secrets, credentials, tokens, and other sensitive data. " "Do not write information that is uncertain, useful only in the current conversation, or not clearly " - "worth preserving.\n\n" + f"worth preserving.{custom_focus}\n\n" "save_memory writes both the detail file and the index. Pass a stable filename and concise " "name/description/summary, and use one of user, feedback, project, or reference for memory_type. " "Keep the description short and general; put detailed information in content. " "If save_memory is unavailable, do not claim that the information was saved.\n" - f"Memory directory: {self._runtime.paths.memory_dir}\n" - f"Index file: {self._runtime.paths.memory_index_path}\n" + f"Memory directory: " + f"{runtime.paths.memory_dir if config.storage_backend == 'local' else config.storage_backend.upper()}\n" + f"Index file: " + f"{runtime.paths.storage_reference('memory_index')}\n" f"\n{index.rstrip()}\n\n" f"") request.append_instructions([instruction]) @@ -95,8 +104,7 @@ def memory_context(self) -> LongTermMemoryContext: async def __call__(self, ctx: "InvocationContext", request: "LlmRequest") -> None: """Inject the long-term memory index before a model request.""" - del ctx - await self._memory_context.apply(request) + await self._memory_context.apply(request, ctx) return None diff --git a/trpc_agent_sdk/memory/advanced_memory/_paths.py b/trpc_agent_sdk/memory/advanced_memory/_paths.py new file mode 100644 index 000000000..768c86ff4 --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_paths.py @@ -0,0 +1,110 @@ +# Tencent is pleased to support the open source ecosystem. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# Licensed under Apache-2.0. +"""Safe path resolution for long-term Advanced Memory.""" + +from __future__ import annotations + +import hashlib +import re +from dataclasses import dataclass +from pathlib import Path + +from ._config import AdvancedMemoryServiceConfig + +_SAFE_COMPONENT_PATTERN = re.compile(r"[^A-Za-z0-9._-]+") + + +def _safe_component(value: str, *, field_name: str) -> str: + if value != value.strip() or any(ord(character) < 32 for character in value): + raise ValueError(f"{field_name} must not contain surrounding or control whitespace") + normalized = _SAFE_COMPONENT_PATTERN.sub("_", value.strip()).strip("._") + if not normalized: + raise ValueError(f"{field_name} must contain at least one safe character") + return normalized + + +def _collision_safe_component(value: str, *, field_name: str) -> str: + stripped = value.strip() + normalized = _safe_component(stripped, field_name=field_name) + if normalized == stripped: + return normalized + digest = hashlib.sha256(stripped.encode("utf-8")).hexdigest()[:12] + return f"{normalized}-{digest}" + + +@dataclass(frozen=True) +class MemoryScope: + """Identify the application and user that own memory.""" + + app_name: str + user_id: str + + def __post_init__(self) -> None: + _safe_component(self.app_name, field_name="app_name") + _safe_component(self.user_id, field_name="user_id") + + @property + def storage_key(self) -> str: + return repr((self.app_name, self.user_id)) + + +@dataclass(frozen=True) +class AdvancedMemoryPaths: + """Build paths for long-term memory only.""" + + config: AdvancedMemoryServiceConfig + scope: MemoryScope | None = None + + def for_scope(self, app_name: str, user_id: str) -> "AdvancedMemoryPaths": + return AdvancedMemoryPaths(self.config, MemoryScope(app_name, user_id)) + + @property + def tenant_root_dir(self) -> Path: + if self.scope is None: + return self.config.root_dir + return (self.config.root_dir / "tenants" / + _collision_safe_component(self.scope.app_name, field_name="app_name") / + _collision_safe_component(self.scope.user_id, field_name="user_id")) + + @property + def scope_key(self) -> str: + return self.scope.storage_key if self.scope is not None else "legacy\0global" + + @property + def memory_dir(self) -> Path: + return self.tenant_root_dir / self.config.memory_dir_name + + @property + def memory_index_path(self) -> Path: + return self.memory_dir / self.config.memory_index_name + + def memory_topic_path(self, topic_name: str) -> Path: + safe_name = _collision_safe_component(topic_name, field_name="topic_name") + if not safe_name.lower().endswith(".md"): + safe_name = f"{safe_name}.md" + if safe_name == self.config.memory_index_name: + raise ValueError("Topic file cannot overwrite the memory index") + return self.memory_dir / safe_name + + def storage_reference(self, resource: str, *, topic_name: str | None = None) -> str: + if resource == "memory_index": + path = self.memory_index_path + elif resource == "memory_topic" and topic_name is not None: + path = self.memory_topic_path(topic_name) + else: + raise ValueError(f"Unknown long-term memory resource: {resource}") + if self.config.storage_backend == "local": + return str(path) + if self.scope is None: + raise ValueError("A scoped path is required for non-local memory storage") + app = _collision_safe_component(self.scope.app_name, field_name="app_name") + user = _collision_safe_component(self.scope.user_id, field_name="user_id") + if self.config.storage_backend == "redis": + key = f"{self.config.redis_key_prefix}:{{{app}:{user}}}:memory:{path.name}" + return f"advanced-memory://redis/{key}" + return f"advanced-memory://sql/{app}/{user}/memory/{path.name}" + + def ensure_base_directories(self) -> None: + self.memory_dir.mkdir(parents=True, exist_ok=True) diff --git a/trpc_agent_sdk/advanced_memory/_preload_memory.py b/trpc_agent_sdk/memory/advanced_memory/_preload_memory.py similarity index 76% rename from trpc_agent_sdk/advanced_memory/_preload_memory.py rename to trpc_agent_sdk/memory/advanced_memory/_preload_memory.py index 4f7f92eee..f2bcfce8b 100644 --- a/trpc_agent_sdk/advanced_memory/_preload_memory.py +++ b/trpc_agent_sdk/memory/advanced_memory/_preload_memory.py @@ -16,18 +16,14 @@ from typing import Protocol from typing import TYPE_CHECKING -from trpc_agent_sdk.agents import LlmAgent from trpc_agent_sdk.log import logger -from trpc_agent_sdk.memory import InMemoryMemoryService -from trpc_agent_sdk.runners import Runner -from trpc_agent_sdk.sessions import InMemorySessionService +from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.memory.advanced_memory._formats import memory_freshness +from trpc_agent_sdk.memory.advanced_memory._formats import parse_memory_updated_at +from ._runtime import AdvancedMemoryRuntime from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import Part -from ._formats import memory_freshness -from ._formats import parse_memory_updated_at -from ._runtime import AdvancedMemoryRuntime - if TYPE_CHECKING: from trpc_agent_sdk.context import InvocationContext @@ -92,15 +88,20 @@ def _candidate_from_content(filename: str, content: str) -> MemoryCandidate: class ModelMemoryRelevanceSelector: - """Use an isolated lightweight Agent to select relevant topic files.""" + """Use one direct LLM call to select relevant topic files.""" def __init__(self, model: object | None = None) -> None: """Store an optional dedicated selector model.""" self._model = model - def _resolve_model(self, ctx: "InvocationContext") -> object: - """Prefer a dedicated selector model and fall back to the main model.""" - model = self._model if self._model is not None else getattr(ctx.agent, "model", None) + async def _resolve_model(self, ctx: "InvocationContext") -> object: + """Prefer a dedicated selector model and resolve the main Agent model.""" + if self._model is not None: + return self._model + resolver = getattr(ctx.agent, "_resolve_model", None) + if callable(resolver): + return await resolver(ctx) + model = getattr(ctx.agent, "model", None) if model is None: raise ValueError("Memory relevance selector cannot resolve an LLM model") return model @@ -158,48 +159,32 @@ async def select( *, limit: int, ) -> list[str]: - """Run the isolated selector Agent and validate its result.""" - app_name = f"{ctx.app_name}_advanced_memory_selector" - agent = LlmAgent( - name="advanced_memory_relevance_selector", - description="Select relevant long-term memories.", - instruction=("You are a strict long-term memory relevance selector. " - "Follow the user's query and output format exactly."), - model=self._resolve_model(ctx), - tools=[], - add_name_to_instruction=False, - ) - runner = Runner( - app_name=app_name, - agent=agent, - session_service=InMemorySessionService(), - memory_service=InMemoryMemoryService(), - enable_post_turn_processing=False, + """Run one direct LLM call and validate its result.""" + model = await self._resolve_model(ctx) + generate_async = getattr(model, "generate_async", None) + if not callable(generate_async): + raise TypeError("Memory relevance selector requires an LLMModel instance") + model_name = getattr(model, "name", None) + if not isinstance(model_name, str) or not model_name: + raise ValueError("Memory relevance selector model has no valid name") + request = LlmRequest( + model=model_name, + contents=[ + Content( + role="user", + parts=[Part.from_text(text=self._build_prompt(query, candidates, limit))], + ) + ], ) - try: - session = await runner.session_service.create_session( - app_name=app_name, - user_id="advanced-memory-selector", - state={}, - ) - content = Content( - role="user", - parts=[Part.from_text(text=self._build_prompt(query, candidates, limit))], - ) - last_event = None - async for event in runner.run_async( - user_id=session.user_id, - session_id=session.id, - new_message=content, - ): - if not event.partial: - last_event = event - if not last_event or not last_event.content or not last_event.content.parts: - raise ValueError("Memory relevance selector returned no final content") - text = "\n".join(part.text for part in last_event.content.parts if part.text) - return self._parse_selection(text, candidates, limit) - finally: - await runner.close() + response_text: list[str] = [] + async for response in generate_async(request, stream=False, ctx=None): + if response.error_code: + raise ValueError(response.error_message or "Memory relevance selector failed") + if response.content and response.content.parts: + response_text.extend(part.text for part in response.content.parts if part.text) + if not response_text: + raise ValueError("Memory relevance selector returned no final content") + return self._parse_selection("\n".join(response_text), candidates, limit) async def select_relevant_memory_filenames( @@ -228,18 +213,19 @@ def __init__( self._runtime = runtime self._selector = selector or ModelMemoryRelevanceSelector() - async def _candidates(self) -> list[MemoryCandidate]: + async def _candidates(self, ctx: "InvocationContext") -> list[MemoryCandidate]: """Read and sort bounded topic metadata for selection.""" + runtime = self._runtime.for_session(ctx.session) candidates: list[MemoryCandidate] = [] - for path in await self._runtime.long_term_memory.list_topics(): - frontmatter = await self._runtime.long_term_memory.read_topic_frontmatter(path.name) + for path in await runtime.long_term_memory.list_topics(): + frontmatter = await runtime.long_term_memory.read_topic_frontmatter(path.name) if frontmatter is not None: candidates.append(_candidate_from_content(path.name, frontmatter)) candidates.sort( key=lambda candidate: candidate.updated_at or datetime.min.replace(tzinfo=timezone.utc), reverse=True, ) - return candidates[:self._runtime.config.preload_memory_candidate_limit] + return candidates[:runtime.config.preload_memory_candidate_limit] async def preload(self, query: str, ctx: "InvocationContext") -> str | None: """Select and render relevant topic bodies within the configured budget.""" @@ -247,7 +233,7 @@ async def preload(self, query: str, ctx: "InvocationContext") -> str | None: if not config.enabled or not config.preload_memory_enabled or not query.strip(): return None try: - candidates = await self._candidates() + candidates = await self._candidates(ctx) except Exception as exc: # noqa: BLE001 logger.warning("Advanced Memory preload candidate loading failed: %s", exc) return None @@ -272,7 +258,7 @@ async def preload(self, query: str, ctx: "InvocationContext") -> str | None: if candidate is None: continue try: - full_content = await self._runtime.long_term_memory.read_topic(filename) + full_content = await self._runtime.for_session(ctx.session).long_term_memory.read_topic(filename) except Exception as exc: # noqa: BLE001 logger.warning("Advanced Memory preload topic loading failed for %s: %s", filename, exc) continue diff --git a/trpc_agent_sdk/memory/advanced_memory/_redis_stores.py b/trpc_agent_sdk/memory/advanced_memory/_redis_stores.py new file mode 100644 index 000000000..8e0dba64a --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_redis_stores.py @@ -0,0 +1,197 @@ +"""Redis implementations of the Advanced Memory storage contracts.""" + +from __future__ import annotations + +import asyncio +from contextlib import asynccontextmanager +from dataclasses import replace +from datetime import datetime, timezone +from pathlib import Path +from typing import Any +from uuid import uuid4 + +from trpc_agent_sdk.storage import RedisCommand, RedisExpire, RedisStorage +from trpc_agent_sdk.types import Ttl + +from ._config import AdvancedMemoryServiceConfig +from ._formats import MemoryDocument, MemoryIndexEntry +from ._formats import limit_memory_index +from ._paths import AdvancedMemoryPaths +from ._storage import parse_memory_index, prune_memory_index + +_RELEASE_LOCK_SCRIPT = """ +if redis.call('GET', KEYS[1]) == ARGV[1] then + return redis.call('DEL', KEYS[1]) +end +return 0 +""" + + +class _RedisStore: + + def __init__( + self, + config: AdvancedMemoryServiceConfig, + paths: AdvancedMemoryPaths, + storage: RedisStorage, + ) -> None: + if paths.scope is None: + raise ValueError("Redis Advanced Memory storage requires a tenant scope") + self._config, self._paths, self._storage = config, paths, storage + app_component = paths.tenant_root_dir.parent.name + user_component = paths.tenant_root_dir.name + self._user_base = f"{config.redis_key_prefix}:{{{app_component}:{user_component}}}" + + async def _command(self, method: str, *args: Any, **kwargs: Any) -> Any: + command_expire = kwargs.pop("_command_expire", None) + async with self._storage.create_db_session() as connection: + return await self._storage.execute_command( + connection, + RedisCommand(method=method, args=args, kwargs=kwargs, expire=command_expire or RedisExpire()), + ) + + def _memory_registry(self) -> str: + return f"{self._user_base}:memory:keys" + + def _memory_lock_key(self) -> str: + """Return the distributed lock key for this app/user memory scope.""" + return f"{self._user_base}:memory:lock" + + @asynccontextmanager + async def _memory_write_lock(self): + """Serialize long-term memory writes across processes and nodes.""" + token = uuid4().hex + key = self._memory_lock_key() + deadline = asyncio.get_running_loop().time() + self._config.memory_lock_acquire_timeout_seconds + acquired = False + while asyncio.get_running_loop().time() < deadline: + result = await self._command( + "set", + key, + token, + nx=True, + ex=self._config.memory_lock_ttl_seconds, + _command_expire=RedisExpire( + key=key, + ttl=Ttl(ttl_seconds=self._config.memory_lock_ttl_seconds), + ), + ) + if result is True or result in (b"OK", "OK"): + acquired = True + break + await asyncio.sleep(min(0.05, max(0.0, deadline - asyncio.get_running_loop().time()))) + if not acquired: + raise TimeoutError(f"Timed out acquiring Advanced Memory lock for {self._paths.scope.storage_key}") + try: + yield + finally: + await self._command( + "eval", + _RELEASE_LOCK_SCRIPT, + 1, + key, + token, + ) + + async def _refresh_ttl_group( + self, + registry: str, + keys: list[str], + ttl: int | None, + skip_prefixes: tuple[str, ...] = (), + ) -> None: + """Track and refresh every key in one logical memory group.""" + if ttl is None: + return + if keys: + await self._command("sadd", registry, *keys) + tracked = await self._command("smembers", registry) or [] + tracked_keys = {self._text(value) for value in tracked} + tracked_keys.update(keys) + for key in tracked_keys: + if key and not key.startswith(skip_prefixes): + await self._command("expire", key, ttl) + await self._command("expire", registry, ttl) + + async def _refresh_memory_ttl(self, *keys: str) -> None: + await self._refresh_ttl_group( + self._memory_registry(), + list(keys), + self._config.memory_ttl_seconds, + ) + + @staticmethod + def _text(value: Any) -> str | None: + if value is None: + return None + return value.decode("utf-8") if isinstance(value, bytes) else str(value) + + +class RedisLongTermMemoryStore(_RedisStore): + + async def initialize(self) -> None: + key = f"{self._user_base}:memory:index" + await self._command("setnx", key, "") + await self._refresh_memory_ttl(key) + + async def read_index(self) -> str: + key = f"{self._user_base}:memory:index" + value = self._text(await self._command("get", key)) or "" + await self._refresh_memory_ttl() + index = value + valid_filenames = set() + for entry in parse_memory_index(index): + topic_key = f"{self._user_base}:memory:topic:{self._topic_name(entry.filename)}" + if await self._command("exists", topic_key): + valid_filenames.add(entry.filename) + pruned_index = prune_memory_index(index, valid_filenames) + if pruned_index != index: + async with self._memory_write_lock(): + await self._command("set", key, pruned_index) + await self._refresh_memory_ttl(key) + return limit_memory_index( + pruned_index, + max_lines=self._config.memory_index_max_lines, + max_bytes=self._config.memory_index_max_bytes, + encoding=self._config.encoding, + ) + + async def write_index(self, entries: list[MemoryIndexEntry]) -> None: + content = "\n".join(entry.to_markdown() for entry in entries) + key = f"{self._user_base}:memory:index" + async with self._memory_write_lock(): + await self._command("set", key, f"{content}\n" if content else "") + await self._refresh_memory_ttl(key) + + def _topic_name(self, topic_name: str) -> str: + return self._paths.memory_topic_path(topic_name).name + + async def read_topic(self, topic_name: str) -> str | None: + key = f"{self._user_base}:memory:topic:{self._topic_name(topic_name)}" + value = await self._command("get", key) + await self._refresh_memory_ttl() + return self._text(value) + + async def read_topic_frontmatter(self, topic_name: str) -> str | None: + content = await self.read_topic(topic_name) + if content is None: + return None + end = content.find("\n---", 4) if content.startswith("---\n") else -1 + return content[:end + 4] if end >= 0 else content + + async def write_topic(self, topic_name: str, document: MemoryDocument) -> Path: + name = self._topic_name(topic_name) + document = replace(document, updated_at=datetime.now(timezone.utc)) + topic_key = f"{self._user_base}:memory:topic:{name}" + topics_key = f"{self._user_base}:memory:topics" + async with self._memory_write_lock(): + await self._command("set", topic_key, document.to_markdown()) + await self._command("zadd", topics_key, {name: document.updated_at.timestamp()}) + await self._refresh_memory_ttl(topic_key, topics_key) + return Path(name) + + async def list_topics(self) -> list[Path]: + key = f"{self._user_base}:memory:topics" + values = await self._command("zrange", key, 0, -1) + await self._refresh_memory_ttl() + return [Path(self._text(value) or "") for value in values] diff --git a/trpc_agent_sdk/memory/advanced_memory/_runtime.py b/trpc_agent_sdk/memory/advanced_memory/_runtime.py new file mode 100644 index 000000000..c3eb60d1e --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_runtime.py @@ -0,0 +1,202 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Unified runtime entry point for the independent memory mechanism.""" + +from __future__ import annotations + +from dataclasses import dataclass +from dataclasses import field +import shutil +import threading +from typing import Any + +from ._config import AdvancedMemoryServiceConfig +from trpc_agent_sdk.sessions.compact.advanced._coordination import CrossLoopLock +from ._paths import AdvancedMemoryPaths +from ._paths import MemoryScope +from ._storage import LocalAdvancedMemoryCleanup +from ._storage import LongTermMemoryStore + + +@dataclass(frozen=True) +class AdvancedMemoryRuntime: + """Aggregate configuration, paths, and long-term memory storage.""" + + config: AdvancedMemoryServiceConfig + paths: AdvancedMemoryPaths + long_term_memory: LongTermMemoryStore + _scoped_runtimes: dict[MemoryScope, "ScopedAdvancedMemoryRuntime"] = field( + default_factory=dict, + repr=False, + compare=False, + ) + _scoped_runtimes_lock: threading.Lock = field( + default_factory=threading.Lock, + repr=False, + compare=False, + ) + _redis_storage: Any | None = field(default=None, repr=False, compare=False) + _sql_storage: Any | None = field(default=None, repr=False, compare=False) + _sql_cleanup: Any | None = field(default=None, repr=False, compare=False) + _local_cleanup: LocalAdvancedMemoryCleanup | None = field(default=None, repr=False, compare=False) + _close_lock: CrossLoopLock = field( + default_factory=CrossLoopLock, + repr=False, + compare=False, + ) + _closed: bool = field(default=False, repr=False, compare=False) + + @classmethod + def create(cls, config: AdvancedMemoryServiceConfig | None = None) -> "AdvancedMemoryRuntime": + """Create a runtime isolated from the legacy mechanism.""" + resolved_config = config or AdvancedMemoryServiceConfig() + paths = AdvancedMemoryPaths(resolved_config) + redis_storage = None + sql_storage = None + sql_cleanup = None + local_cleanup = None + if resolved_config.storage_backend == "redis": + from trpc_agent_sdk.storage import RedisStorage + redis_storage = RedisStorage(redis_url=resolved_config.redis_url, is_async=resolved_config.redis_is_async) + elif resolved_config.storage_backend == "sql": + from trpc_agent_sdk.storage import SqlStorage + from ._sql_stores import AdvancedMemorySqlBase, SqlAdvancedMemoryCleanup + sql_storage = SqlStorage( + is_async=resolved_config.sql_is_async, + db_url=resolved_config.sql_url, + metadata=AdvancedMemorySqlBase.metadata, + expire_on_commit=False, + ) + sql_cleanup = SqlAdvancedMemoryCleanup(resolved_config, sql_storage) + else: + local_cleanup = LocalAdvancedMemoryCleanup(resolved_config) + return cls( + config=resolved_config, + paths=paths, + long_term_memory=LongTermMemoryStore(resolved_config, paths), + _redis_storage=redis_storage, + _sql_storage=sql_storage, + _sql_cleanup=sql_cleanup, + _local_cleanup=local_cleanup, + ) + + def for_scope(self, app_name: str, user_id: str) -> "ScopedAdvancedMemoryRuntime": + """Return the stores isolated to one application user.""" + scope = MemoryScope(app_name, user_id) + with self._scoped_runtimes_lock: + runtime = self._scoped_runtimes.get(scope) + if runtime is None: + paths = self.paths.for_scope(app_name, user_id) + if self.config.storage_backend == "redis": + from trpc_agent_sdk.storage import RedisStorage + from ._redis_stores import RedisLongTermMemoryStore + + storage = self._redis_storage or RedisStorage( + redis_url=self.config.redis_url, + is_async=self.config.redis_is_async, + ) + long_term_memory = RedisLongTermMemoryStore(self.config, paths, storage) + elif self.config.storage_backend == "sql": + from ._sql_stores import SqlLongTermMemoryStore + storage = self._sql_storage + if storage is None: + raise RuntimeError("SQL Advanced Memory storage is not initialized") + long_term_memory = SqlLongTermMemoryStore(self.config, paths, storage) + else: + long_term_memory = LongTermMemoryStore(self.config, paths) + runtime = ScopedAdvancedMemoryRuntime( + root=self, + scope=scope, + paths=paths, + long_term_memory=long_term_memory, + ) + self._scoped_runtimes[scope] = runtime + return runtime + + def for_session(self, session: object) -> "ScopedAdvancedMemoryRuntime": + """Return the scoped runtime for a SessionABC-compatible object.""" + app_name = getattr(session, "app_name", None) + user_id = getattr(session, "user_id", None) + if not isinstance(app_name, str) or not isinstance(user_id, str): + raise ValueError("Advanced Memory requires session app_name and user_id") + return self.for_scope(app_name, user_id) + + def migrate_legacy(self, app_name: str, user_id: str) -> "ScopedAdvancedMemoryRuntime": + """Move an old flat Advanced Memory layout into one explicit tenant. + + Refuses to overwrite a tenant that already contains data. + """ + scoped = self.for_scope(app_name, user_id) + legacy_paths = self.paths + target_root = scoped.paths.tenant_root_dir + if target_root.exists(): + raise FileExistsError(f"Target Advanced Memory tenant already exists: {target_root}") + if not legacy_paths.memory_dir.exists(): + raise FileNotFoundError("No legacy Advanced Memory directories exist") + target_root.mkdir(parents=True) + if legacy_paths.memory_dir.exists(): + shutil.move(str(legacy_paths.memory_dir), str(scoped.paths.memory_dir)) + return scoped + + async def initialize(self) -> bool: + """Create memory directories only when the mechanism is enabled.""" + if not self.config.enabled: + return False + if self.config.storage_backend == "sql": + if self._sql_storage is None: + raise RuntimeError("SQL Advanced Memory storage is not initialized") + if self._sql_cleanup is not None: + await self._sql_cleanup.start() + async with self._sql_storage.create_db_session(): + pass + return True + if self.config.storage_backend == "redis": + return True + if self._local_cleanup is not None: + await self._local_cleanup.start() + await self.long_term_memory.initialize() + return True + + async def close(self) -> None: + """Release shared external backend resources.""" + async with self._close_lock: + if self._closed: + return + if self._local_cleanup is not None: + await self._local_cleanup.close() + if self._redis_storage is not None: + await self._redis_storage.close() + if self._sql_cleanup is not None: + await self._sql_cleanup.close() + if self._sql_storage is not None: + await self._sql_storage.close() + object.__setattr__(self, "_closed", True) + + +@dataclass(frozen=True) +class ScopedAdvancedMemoryRuntime: + """A tenant-bound view of an :class:`AdvancedMemoryRuntime`.""" + + root: AdvancedMemoryRuntime + scope: MemoryScope + paths: AdvancedMemoryPaths + long_term_memory: LongTermMemoryStore + + @property + def config(self) -> AdvancedMemoryServiceConfig: + """Return the root runtime configuration.""" + return self.root.config + + async def initialize(self) -> bool: + """Initialize only this tenant's local directories.""" + if not self.config.enabled: + return False + if self.config.storage_backend == "local" and self.root._local_cleanup is not None: + await self.root._local_cleanup.start() + if self.config.storage_backend == "sql" and self.root._sql_cleanup is not None: + await self.root._sql_cleanup.start() + await self.long_term_memory.initialize() + return True diff --git a/trpc_agent_sdk/memory/advanced_memory/_sql_stores.py b/trpc_agent_sdk/memory/advanced_memory/_sql_stores.py new file mode 100644 index 000000000..05874e849 --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_sql_stores.py @@ -0,0 +1,315 @@ +"""SQL implementations of the Advanced Memory storage contracts.""" + +from __future__ import annotations + +import asyncio +from datetime import datetime, timedelta, timezone +from dataclasses import replace +from pathlib import Path +from typing import Any + +from sqlalchemy import DateTime, String, Text, func +from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column + +from trpc_agent_sdk.storage import ( + DEFAULT_MAX_KEY_LENGTH, + DEFAULT_MAX_VARCHAR_LENGTH, + PreciseTimestamp, + SqlCondition, + SqlKey, + SqlStorage, +) + +from ._config import AdvancedMemoryServiceConfig +from ._formats import MemoryDocument, MemoryIndexEntry +from ._formats import limit_memory_index +from ._paths import AdvancedMemoryPaths +from ._storage import prune_memory_index + + +class AdvancedMemorySqlBase(DeclarativeBase): + """Metadata owned exclusively by Advanced Memory SQL stores.""" + + +class SqlMemoryIndex(AdvancedMemorySqlBase): + __tablename__ = "advanced_memory_indexes" + + app_name: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + user_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + content: Mapped[str] = mapped_column(Text, default="") + updated_at: Mapped[datetime] = mapped_column(PreciseTimestamp, default=func.now(), onupdate=func.now()) + expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) + + +class SqlMemoryTopic(AdvancedMemorySqlBase): + __tablename__ = "advanced_memory_topics" + + app_name: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + user_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + topic_name: Mapped[str] = mapped_column(String(DEFAULT_MAX_VARCHAR_LENGTH), primary_key=True) + content: Mapped[str] = mapped_column(Text) + updated_at: Mapped[datetime] = mapped_column(PreciseTimestamp, default=func.now(), onupdate=func.now()) + expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) + + +class _SqlStore: + + def __init__( + self, + config: AdvancedMemoryServiceConfig, + paths: AdvancedMemoryPaths, + storage: SqlStorage, + ) -> None: + if paths.scope is None: + raise ValueError("SQL Advanced Memory storage requires a tenant scope") + self._config = config + self._paths = paths + self._storage = storage + self._app_name = paths.scope.app_name + self._user_id = paths.scope.user_id + + @staticmethod + def _now() -> datetime: + return datetime.now(timezone.utc).replace(tzinfo=None) + + def _expiry(self, ttl: int | None) -> datetime | None: + return self._now() + timedelta(seconds=ttl) if ttl is not None else None + + @staticmethod + def _expired(value: datetime | None) -> bool: + if value is None: + return False + return value.replace(tzinfo=None) <= datetime.now(timezone.utc).replace(tzinfo=None) + + async def initialize(self) -> None: + async with self._storage.create_db_session(): + pass + + async def _refresh_memory_scope(self, db: Any) -> None: + expiry = self._expiry(self._config.memory_ttl_seconds) + if expiry is None: + return + index = await self._storage.get(db, SqlKey( + key=(self._app_name, self._user_id), + storage_cls=SqlMemoryIndex, + )) + if index is not None: + index.expires_at = expiry + topics = await self._storage.query( + db, + SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryTopic), + SqlCondition(filters=[ + SqlMemoryTopic.app_name == self._app_name, + SqlMemoryTopic.user_id == self._user_id, + SqlMemoryTopic.expires_at.is_(None) | (SqlMemoryTopic.expires_at > self._now()), + ]), + ) + for topic in topics: + topic.expires_at = expiry + + +class SqlLongTermMemoryStore(_SqlStore): + + async def initialize(self) -> None: + await super().initialize() + async with self._storage.create_db_session() as db: + key = SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryIndex) + row = await self._storage.get(db, key) + if row is None: + await self._storage.add( + db, + SqlMemoryIndex( + app_name=self._app_name, + user_id=self._user_id, + content="", + expires_at=self._expiry(self._config.memory_ttl_seconds), + )) + await self._storage.commit(db) + + async def read_index(self) -> str: + async with self._storage.create_db_session() as db: + row = await self._storage.get(db, SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryIndex)) + if row is None or self._expired(row.expires_at): + return "" + await self._refresh_memory_scope(db) + content = row.content + valid_topics = await self._storage.query( + db, + SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryTopic), + SqlCondition(filters=[ + SqlMemoryTopic.app_name == self._app_name, + SqlMemoryTopic.user_id == self._user_id, + SqlMemoryTopic.expires_at.is_(None) | (SqlMemoryTopic.expires_at > self._now()), + ]), + ) + valid_filenames = {topic.topic_name for topic in valid_topics} + pruned_content = prune_memory_index(content, valid_filenames) + if pruned_content != content: + row.content = pruned_content + content = pruned_content + await self._storage.commit(db) + return limit_memory_index( + content, + max_lines=self._config.memory_index_max_lines, + max_bytes=self._config.memory_index_max_bytes, + encoding=self._config.encoding, + ) + + async def write_index(self, entries: list[MemoryIndexEntry]) -> None: + content = "\n".join(entry.to_markdown() for entry in entries) + if content: + content += "\n" + async with self._storage.create_db_session() as db: + # Keep the tenant's lock row locked until this transaction commits. + await self._storage.get_for_update( + db, + SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryIndex), + ) + key = SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryIndex) + row = await self._storage.get(db, key) + if row is None: + row = SqlMemoryIndex(app_name=self._app_name, user_id=self._user_id) + await self._storage.add(db, row) + row.content = content + row.expires_at = self._expiry(self._config.memory_ttl_seconds) + await self._refresh_memory_scope(db) + await self._storage.commit(db) + + def _topic_key(self, topic_name: str) -> tuple[str, str, str]: + return self._app_name, self._user_id, self._paths.memory_topic_path(topic_name).name + + async def read_topic(self, topic_name: str) -> str | None: + async with self._storage.create_db_session() as db: + row = await self._storage.get(db, SqlKey(key=self._topic_key(topic_name), storage_cls=SqlMemoryTopic)) + if row is None or self._expired(row.expires_at): + return None + await self._refresh_memory_scope(db) + await self._storage.commit(db) + return row.content + + async def read_topic_frontmatter(self, topic_name: str) -> str | None: + content = await self.read_topic(topic_name) + if content is None: + return None + end = content.find("\n---", 4) if content.startswith("---\n") else -1 + return content[:end + 4] if end >= 0 else content + + async def write_topic(self, topic_name: str, document: MemoryDocument) -> Path: + name = self._paths.memory_topic_path(topic_name).name + async with self._storage.create_db_session() as db: + # Serialize all long-term writes for this app/user scope. + await self._storage.get_for_update( + db, + SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryIndex), + ) + key = self._topic_key(name) + row = await self._storage.get(db, SqlKey(key=key, storage_cls=SqlMemoryTopic)) + if row is None: + row = SqlMemoryTopic(app_name=key[0], user_id=key[1], topic_name=key[2]) + await self._storage.add(db, row) + row.content = replace(document, updated_at=self._now().replace(tzinfo=timezone.utc)).to_markdown() + row.updated_at = self._now() + row.expires_at = self._expiry(self._config.memory_ttl_seconds) + await self._refresh_memory_scope(db) + await self._storage.commit(db) + return Path(name) + + async def list_topics(self) -> list[Path]: + async with self._storage.create_db_session() as db: + rows = await self._storage.query( + db, + SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryTopic), + SqlCondition(filters=[ + SqlMemoryTopic.app_name == self._app_name, + SqlMemoryTopic.user_id == self._user_id, + ]), + ) + rows = [row for row in rows if not self._expired(row.expires_at)] + await self._refresh_memory_scope(db) + await self._storage.commit(db) + return [Path(row.topic_name) for row in sorted(rows, key=lambda item: item.topic_name)] + + +class SqlAdvancedMemoryCleanup: + """Periodically remove expired Advanced Memory SQL rows.""" + + _models = ( + SqlMemoryIndex, + SqlMemoryTopic, + ) + + def __init__(self, config: AdvancedMemoryServiceConfig, storage: SqlStorage) -> None: + self._config = config + self._storage = storage + self._task: asyncio.Task[None] | None = None + self._stop_event: asyncio.Event | None = None + + async def start(self) -> None: + if self._task is not None or self._config.memory_ttl_seconds is None: + return + self._stop_event = asyncio.Event() + self._task = asyncio.create_task(self._run()) + + async def cleanup_once(self) -> None: + now = datetime.now(timezone.utc).replace(tzinfo=None) + async with self._storage.create_db_session() as db: + for model in self._models: + await self._storage.delete( + db, + SqlKey(key=tuple(), storage_cls=model), + SqlCondition(filters=[model.expires_at.is_not(None), model.expires_at <= now]), + ) + indexes = await self._storage.query( + db, + SqlKey(key=tuple(), storage_cls=SqlMemoryIndex), + ) + for index in indexes: + topics = await self._storage.query( + db, + SqlKey( + key=(index.app_name, index.user_id), + storage_cls=SqlMemoryTopic, + ), + SqlCondition(filters=[ + SqlMemoryTopic.app_name == index.app_name, + SqlMemoryTopic.user_id == index.user_id, + SqlMemoryTopic.expires_at.is_(None) | (SqlMemoryTopic.expires_at > now), + ]), + ) + valid_filenames = {topic.topic_name for topic in topics} + index.content = prune_memory_index(index.content, valid_filenames) + await self._storage.commit(db) + + async def _run(self) -> None: + if self._stop_event is None: + return + try: + while not self._stop_event.is_set(): + try: + await asyncio.wait_for( + self._stop_event.wait(), + timeout=self._config.sql_cleanup_interval_seconds, + ) + except asyncio.TimeoutError: + await self.cleanup_once() + except asyncio.CancelledError: + raise + + async def close(self) -> None: + if self._stop_event is not None: + self._stop_event.set() + if self._task is not None and not self._task.done(): + self._task.cancel() + try: + await self._task + except asyncio.CancelledError: + pass + self._task = None + self._stop_event = None + + +__all__ = [ + "AdvancedMemorySqlBase", + "SqlAdvancedMemoryCleanup", + "SqlLongTermMemoryStore", +] diff --git a/trpc_agent_sdk/memory/advanced_memory/_storage.py b/trpc_agent_sdk/memory/advanced_memory/_storage.py new file mode 100644 index 000000000..d5a0e9f08 --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_storage.py @@ -0,0 +1,221 @@ +# Tencent is pleased to support the open source ecosystem. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# Licensed under Apache-2.0. +"""Long-term memory storage owned by AdvancedMemoryService.""" + +from __future__ import annotations + +import asyncio +import os +import re +import tempfile +import time +from dataclasses import replace +from datetime import datetime +from datetime import timezone +from pathlib import Path + +from ._formats import MemoryDocument +from ._formats import MemoryIndexEntry +from ._formats import limit_memory_index + +from ._config import AdvancedMemoryServiceConfig +from ._paths import AdvancedMemoryPaths + +_MEMORY_INDEX_PATTERN = re.compile(r"^- \[(?P.+?)\]((?P.+?)):(?P.+)$") + + +def parse_memory_index(index: str) -> list[MemoryIndexEntry]: + """Parse standard entries from a MEMORY.md index.""" + entries: list[MemoryIndexEntry] = [] + for line in index.splitlines(): + match = _MEMORY_INDEX_PATTERN.match(line.strip()) + if match is not None: + entries.append(MemoryIndexEntry(**match.groupdict())) + return entries + + +def prune_memory_index(index: str, valid_filenames: set[str]) -> str: + """Remove index entries whose topic files no longer exist.""" + lines = [ + line for line in index.splitlines() + if (match := _MEMORY_INDEX_PATTERN.match(line.strip())) is None or match.group("filename") in valid_filenames + ] + if lines == index.splitlines(): + return index + return "\n".join(lines) + ("\n" if lines else "") + + +def _atomic_write_text(path: Path, content: str, *, encoding: str) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + descriptor, temporary_name = tempfile.mkstemp(prefix=f".{path.name}.", dir=path.parent) + try: + with os.fdopen(descriptor, "w", encoding=encoding) as output: + output.write(content) + output.flush() + os.fsync(output.fileno()) + os.replace(temporary_name, path) + except BaseException: + try: + os.unlink(temporary_name) + except FileNotFoundError: + pass + raise + + +def _is_expired(path: Path, ttl: int | None) -> bool: + return ttl is not None and path.exists() and time.time() - path.stat().st_mtime >= ttl + + +class LongTermMemoryStore: + """Read and write MEMORY.md and its topic files.""" + + def __init__( + self, + config: AdvancedMemoryServiceConfig, + paths: AdvancedMemoryPaths | None = None, + ) -> None: + self._config = config + self._paths = paths or AdvancedMemoryPaths(config) + + @property + def index_path(self) -> Path: + return self._paths.memory_index_path + + async def initialize(self) -> None: + await asyncio.to_thread(self._initialize_sync) + + def _initialize_sync(self) -> None: + self._paths.ensure_base_directories() + if not self.index_path.exists(): + _atomic_write_text(self.index_path, "", encoding=self._config.encoding) + + async def read_index(self) -> str: + return await asyncio.to_thread(self._read_index_sync) + + def _read_index_sync(self) -> str: + if _is_expired(self.index_path, self._config.memory_ttl_seconds): + for path in self._paths.memory_dir.glob("*.md"): + path.unlink(missing_ok=True) + return "" + if not self.index_path.exists(): + return "" + with self.index_path.open(encoding=self._config.encoding) as source: + index = source.read() + valid_filenames = { + path.name + for path in self._paths.memory_dir.glob("*.md") + if path.name != self._config.memory_index_name and not _is_expired(path, self._config.memory_ttl_seconds) + } + pruned_index = prune_memory_index(index, valid_filenames) + if pruned_index != index: + _atomic_write_text( + self.index_path, + pruned_index, + encoding=self._config.encoding, + ) + return limit_memory_index( + pruned_index, + max_lines=self._config.memory_index_max_lines, + max_bytes=self._config.memory_index_max_bytes, + encoding=self._config.encoding, + ) + + async def write_index(self, entries: list[MemoryIndexEntry]) -> None: + content = "\n".join(entry.to_markdown() for entry in entries) + await asyncio.to_thread( + _atomic_write_text, + self.index_path, + f"{content}\n" if content else "", + encoding=self._config.encoding, + ) + + async def read_topic(self, topic_name: str) -> str | None: + path = self._paths.memory_topic_path(topic_name) + return await asyncio.to_thread(lambda: path.read_text(encoding=self._config.encoding) + if path.exists() else None) + + async def read_topic_frontmatter(self, topic_name: str) -> str | None: + content = await self.read_topic(topic_name) + if content is None: + return None + lines: list[str] = [] + for line in content.splitlines(keepends=True): + lines.append(line) + if len(lines) > 1 and line.rstrip("\r\n") == "---": + break + return "".join(lines) + + async def write_topic(self, topic_name: str, document: MemoryDocument) -> Path: + path = self._paths.memory_topic_path(topic_name) + updated = replace(document, updated_at=datetime.now(timezone.utc)) + await asyncio.to_thread( + _atomic_write_text, + path, + updated.to_markdown(), + encoding=self._config.encoding, + ) + return path + + async def list_topics(self) -> list[Path]: + return await asyncio.to_thread(lambda: sorted(path for path in self._paths.memory_dir.glob("*.md") + if path.name != self._config.memory_index_name)) + + +class LocalAdvancedMemoryCleanup: + """Remove expired long-term memory files for the local backend.""" + + def __init__(self, config: AdvancedMemoryServiceConfig) -> None: + self._config = config + self._task: asyncio.Task[None] | None = None + self._stop_event: asyncio.Event | None = None + + async def start(self) -> None: + if self._task is not None or self._config.memory_ttl_seconds is None: + return + self._stop_event = asyncio.Event() + await self.cleanup_once() + self._task = asyncio.create_task(self._run()) + + async def cleanup_once(self) -> None: + await asyncio.to_thread(self._cleanup_sync) + + def _cleanup_sync(self) -> None: + root = self._config.root_dir + memory_dirs = [root / self._config.memory_dir_name] + tenants_root = root / "tenants" + if tenants_root.exists(): + for app_dir in tenants_root.iterdir(): + if app_dir.is_dir(): + memory_dirs.extend(user_dir / self._config.memory_dir_name for user_dir in app_dir.iterdir() + if user_dir.is_dir()) + for memory_dir in memory_dirs: + index_path = memory_dir / self._config.memory_index_name + if _is_expired(index_path, self._config.memory_ttl_seconds): + for path in memory_dir.glob("*.md"): + path.unlink(missing_ok=True) + + async def _run(self) -> None: + if self._stop_event is None: + return + try: + while not self._stop_event.is_set(): + try: + await asyncio.wait_for( + self._stop_event.wait(), + timeout=self._config.memory_ttl_seconds or 60, + ) + except asyncio.TimeoutError: + await self.cleanup_once() + except asyncio.CancelledError: + raise + + async def close(self) -> None: + if self._stop_event is not None: + self._stop_event.set() + if self._task is not None and not self._task.done(): + self._task.cancel() + await asyncio.gather(self._task, return_exceptions=True) + self._task = None + self._stop_event = None diff --git a/trpc_agent_sdk/runners.py b/trpc_agent_sdk/runners.py index e93023cb5..16fa9a353 100644 --- a/trpc_agent_sdk/runners.py +++ b/trpc_agent_sdk/runners.py @@ -223,16 +223,11 @@ def __init__( the memory service. Set to False when the service is managed outside the runner. """ - # Advanced Memory needs the agent and session service in addition to - # the traditional memory-service hook. Bind it here so callers can - # use the same construction pattern as Redis/Mem0 memory services. - from trpc_agent_sdk.memory import AdvancedMemoryService - from trpc_agent_sdk.sessions import AdvancedMemorySessionService - - if isinstance(memory_service, AdvancedMemoryService): - session_service = memory_service.bind(agent, session_service) - elif isinstance(session_service, AdvancedMemorySessionService): - session_service = session_service.bind(agent) + if memory_service is not None: + from trpc_agent_sdk.memory import AdvancedMemoryService + + if isinstance(memory_service, AdvancedMemoryService): + session_service = memory_service.bind(agent, session_service) self.app_name = app_name self.agent = agent self.artifact_service = artifact_service diff --git a/trpc_agent_sdk/sessions/__init__.py b/trpc_agent_sdk/sessions/__init__.py index 1b8d84418..8cb7f9f6f 100644 --- a/trpc_agent_sdk/sessions/__init__.py +++ b/trpc_agent_sdk/sessions/__init__.py @@ -16,29 +16,39 @@ from ._base_session_service import BaseSessionService from ._history_record import HistoryRecord +from .compact.default import DefaultSessionSummarizer +from .compact.default import DefaultSessionSummary +from .compact.default import DefaultSessionSummarizerManager +from .compact.default import CheckSummarizerFunction +from .compact.default import set_summarizer_check_functions_by_and +from .compact.default import set_summarizer_check_functions_by_or +from .compact.default import set_summarizer_conversation_threshold +from .compact.default import set_summarizer_events_count_threshold +from .compact.default import set_summarizer_important_content_threshold +from .compact.default import set_summarizer_time_interval_threshold +from .compact.default import set_summarizer_token_threshold +from .compact.advanced import AdvancedAutoCompactSummarizer +from .compact.advanced import AdvancedAutoCompactSummarizerManager +from .compact.advanced import BaseCompactSummarizerHandler +from .compact.advanced import BaseTokenEstimator +from .compact.advanced import BaseModelContextWindowResolver +from .compact.advanced import AutoCompactSummarizerConfig +from .compact.advanced import HistorySnipConfig +from .compact.advanced import TokenContextTrackerConfig +from .compact.advanced import MicroCompactConfig +from .compact.advanced import AdvancedAutoCompactSummarizerConfig from ._in_memory_session_service import InMemorySessionService from ._in_memory_session_service import SessionWithTTL from ._in_memory_session_service import StateWithTTL from ._redis_session_service import RedisSessionService from ._redis_cluster_session_service import RedisClusterSessionService from ._session import Session -from ._session_summarizer import SessionSummarizer -from ._session_summarizer import SessionSummary from ._sql_session_service import SessionStorageBase from ._sql_session_service import SessionStorageEvent from ._sql_session_service import SqlSessionService from ._sql_session_service import StorageAppState from ._sql_session_service import StorageSession from ._sql_session_service import StorageUserState -from ._summarizer_checker import CheckSummarizerFunction -from ._summarizer_checker import set_summarizer_check_functions_by_and -from ._summarizer_checker import set_summarizer_check_functions_by_or -from ._summarizer_checker import set_summarizer_conversation_threshold -from ._summarizer_checker import set_summarizer_events_count_threshold -from ._summarizer_checker import set_summarizer_important_content_threshold -from ._summarizer_checker import set_summarizer_time_interval_threshold -from ._summarizer_checker import set_summarizer_token_threshold -from ._summarizer_manager import SummarizerSessionManager from ._types import SessionServiceConfig from ._utils import StateStorageEntry from ._utils import app_state_key @@ -49,11 +59,15 @@ from ._utils import session_key from ._utils import user_state_key +# Default compact session summarizer for backward compatibility +SessionSummary = DefaultSessionSummary +SessionSummarizer = DefaultSessionSummarizer +SummarizerSessionManager = DefaultSessionSummarizerManager + __all__ = [ "ListSessionsResponse", "State", "BaseSessionService", - "AdvancedMemorySessionService", "HistoryRecord", "InMemorySessionService", "SessionWithTTL", @@ -61,8 +75,6 @@ "RedisSessionService", "RedisClusterSessionService", "Session", - "SessionSummarizer", - "SessionSummary", "SessionStorageBase", "SessionStorageEvent", "SqlSessionService", @@ -77,6 +89,9 @@ "set_summarizer_important_content_threshold", "set_summarizer_time_interval_threshold", "set_summarizer_token_threshold", + "DefaultSessionSummarizer", + "DefaultSessionSummary", + "DefaultSessionSummarizerManager", "SummarizerSessionManager", "SessionServiceConfig", "StateStorageEntry", @@ -87,13 +102,14 @@ "is_summary_anchor", "session_key", "user_state_key", + "AdvancedAutoCompactSummarizer", + "AdvancedAutoCompactSummarizerManager", + "BaseCompactSummarizerHandler", + "BaseTokenEstimator", + "BaseModelContextWindowResolver", + "AutoCompactSummarizerConfig", + "HistorySnipConfig", + "TokenContextTrackerConfig", + "MicroCompactConfig", + "AdvancedAutoCompactSummarizerConfig", ] - - -def __getattr__(name: str): - """Lazily expose Advanced Memory without creating an import cycle.""" - if name == "AdvancedMemorySessionService": - from ._advanced_memory_session_service import AdvancedMemorySessionService - - return AdvancedMemorySessionService - raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/trpc_agent_sdk/sessions/_advanced_memory_session_service.py b/trpc_agent_sdk/sessions/_advanced_memory_session_service.py deleted file mode 100644 index feb14fb14..000000000 --- a/trpc_agent_sdk/sessions/_advanced_memory_session_service.py +++ /dev/null @@ -1,403 +0,0 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. -# -# Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Session service backed by Advanced Memory transcript storage.""" - -from __future__ import annotations - -import asyncio -import json -import os -import shutil -import tempfile -import time -import uuid -from pathlib import Path -from typing import Any -from typing import Optional - -from trpc_agent_sdk.abc import ListSessionsResponse -from trpc_agent_sdk.context import AgentContext -from trpc_agent_sdk.context import InvocationContext -from trpc_agent_sdk.events import Event -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig -from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime -from trpc_agent_sdk.advanced_memory._coordination import CrossLoopLock -from trpc_agent_sdk.advanced_memory._transcript import build_event_transcript_record -from trpc_agent_sdk.advanced_memory._transcript import find_last_event_id - -from ._base_session_service import BaseSessionService -from ._session import Session -from ._types import SessionServiceConfig -from ._utils import extract_state_delta -from ._utils import merge_state - - -class _AdvancedMemorySessionBackend(BaseSessionService): - """Persist Session metadata while TranscriptSessionService persists events.""" - - def __init__(self, runtime: AdvancedMemoryRuntime, session_config: SessionServiceConfig | None = None) -> None: - super().__init__(session_config=session_config) - self._runtime = runtime - self._lock = CrossLoopLock() - self._cleanup_task: asyncio.Task[None] | None = None - self._cleanup_stop_event: asyncio.Event | None = None - self._transcript_enabled = True - self._start_cleanup_task() - - def set_transcript_enabled(self, enabled: bool) -> None: - """Enable or disable transcript persistence for this backend.""" - self._transcript_enabled = enabled - - def _metadata_path(self, session_id: str) -> Path: - return self._runtime.paths.session_dir(session_id) / "session.json" - - @property - def _state_path(self) -> Path: - return self._runtime.paths.session_root_dir / "_state.json" - - async def _write_session(self, session: Session) -> None: - payload = session.model_dump(mode="json", by_alias=True, exclude={"events", "historical_events"}) - payload["state"] = extract_state_delta(session.state).session_state - path = self._metadata_path(session.id) - await asyncio.to_thread(self._write_json, path, payload, self._runtime.config.encoding) - - @staticmethod - def _write_json(path: Path, payload: dict[str, Any], encoding: str) -> None: - path.parent.mkdir(parents=True, exist_ok=True) - file_descriptor, temporary_name = tempfile.mkstemp(prefix=f".{path.name}.", dir=path.parent) - try: - with os.fdopen(file_descriptor, "w", encoding=encoding) as temporary_file: - temporary_file.write(json.dumps(payload, ensure_ascii=False, separators=(",", ":"))) - temporary_file.flush() - os.fsync(temporary_file.fileno()) - os.replace(temporary_name, path) - except BaseException: - try: - os.unlink(temporary_name) - except FileNotFoundError: - pass - raise - - async def _read_session(self, session_id: str) -> Session | None: - path = self._metadata_path(session_id) - if not path.exists(): - return None - payload = await asyncio.to_thread(path.read_text, encoding=self._runtime.config.encoding) - await asyncio.to_thread(path.touch) - return Session.model_validate(json.loads(payload)) - - def _start_cleanup_task(self) -> None: - """Start persistent session cleanup when TTL is enabled.""" - if not self.session_config.need_ttl_expire() or self._cleanup_task is not None: - return - try: - loop = asyncio.get_running_loop() - except RuntimeError: - return - self._cleanup_stop_event = asyncio.Event() - self._cleanup_task = loop.create_task(self._cleanup_loop()) - - async def _cleanup_loop(self) -> None: - """Periodically remove expired session directories.""" - assert self._cleanup_stop_event is not None - try: - while not self._cleanup_stop_event.is_set(): - try: - await asyncio.wait_for( - self._cleanup_stop_event.wait(), - timeout=self.session_config.ttl.cleanup_interval_seconds, - ) - break - except asyncio.TimeoutError: - async with self._lock: - await asyncio.to_thread(self._cleanup_expired_sessions) - except asyncio.CancelledError: - raise - - def _cleanup_expired_sessions(self) -> None: - """Delete session directories idle longer than the configured TTL.""" - cutoff = time.time() - self.session_config.ttl.ttl_seconds - root = self._runtime.paths.session_root_dir - if not root.exists(): - return - for metadata_path in root.glob("*/session.json"): - try: - if metadata_path.stat().st_mtime < cutoff: - shutil.rmtree(metadata_path.parent, ignore_errors=True) - except FileNotFoundError: - continue - - async def _stop_cleanup_task(self) -> None: - """Stop the background TTL cleanup task.""" - task = self._cleanup_task - self._cleanup_task = None - if task is None: - return - if self._cleanup_stop_event is not None: - self._cleanup_stop_event.set() - task.cancel() - await asyncio.gather(task, return_exceptions=True) - self._cleanup_stop_event = None - - async def _read_global_state(self) -> dict[str, dict[str, Any]]: - if not self._state_path.exists(): - return {"app": {}, "user": {}} - payload = await asyncio.to_thread( - self._state_path.read_text, - encoding=self._runtime.config.encoding, - ) - parsed = json.loads(payload) - return { - "app": dict(parsed.get("app", {})), - "user": dict(parsed.get("user", {})), - } - - async def _write_global_state(self, state: dict[str, dict[str, Any]]) -> None: - await asyncio.to_thread(self._write_json, self._state_path, state, self._runtime.config.encoding) - - async def _restore_events(self, session: Session) -> Session: - records = await self._runtime.transcripts.read_all(session.id) - events: list[Event] = [] - for record in records: - event_payload = record.get("event") - if record.get("kind") != "event" or not isinstance(event_payload, dict): - continue - events.append(Event.model_validate(event_payload)) - session.events = events - if events: - session.last_update_time = events[-1].timestamp - return session - - async def create_session( - self, - *, - app_name: str, - user_id: str, - state: Optional[dict[str, Any]] = None, - session_id: Optional[str] = None, - agent_context: Optional[AgentContext] = None, - ) -> Session: - self._start_cleanup_task() - resolved_id = session_id.strip() if session_id and session_id.strip() else str(uuid.uuid4()) - state_delta = extract_state_delta(state) - session = Session( - id=resolved_id, - app_name=app_name, - user_id=user_id, - state=state_delta.session_state, - save_key=f"{app_name}/{user_id}", - ) - async with self._lock: - await self._runtime.initialize() - existing = await self._read_session(resolved_id) - if existing is not None and (existing.app_name != app_name or existing.user_id != user_id): - raise ValueError(f"Session ID {resolved_id!r} is already used by another app or user") - global_state = await self._read_global_state() - global_state["app"].setdefault(app_name, {}).update(state_delta.app_state_delta) - global_state["user"].setdefault(f"{app_name}/{user_id}", {}).update(state_delta.user_state_delta) - await self._write_global_state(global_state) - await self._write_session(session) - session.state = merge_state( - extract_state_delta(session.state), - need_copy=True, - ) - session.state.update({f"app:{key}": value for key, value in global_state["app"].get(app_name, {}).items()}) - session.state.update({ - f"user:{key}": value - for key, value in global_state["user"].get(f"{app_name}/{user_id}", {}).items() - }) - return session - - async def get_session( - self, - *, - app_name: str, - user_id: str, - session_id: str, - agent_context: Optional[AgentContext] = None, - ) -> Session | None: - self._start_cleanup_task() - async with self._lock: - session = await self._read_session(session_id) - if session is None or session.app_name != app_name or session.user_id != user_id: - return None - global_state = await self._read_global_state() - app_state = global_state["app"].get(app_name, {}) - user_state = global_state["user"].get(f"{app_name}/{user_id}", {}) - session.state = merge_state( - extract_state_delta(session.state), - need_copy=True, - ) - session.state.update({f"app:{key}": value for key, value in app_state.items()}) - session.state.update({f"user:{key}": value for key, value in user_state.items()}) - return self.filter_events(await self._restore_events(session), need_copy=True) - - async def list_sessions( - self, - *, - app_name: str, - user_id: Optional[str] = None, - ) -> ListSessionsResponse: - self._start_cleanup_task() - if not self._runtime.paths.session_root_dir.exists(): - return ListSessionsResponse() - sessions: list[Session] = [] - for path in await asyncio.to_thread(lambda: list(self._runtime.paths.session_root_dir.glob("*/session.json"))): - try: - session = await asyncio.to_thread(lambda path=path: Session.model_validate( - json.loads(path.read_text(encoding=self._runtime.config.encoding)))) - except (OSError, ValueError, TypeError): - continue - if session.app_name == app_name and (user_id is None or session.user_id == user_id): - session.events = [] - session.historical_events = [] - sessions.append(session) - return ListSessionsResponse(sessions=sessions) - - async def delete_session(self, *, app_name: str, user_id: str, session_id: str) -> None: - self._start_cleanup_task() - session = await self.get_session( - app_name=app_name, - user_id=user_id, - session_id=session_id, - ) - if session is not None: - async with self._lock: - await asyncio.to_thread(shutil.rmtree, self._runtime.paths.session_dir(session_id), True) - - async def append_event(self, session: Session, event: Event) -> Event: - self._start_cleanup_task() - async with self._lock: - persisted = await super().append_event(session, event) - if not event.partial: - state_delta = extract_state_delta(event.actions.state_delta if event.actions else None) - if state_delta.app_state_delta or state_delta.user_state_delta: - global_state = await self._read_global_state() - global_state["app"].setdefault(session.app_name, {}).update(state_delta.app_state_delta) - global_state["user"].setdefault(f"{session.app_name}/{session.user_id}", - {}).update(state_delta.user_state_delta) - await self._write_global_state(global_state) - session.state.update({f"app:{key}": value for key, value in state_delta.app_state_delta.items()}) - session.state.update({f"user:{key}": value for key, value in state_delta.user_state_delta.items()}) - await self._write_session(session) - if not event.partial and self._transcript_enabled: - records = await self._runtime.transcripts.read_all(session.id) - record = build_event_transcript_record( - session, - persisted, - parent_event_id=find_last_event_id(records), - ) - await self._runtime.transcripts.append_unique( - session.id, - record, - unique_key="event_id", - ) - return persisted - - async def update_session(self, session: Session) -> None: - self._start_cleanup_task() - async with self._lock: - await self._write_session(session) - - async def create_session_summary( - self, - session: Session, - ctx: InvocationContext | None = None, - ) -> None: - await super().create_session_summary(session, ctx=ctx) - await self.update_session(session) - - async def get_session_summary(self, session: Session) -> str | None: - return await super().get_session_summary(session) - - async def close(self) -> None: - await self._stop_cleanup_task() - - -class AdvancedMemorySessionService(BaseSessionService): - """Persist sessions and raw events in the Advanced Memory directory.""" - - def __init__( - self, - runtime: AdvancedMemoryRuntime | None = None, - *, - config: AdvancedMemoryConfig | None = None, - session_config: SessionServiceConfig | None = None, - preload_memory_model: Any | None = None, - ) -> None: - if runtime is not None and config is not None and runtime.config != config: - raise ValueError("runtime and config must describe the same Advanced Memory configuration") - self._runtime = runtime or AdvancedMemoryRuntime.create(config) - self._preload_memory_model = preload_memory_model - self._backend = _AdvancedMemorySessionBackend(self._runtime, session_config=session_config) - self._integration: Any | None = None - self._bound_agent: Any | None = None - super().__init__(session_config=session_config) - - @property - def runtime(self) -> AdvancedMemoryRuntime: - """Return the Advanced Memory runtime used by this service.""" - return self._runtime - - @property - def integration(self) -> Any | None: - """Return the Advanced Memory binding, when attached to a Runner.""" - return self._integration - - @property - def backend(self) -> BaseSessionService: - """Return the persistent backend used by the transcript decorator.""" - return self._backend - - def bind(self, agent: Any) -> BaseSessionService: - """Install Advanced Memory callbacks and return the wrapped service.""" - from trpc_agent_sdk.advanced_memory import setup_advanced_memory - - if self._integration is not None: - if agent is not self._bound_agent: - raise ValueError("AdvancedMemorySessionService is already bound to another agent") - return self._integration.session_service - integration = setup_advanced_memory( - agent, - self, - self._runtime, - preload_memory_model=self._preload_memory_model, - ) - self._backend.set_transcript_enabled(False) - self._integration = integration - self._bound_agent = agent - return self._integration.session_service - - async def create_session(self, **kwargs: Any) -> Session: - return await self._backend.create_session(**kwargs) - - async def get_session(self, **kwargs: Any) -> Session | None: - return await self._backend.get_session(**kwargs) - - async def list_sessions(self, **kwargs: Any) -> ListSessionsResponse: - return await self._backend.list_sessions(**kwargs) - - async def delete_session(self, **kwargs: Any) -> None: - await self._backend.delete_session(**kwargs) - - async def append_event(self, session: Session, event: Event) -> Event: - return await self._backend.append_event(session, event) - - async def update_session(self, session: Session) -> None: - await self._backend.update_session(session) - - async def create_session_summary( - self, - session: Session, - ctx: InvocationContext | None = None, - ) -> None: - await self._backend.create_session_summary(session, ctx=ctx) - - async def get_session_summary(self, session: Session) -> str | None: - return await self._backend.get_session_summary(session) - - async def close(self) -> None: - await self._backend.close() diff --git a/trpc_agent_sdk/sessions/_base_session_service.py b/trpc_agent_sdk/sessions/_base_session_service.py index 979523f46..84c02881d 100644 --- a/trpc_agent_sdk/sessions/_base_session_service.py +++ b/trpc_agent_sdk/sessions/_base_session_service.py @@ -24,16 +24,17 @@ """Base session service interface.""" from __future__ import annotations + from typing import Optional from typing_extensions import override from trpc_agent_sdk.abc import SessionServiceABC +from trpc_agent_sdk.abc import CompactSummarizerManagerABC from trpc_agent_sdk.context import InvocationContext from trpc_agent_sdk.events import Event from trpc_agent_sdk.types import State from ._session import Session -from ._summarizer_manager import SummarizerSessionManager from ._types import SessionServiceConfig @@ -44,7 +45,7 @@ class BaseSessionService(SessionServiceABC): """ def __init__(self, - summarizer_manager: Optional[SummarizerSessionManager] = None, + summarizer_manager: Optional[CompactSummarizerManagerABC] = None, session_config: Optional[SessionServiceConfig] = None): """Initialize the base session service. @@ -62,7 +63,7 @@ def __init__(self, self._summarizer_manager.set_session_service(self) @property - def summarizer_manager(self) -> Optional[SummarizerSessionManager]: + def summarizer_manager(self) -> Optional[CompactSummarizerManagerABC]: """Get the summarizer manager.""" return self._summarizer_manager @@ -71,7 +72,7 @@ def session_config(self) -> SessionServiceConfig: """Get the session service configuration.""" return self._session_config - def set_summarizer_manager(self, summarizer_manager: SummarizerSessionManager, force: bool = False) -> None: + def set_summarizer_manager(self, summarizer_manager: CompactSummarizerManagerABC, force: bool = False) -> None: """Set the summarizer manager to use. Args: @@ -80,7 +81,7 @@ def set_summarizer_manager(self, summarizer_manager: SummarizerSessionManager, f """ if not self._summarizer_manager or force: self._summarizer_manager = summarizer_manager - self._summarizer_manager.set_session_service(self) + self._summarizer_manager.set_session_service(self, force) @override async def append_event(self, session: Session, event: Event) -> Event: @@ -187,7 +188,9 @@ async def get_session_summary(self, session: Session) -> Optional[str]: """ if self._summarizer_manager: summary = await self._summarizer_manager.get_session_summary(session) - if summary: + if isinstance(summary, str): + return summary + if summary is not None: return summary.summary_text return None @@ -211,4 +214,5 @@ def filter_events(self, session: Session, need_copy: bool = False) -> Session: @override async def close(self) -> None: """Closes the session service and releases any resources.""" - pass + if self._summarizer_manager: + await self._summarizer_manager.close() diff --git a/trpc_agent_sdk/sessions/_in_memory_session_service.py b/trpc_agent_sdk/sessions/_in_memory_session_service.py index 567a52d16..c19d56ab1 100644 --- a/trpc_agent_sdk/sessions/_in_memory_session_service.py +++ b/trpc_agent_sdk/sessions/_in_memory_session_service.py @@ -37,6 +37,7 @@ from pydantic import Field from trpc_agent_sdk.abc import ListSessionsResponse +from trpc_agent_sdk.abc import CompactSummarizerManagerABC from trpc_agent_sdk.context import AgentContext from trpc_agent_sdk.events import Event from trpc_agent_sdk.log import logger @@ -45,7 +46,6 @@ from ._base_session_service import BaseSessionService from ._session import Session -from ._summarizer_manager import SummarizerSessionManager from ._types import SessionServiceConfig from ._utils import StateStorageEntry from ._utils import extract_state_delta @@ -107,7 +107,7 @@ class InMemorySessionService(BaseSessionService): """ def __init__(self, - summarizer_manager: Optional[SummarizerSessionManager] = None, + summarizer_manager: Optional[CompactSummarizerManagerABC] = None, session_config: Optional[SessionServiceConfig] = None): super().__init__(summarizer_manager=summarizer_manager, session_config=session_config) # Storage with TTL support @@ -213,9 +213,8 @@ async def list_sessions(self, *, app_name: str, user_id: Optional[str] = None) - @override async def delete_session(self, *, app_name: str, user_id: str, session_id: str) -> None: - if not self._is_session_exist(app_name=app_name, user_id=user_id, session_id=session_id): - return - del self._sessions[app_name][user_id][session_id] + if self._is_session_exist(app_name=app_name, user_id=user_id, session_id=session_id): + del self._sessions[app_name][user_id][session_id] @override async def append_event(self, session: Session, event: Event) -> Event: @@ -270,6 +269,26 @@ def _warning(message: str) -> None: return event + @override + async def update_session_state( + self, + session: Session, + state_delta: dict[str, Any], + ) -> None: + """Patch stored session state without replacing its Event window.""" + if not state_delta: + return + session.state.update(state_delta) + + app_sessions = self._sessions.get(session.app_name) + user_sessions = app_sessions.get(session.user_id) if app_sessions else None + stored = user_sessions.get(session.id) if user_sessions else None + if stored is None: + logger.warning("Session %s not found while updating state", session.id) + return + stored.session.state.update(state_delta) + stored.ttl.update_expired_at() + @override async def update_session(self, session: Session) -> None: """Update a session in storage. diff --git a/trpc_agent_sdk/sessions/_redis_session_service.py b/trpc_agent_sdk/sessions/_redis_session_service.py index 8bec47af1..9190f898f 100644 --- a/trpc_agent_sdk/sessions/_redis_session_service.py +++ b/trpc_agent_sdk/sessions/_redis_session_service.py @@ -8,6 +8,7 @@ from __future__ import annotations +import json import time import uuid from typing import Any @@ -15,6 +16,7 @@ from typing_extensions import override from trpc_agent_sdk.abc import ListSessionsResponse +from trpc_agent_sdk.abc import CompactSummarizerManagerABC from trpc_agent_sdk.context import AgentContext from trpc_agent_sdk.events import Event from trpc_agent_sdk.log import logger @@ -26,7 +28,6 @@ from ._base_session_service import BaseSessionService from ._session import Session -from ._summarizer_manager import SummarizerSessionManager from ._types import SessionServiceConfig from ._utils import StateStorageEntry from ._utils import app_state_key @@ -54,6 +55,15 @@ def _session_key_prefix(app_name: str, user_id: Optional[str] = None) -> str: return f"session:{app_name}:{user_id}:*" +def _session_from_storage_json(value: Any) -> Session: + """Decode a Session and repair empty arrays changed to objects by Lua cjson.""" + payload = json.loads(value) + for field_name in ("events", "historical_events", "historicalEvents"): + if payload.get(field_name) == {}: + payload[field_name] = [] + return Session.model_validate(payload) + + class RedisSessionService(BaseSessionService): """A Redis implementation of the session service. @@ -76,18 +86,33 @@ class RedisSessionService(BaseSessionService): def __init__(self, db_url: str, - summarizer_manager: Optional[SummarizerSessionManager] = None, + summarizer_manager: Optional[CompactSummarizerManagerABC] = None, session_config: Optional[SessionServiceConfig] = None, is_async: bool = False, **kwargs: Any): + self._db_url = db_url + self._is_async = is_async is_default_config = session_config is None - super().__init__(summarizer_manager=summarizer_manager, session_config=session_config) + super().__init__( + summarizer_manager=summarizer_manager, + session_config=session_config, + ) if is_default_config: # Default to store historical events for persistent backends. self._session_config.store_historical_events = True # Redis needs default TTL configuration self._redis_storage = self._create_storage(db_url=db_url, is_async=is_async, **kwargs) + @property + def db_url(self) -> str: + """Return the configured Redis connection URL.""" + return self._db_url + + @property + def is_async(self) -> bool: + """Return whether this service uses the asynchronous Redis client.""" + return self._is_async + def _create_storage(self, db_url: str, is_async: bool, **kwargs: Any) -> RedisStorage: """Create the backing storage. @@ -235,6 +260,29 @@ def _warning(message: str) -> None: return event + @override + async def update_session_state( + self, + session: Session, + state_delta: dict[str, Any], + ) -> None: + """Persist session-scoped state without replacing the caller's Event window.""" + if not state_delta: + return + session.state.update(state_delta) + + async with self._redis_storage.create_db_session() as redis_session: + key = session_key(session.app_name, session.user_id, session.id) + storage_session = await self._get_session(redis_session, key) + if not storage_session: + logger.warning( + "Session %s not found in Redis while updating state", + session.id, + ) + return + storage_session.state.update(state_delta) + await self._set_session(redis_session, storage_session) + @override async def update_session(self, session: Session) -> None: """Update a session in storage. @@ -410,7 +458,7 @@ async def _get_session(self, redis_session: RedisSession, session_key: str) -> O storage_session_data = await self._redis_storage.execute_command(redis_session, command) if storage_session_data: await self._refresh_ttl(redis_session, session_key) - session = Session.model_validate_json(storage_session_data) + session = _session_from_storage_json(storage_session_data) if not self._session_config.store_historical_events: session.historical_events = [] return session diff --git a/trpc_agent_sdk/sessions/_session.py b/trpc_agent_sdk/sessions/_session.py index b0fd094f6..061c9e2b6 100644 --- a/trpc_agent_sdk/sessions/_session.py +++ b/trpc_agent_sdk/sessions/_session.py @@ -136,3 +136,51 @@ def insert_events(self, events: List[Event], idx: Optional[int] = None) -> None: if idx is None: idx = 0 self.events[idx:idx] = events + + def compact_events( + self, + summary_event: Event, + boundary_event_id: str, + *, + compaction_id: str, + ) -> bool: + """Replace the active prefix through ``boundary_event_id`` with a summary. + + The replaced active Events remain recoverable in ``historical_events``. + ``compaction_id`` makes retries idempotent when a persistence operation + succeeds but its caller does not observe the result. + """ + for event in self.events: + metadata = event.custom_metadata or {} + if metadata.get("session_compaction_id") == compaction_id: + return False + + boundary_index = next( + (index for index, event in enumerate(self.events) if event.id == boundary_event_id), + None, + ) + if boundary_index is None: + raise ValueError(f"Session compaction boundary Event {boundary_event_id!r} " + "is not in the active event window") + + replaced = self.events[:boundary_index + 1] + if not replaced: + return False + + metadata = dict(summary_event.custom_metadata or {}) + metadata.update({ + "session_compaction_id": compaction_id, + "session_compaction_boundary_event_id": boundary_event_id, + }) + summary_event.custom_metadata = metadata + summary_event.set_summary_event(True) + # SQL backends restore active Events in timestamp order. Give the + # replacement summary the prefix's timestamp so it remains the anchor + # before every retained Event after persistence. + summary_event.timestamp = replaced[0].timestamp + + historical_ids = {event.id for event in self.historical_events} + self.historical_events.extend(event for event in replaced if event.id not in historical_ids) + self.events = [summary_event, *self.events[boundary_index + 1:]] + self.last_update_time = max(self.last_update_time, summary_event.timestamp) + return True diff --git a/trpc_agent_sdk/sessions/_sql_session_service.py b/trpc_agent_sdk/sessions/_sql_session_service.py index 4333ffeb1..e4f5225a0 100644 --- a/trpc_agent_sdk/sessions/_sql_session_service.py +++ b/trpc_agent_sdk/sessions/_sql_session_service.py @@ -50,6 +50,7 @@ from sqlalchemy.types import Integer from trpc_agent_sdk.abc import ListSessionsResponse +from trpc_agent_sdk.abc import CompactSummarizerManagerABC from trpc_agent_sdk.context import AgentContext from trpc_agent_sdk.events import Event from trpc_agent_sdk.log import logger @@ -71,7 +72,6 @@ from ._base_session_service import BaseSessionService from ._session import Session -from ._summarizer_manager import SummarizerSessionManager from ._types import SessionServiceConfig from ._utils import StateStorageEntry from ._utils import extract_state_delta @@ -388,21 +388,40 @@ class SqlSessionService(BaseSessionService): def __init__(self, db_url: str, - summarizer_manager: Optional[SummarizerSessionManager] = None, + summarizer_manager: Optional[CompactSummarizerManagerABC] = None, is_async: bool = False, session_config: Optional[SessionServiceConfig] = None, **kwargs: Any): + self._db_url = db_url + self._is_async = is_async is_default_config = session_config is None - super().__init__(summarizer_manager=summarizer_manager, session_config=session_config) + super().__init__( + summarizer_manager=summarizer_manager, + session_config=session_config, + ) if is_default_config: # Default to store historical events for persistent backends. self._session_config.store_historical_events = True + # AsyncSession cannot perform an implicit refresh when an ORM + # attribute is accessed after commit. Keep committed values available + # because this service reads StorageSession state after committing. + kwargs.setdefault("expire_on_commit", False) self._sql_storage = SqlStorage(is_async=is_async, db_url=db_url, metadata=SessionStorageBase.metadata, **kwargs) self.__cleanup_task: Optional[asyncio.Task] = None self.__cleanup_stop_event: Optional[asyncio.Event] = None self._start_cleanup_task() + @property + def db_url(self) -> str: + """Return the configured SQL connection URL.""" + return self._db_url + + @property + def is_async(self) -> bool: + """Return whether this service uses asynchronous SQL sessions.""" + return self._is_async + @override async def create_session( self, @@ -543,7 +562,10 @@ async def append_event(self, session: Session, event: Event) -> Event: async with self._sql_storage.create_db_session() as sql_session: session_key = SqlKey(key=(app_name, user_id, session_id), storage_cls=StorageSession) - storage_session: Optional[StorageSession] = await self._sql_storage.get(sql_session, session_key) + storage_session: Optional[StorageSession] = await self._sql_storage.get_for_update( + sql_session, + session_key, + ) if not storage_session: logger.warning("Session %s not found in storage, it will be created", session_id) return event @@ -616,6 +638,40 @@ async def append_event(self, session: Session, event: Event) -> Event: return event + @override + async def update_session_state( + self, + session: Session, + state_delta: dict[str, Any], + ) -> None: + """Persist session-scoped state without rewriting Event rows.""" + if not state_delta: + return + session.state.update(state_delta) + + async with self._sql_storage.create_db_session() as sql_session: + session_key = SqlKey( + key=(session.app_name, session.user_id, session.id), + storage_cls=StorageSession, + ) + storage_session: Optional[StorageSession] = await self._sql_storage.get_for_update( + sql_session, + session_key, + ) + if storage_session is None: + logger.warning( + "Session %s not found in storage while updating state", + session.id, + ) + return + + persisted_state = dict(storage_session.state or {}) + persisted_state.update(state_delta) + storage_session.state = persisted_state # type: ignore + await self._sql_storage.commit(sql_session) + await self._sql_storage.refresh(sql_session, storage_session) + session.last_update_time = storage_session.update_timestamp_tz + @override async def update_session(self, session: Session) -> None: app_name = session.app_name @@ -704,6 +760,7 @@ async def _get_app_state(self, sql_session: SqlSession, app_name: str) -> dict[s app_state = storage_app_state.state storage_app_state.update_time = func.now() await self._sql_storage.commit(sql_session) + await self._sql_storage.refresh(sql_session, storage_app_state) return app_state @@ -717,6 +774,7 @@ async def _get_user_state(self, sql_session: SqlSession, app_name: str, user_id: user_state = storage_user_state.state storage_user_state.update_time = func.now() await self._sql_storage.commit(sql_session) + await self._sql_storage.refresh(sql_session, storage_user_state) return user_state @@ -733,6 +791,10 @@ async def _get_session(self, sql_session: SqlSession, app_name: str, user_id: st storage_session.update_time = func.now() await self._sql_storage.commit(sql_session) + # Assigning a SQL expression expires the server-generated timestamp + # even when expire_on_commit=False. Refresh it before callers access + # update_time outside SQLAlchemy's async greenlet. + await self._sql_storage.refresh(sql_session, storage_session) return storage_session diff --git a/trpc_agent_sdk/sessions/compact/__init__.py b/trpc_agent_sdk/sessions/compact/__init__.py new file mode 100644 index 000000000..77a45bb99 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/__init__.py @@ -0,0 +1,77 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Canonical context-compression package for session management.""" + +from trpc_agent_sdk.abc import CompactTrigger + +from .advanced import AdvancedAutoCompactSummarizer +from .advanced import AdvancedAutoCompactSummarizerManager +from .advanced import BaseCompactSummarizerHandler +from .advanced import BaseTokenEstimator +from .advanced import BaseModelContextWindowResolver +from .advanced import AutoCompactSummarizerConfig +from .advanced import HistorySnipConfig +from .advanced import TokenContextTrackerConfig +from .advanced import MicroCompactConfig +from .advanced import ToolResultBudgetConfig +from .advanced import SessionMemoryExtractorConfig +from .advanced import AdvancedAutoCompactSummarizerConfig +from .advanced import AdvancedAutoCompactSummarizerFilter +from .advanced import HistorySnip +from .advanced import MicroCompact +from .advanced import SessionMemoryDocument +from .advanced import SessionMemoryExtractor +from .advanced import AdvancedAutoCompactSummarizerRuntime +from .advanced import TokenContextTracker +from .advanced import ToolResultBudget +from .default import DEFAULT_SUMMARIZER_PROMPT +from .default import DefaultSessionSummarizer +from .default import DefaultSessionSummarizerManager +from .default import DefaultSessionSummary +from .default import CheckSummarizerFunction +from .default import set_summarizer_token_threshold +from .default import set_summarizer_events_count_threshold +from .default import set_summarizer_time_interval_threshold +from .default import set_summarizer_important_content_threshold +from .default import set_summarizer_conversation_threshold +from .default import set_summarizer_check_functions_by_and +from .default import set_summarizer_check_functions_by_or + +__all__ = [ + "CompactTrigger", + "AdvancedAutoCompactSummarizer", + "AdvancedAutoCompactSummarizerManager", + "BaseCompactSummarizerHandler", + "BaseTokenEstimator", + "BaseModelContextWindowResolver", + "AutoCompactSummarizerConfig", + "HistorySnipConfig", + "TokenContextTrackerConfig", + "MicroCompactConfig", + "ToolResultBudgetConfig", + "SessionMemoryExtractorConfig", + "AdvancedAutoCompactSummarizerConfig", + "AdvancedAutoCompactSummarizerFilter", + "HistorySnip", + "MicroCompact", + "SessionMemoryDocument", + "SessionMemoryExtractor", + "AdvancedAutoCompactSummarizerRuntime", + "TokenContextTracker", + "ToolResultBudget", + "DEFAULT_SUMMARIZER_PROMPT", + "DefaultSessionSummarizer", + "DefaultSessionSummarizerManager", + "DefaultSessionSummary", + "CheckSummarizerFunction", + "set_summarizer_token_threshold", + "set_summarizer_events_count_threshold", + "set_summarizer_time_interval_threshold", + "set_summarizer_important_content_threshold", + "set_summarizer_conversation_threshold", + "set_summarizer_check_functions_by_and", + "set_summarizer_check_functions_by_or", +] diff --git a/trpc_agent_sdk/advanced_memory/_callbacks.py b/trpc_agent_sdk/sessions/compact/_callbacks.py similarity index 92% rename from trpc_agent_sdk/advanced_memory/_callbacks.py rename to trpc_agent_sdk/sessions/compact/_callbacks.py index 99b422346..44249496a 100644 --- a/trpc_agent_sdk/advanced_memory/_callbacks.py +++ b/trpc_agent_sdk/sessions/compact/_callbacks.py @@ -1,7 +1,7 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# Tencent is pleased to support the open source community by making +# contributions to the open source ecosystem. # # Copyright (C) 2026 Tencent. All rights reserved. -# # tRPC-Agent-Python is licensed under Apache-2.0. """Shared callback installation and stage ordering for Advanced Memory.""" @@ -9,8 +9,6 @@ from typing import Any -from ._runtime import AdvancedMemoryRuntime - def install_staged_callback( agent: Any, @@ -18,7 +16,7 @@ def install_staged_callback( *, callback_type: type, component_attribute: str, - memory_runtime: AdvancedMemoryRuntime, + memory_runtime: Any, conflict_message: str, ) -> Any | None: """Install a staged callback idempotently and validate runtime ownership.""" diff --git a/trpc_agent_sdk/sessions/compact/advanced/__init__.py b/trpc_agent_sdk/sessions/compact/advanced/__init__.py new file mode 100644 index 000000000..6d20ea48c --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/__init__.py @@ -0,0 +1,50 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Advanced compact session manager.""" + +from ._auto_compact import AdvancedAutoCompactSummarizer +from ._base import BaseCompactSummarizerHandler +from ._base import BaseTokenEstimator +from ._base import BaseModelContextWindowResolver +from ._config import AutoCompactSummarizerConfig +from ._config import HistorySnipConfig +from ._config import TokenContextTrackerConfig +from ._config import MicroCompactConfig +from ._config import ToolResultBudgetConfig +from ._config import SessionMemoryExtractorConfig +from ._config import AdvancedAutoCompactSummarizerConfig +from ._filters import AdvancedAutoCompactSummarizerFilter +from ._formats import SessionMemoryDocument +from ._history_snip import HistorySnip +from ._micro_compact import MicroCompact +from ._manager import AdvancedAutoCompactSummarizerManager +from ._compaction_memory_extractor import SessionMemoryExtractor +from ._runtime import AdvancedAutoCompactSummarizerRuntime +from ._token_budget import TokenContextTracker +from ._tool_result_budget import ToolResultBudget + +__all__ = [ + "AdvancedAutoCompactSummarizer", + "AdvancedAutoCompactSummarizerManager", + "BaseCompactSummarizerHandler", + "BaseTokenEstimator", + "BaseModelContextWindowResolver", + "AutoCompactSummarizerConfig", + "HistorySnipConfig", + "TokenContextTrackerConfig", + "MicroCompactConfig", + "ToolResultBudgetConfig", + "SessionMemoryExtractorConfig", + "AdvancedAutoCompactSummarizerConfig", + "AdvancedAutoCompactSummarizerFilter", + "HistorySnip", + "MicroCompact", + "SessionMemoryDocument", + "SessionMemoryExtractor", + "AdvancedAutoCompactSummarizerRuntime", + "TokenContextTracker", + "ToolResultBudget", +] diff --git a/trpc_agent_sdk/sessions/compact/advanced/_auto_compact.py b/trpc_agent_sdk/sessions/compact/advanced/_auto_compact.py new file mode 100644 index 000000000..4b43f0c0b --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/_auto_compact.py @@ -0,0 +1,811 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Automatically compact history before model requests and circuit-break failures.""" + +from __future__ import annotations + +import asyncio +import copy +import json +import re +import uuid +from dataclasses import dataclass +from typing_extensions import override + +from trpc_agent_sdk.abc import CompactSummarizerABC +from trpc_agent_sdk.abc import RequestABC +from trpc_agent_sdk.abc import ResponseABC +from trpc_agent_sdk.abc import SessionABC +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.events import Event +from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.models import LlmResponse +from trpc_agent_sdk.models import LLMModel +from trpc_agent_sdk.types import Content +from trpc_agent_sdk.types import Part + +from ._base import BaseCompactSummarizerHandler +from ._config import AdvancedAutoCompactSummarizerConfig +from ._formats import SESSION_MEMORY_SECTIONS +from ._formats import SESSION_MEMORY_STATE_KEY +from ._formats import SessionMemoryDocument +from ._formats import parse_session_memory_state +from ._history_snip import estimate_request_chars +from ._runtime import AdvancedAutoCompactSummarizerRuntime +from ._token_budget import TokenContextTracker +from ._compaction_memory_extractor import SessionMemoryExtractor +from ._utils import content_signature +from ._utils import internal_compaction_call + +ADVANCED_AUTOCOMPACT_BLOCKED_MESSAGE = ( + "Automatic context compaction has failed repeatedly and the request is near the hard context limit. " + "To avoid sending a request that will certainly fail, reduce the input, start a new session, " + "or manually organize session memory before retrying.") +ADVANCED_AUTOCOMPACT_SUMMARY_PREFIX = """This session is being continued from a compacted context. +The following summary contains the important information from earlier messages. +The complete original events remain available in the SessionService. + +""" +_LEGACY_SESSION_MEMORY_SECTION_LIST = "\n".join(f"- # {section}" for section in SESSION_MEMORY_SECTIONS) + +LEGACY_SUMMARY_INSTRUCTION = """You are an isolated context-compaction Agent. +Compress the provided old conversation into a dense Markdown summary that another Agent can continue seamlessly. +Preserve the user's goals, explicit requirements, key technical decisions, files and functions, commands, +errors and fixes, verified results, current state, and next steps. +Do not answer questions from the old conversation, mention this compaction prompt, or invent information. +Return exactly two XML blocks: first use ... to check coverage, then +... for the final Markdown summary. The summary must contain these ten Markdown sections +in this order: +""" + _LEGACY_SESSION_MEMORY_SECTION_LIST + """ +The analysis is only for organization; keep only the summary.""" + + +@dataclass(frozen=True) +class AdvancedAutoCompactRecord: + """Store stable replay information for the latest successful compaction.""" + + boundary_signature: str + boundary_occurrence: int + summary: str + source: str + boundary_event_id: str | None = None + compaction_id: str | None = None + + +@dataclass +class AdvancedAutoCompactState: + """Store the latest compaction record and consecutive failure count.""" + + latest_compaction: AdvancedAutoCompactRecord | None + consecutive_failures: int + + +@dataclass(frozen=True) +class AdvancedAutoCompactResult: + """Summarize one compaction, replay, or hard-block result.""" + + compacted: bool + reapplied: bool + blocked: bool + source: str | None + request_chars_before: int + request_chars_after: int + consecutive_failures: int + error: str | None = None + request_tokens_before: int | None = None + request_tokens_after: int | None = None + token_source: str | None = None + summary: str | None = None + + +def _content_text(content: Content) -> str: + """Render one model content item as legacy summary input.""" + return json.dumps( + content.model_dump(mode="json", by_alias=True, exclude_none=True), + ensure_ascii=False, + sort_keys=True, + separators=(",", ":"), + default=str, + ) + + +class AdvancedAutoCompactSummarizer(CompactSummarizerABC): + """Compact with session memory first, then fall back to a legacy summary.""" + + def __init__( + self, + config: AdvancedAutoCompactSummarizerConfig | None = None, + *, + model: LLMModel | None = None, + session_memory_extractor: SessionMemoryExtractor | None = None, + ) -> None: + """Initialize the compressor, summary generator, and session locks.""" + self._model = model + self._runtime = AdvancedAutoCompactSummarizerRuntime(config=config or AdvancedAutoCompactSummarizerConfig()) + self._auto_compact_config = config.auto_compact + self._session_memory_extractor = self._create_extractor(session_memory_extractor) + self._states: dict[str, AdvancedAutoCompactState] = {} + self._session_locks: dict[str, asyncio.Lock] = {} + self._scoped_processors: dict[object, "AdvancedAutoCompactSummarizer"] = {} + + def _resolve_model(self, ctx: InvocationContext) -> LLMModel: + """Resolve the model used for legacy compaction.""" + if self._model is not None: + return self._model + if ctx.agent is None: + raise ValueError("Autocompact summary generator cannot resolve an LLM model") + return ctx.agent.model + + async def _generate_summary(self, history: str, ctx: InvocationContext | None = None) -> str: + """Generate a summary using the LLM model. + + Args: + history: The conversation text to summarize + + Returns: + Generated summary text + """ + request = LlmRequest() + request.append_instructions([LEGACY_SUMMARY_INSTRUCTION]) + prompt = ("Compress the following old conversation. The input may contain JSON representations " + "of tool calls and results:\n\n" + f"\n{history}\n") + request.contents.append(Content(role="user", parts=[Part.from_text(text=prompt)])) + + output = "" + with internal_compaction_call(getattr(ctx, "agent_context", None)): + async for llm_response in self._resolve_model(ctx).generate_async( + request, + stream=False, + ctx=ctx, + ): + if llm_response.content and llm_response.content.parts: + for part in llm_response.content.parts: + if part.text: + output += part.text + output = output.strip() + if not output: + raise ValueError("AdvancedAutoCompactSummarizer returned no final content") + summary_match = re.search( + r"\s*(.*?)\s*", + output, + flags=re.DOTALL | re.IGNORECASE, + ) + if summary_match is None or not summary_match.group(1).strip(): + raise ValueError("AdvancedAutoCompactSummarizer returned no block") + return summary_match.group(1).strip() + + def _create_extractor(self, + session_memory_extractor: SessionMemoryExtractor | None = None) -> SessionMemoryExtractor: + """Create the session memory extractor.""" + if session_memory_extractor is not None: + return session_memory_extractor + return SessionMemoryExtractor( + runtime=self._runtime, + model=self._model, + ) + + @property + def session_memory_extractor(self) -> SessionMemoryExtractor: + """Return the session memory extractor.""" + return self._session_memory_extractor + + @property + def runtime(self) -> AdvancedAutoCompactSummarizerRuntime: + """Return the runtime bound to this compressor.""" + return self._runtime + + def _session_lock(self, session_id: str) -> asyncio.Lock: + """Return the unique compaction lock for a session.""" + key = self._runtime.session_key(session_id) + lock = self._session_locks.get(key) + if lock is None: + lock = asyncio.Lock() + self._session_locks[key] = lock + return lock + + async def _load_state(self, session_id: str) -> AdvancedAutoCompactState: + """Restore process-local compaction state.""" + state_key = self._runtime.session_key(session_id) + state = self._states.get(state_key) + if state is not None: + return state + state = AdvancedAutoCompactState(latest_compaction=None, consecutive_failures=0) + self._states[state_key] = state + return state + + def _summary_content(self, summary: str) -> Content: + """Wrap a compaction summary in stable model-visible user content.""" + return Content( + role="user", + parts=[Part.from_text(text=ADVANCED_AUTOCOMPACT_SUMMARY_PREFIX + summary)], + ) + + def _summary_with_recovery_path(self, summary: str) -> str: + """Tell the model where the authoritative compacted data lives.""" + return (f"{summary.rstrip()}\n\n" + "For exact content from before compaction, read the original " + "SessionService Events. Current session memory is stored in " + f"session.state[{SESSION_MEMORY_STATE_KEY!r}].") + + def _find_signature_index( + self, + contents: list[Content], + signature: str, + occurrence: int, + ) -> int | None: + """Locate a persisted compaction boundary by signature occurrence.""" + seen = 0 + for index, content in enumerate(contents): + if content_signature(contents[index]) == signature: + seen += 1 + if seen == occurrence: + return index + return None + + def _find_last_signature_index( + self, + contents: list[Content], + signature: str, + ) -> int | None: + """Find the newest matching boundary after an earlier replay.""" + for index in range(len(contents) - 1, -1, -1): + if content_signature(contents[index]) == signature: + return index + return None + + def _signature_occurrence( + self, + contents: list[Content], + signature: str, + boundary_index: int, + ) -> int: + """Count a boundary signature's occurrences from the request start.""" + return sum(1 for content in contents[:boundary_index + 1] if content_signature(content) == signature) + + def _adjust_start_for_tool_pairing(self, contents: list[Content], start: int) -> int: + """Extend the retained range to keep calls paired with responses.""" + if start <= 0 or start >= len(contents): + return max(0, start) + response_ids = { + getattr(part.function_response, "id", None) + for content in contents[start:] + for part in content.parts or [] if part.function_response is not None + } + response_ids.discard(None) + if not response_ids: + return start + for index in range(start - 1, -1, -1): + call_ids = { + getattr(part.function_call, "id", None) + for part in contents[index].parts or [] if part.function_call is not None + } + if call_ids & response_ids: + start = index + response_ids -= call_ids + if not response_ids: + break + return start + + def _compaction_start(self, contents: list[Content], boundary_index: int) -> int: + """Return the retained-content start for a legacy compaction.""" + start = min( + boundary_index + 1, + len(contents) - self._auto_compact_config.keep_recent_contents, + ) + return self._adjust_start_for_tool_pairing(contents, start) + + def _session_memory_compaction_start(self, boundary_index: int) -> int: + """Drop everything through the session-memory checkpoint boundary.""" + return boundary_index + 1 + + def _apply_record(self, request: LlmRequest, record: AdvancedAutoCompactRecord) -> bool: + """Replay a persisted compaction record into a rebuilt request.""" + boundary_index = self._find_signature_index( + request.contents, + record.boundary_signature, + record.boundary_occurrence, + ) + if boundary_index is None: + return False + start = (self._session_memory_compaction_start(boundary_index) + if record.source == "session-memory" else self._compaction_start(request.contents, boundary_index)) + request.contents = [ + self._summary_content(record.summary), + *request.contents[start:], + ] + return True + + async def _latest_session_memory_record( + self, + ctx: InvocationContext, + ) -> tuple[str, str, int, str] | None: + """Read Session Memory and its checkpoint from Session.state.""" + state = ctx.session.state + parsed = parse_session_memory_state(state.get(SESSION_MEMORY_STATE_KEY)) + if parsed is None: + return None + document, checkpoint, _ = parsed + signature = checkpoint.get("boundary_signature") + occurrence = checkpoint.get("boundary_occurrence") + event_id = checkpoint.get("last_event_id") + if (not isinstance(signature, str) or not isinstance(occurrence, int) or occurrence <= 0 + or not isinstance(event_id, str)): + return None + memory = document.to_markdown() + if memory.strip() == SessionMemoryDocument().to_markdown().strip(): + return None + return memory, signature, occurrence, event_id + + def _compact_with_summary( + self, + request: LlmRequest, + *, + summary: str, + boundary_index: int, + source: str, + strict_boundary: bool = False, + boundary_event_id: str | None = None, + ) -> AdvancedAutoCompactRecord: + """Replace the old prefix with a summary and return a replay record.""" + boundary_signature = content_signature(request.contents[boundary_index]) + boundary_occurrence = self._signature_occurrence( + request.contents, + boundary_signature, + boundary_index, + ) + start = (self._session_memory_compaction_start(boundary_index) if strict_boundary else self._compaction_start( + request.contents, boundary_index)) + request.contents = [self._summary_content(summary), *request.contents[start:]] + return AdvancedAutoCompactRecord( + boundary_signature, + boundary_occurrence, + summary, + source, + boundary_event_id, + f"advanced_autocompact:{uuid.uuid4().hex}", + ) + + def _resolve_boundary_event_id( + self, + ctx: InvocationContext, + signature: str, + occurrence: int, + ) -> str | None: + """Map one request-content boundary back to an active Session Event.""" + seen = 0 + for event in ctx.session.events: + content = getattr(event, "content", None) + if content is None or content_signature(content) != signature: + continue + seen += 1 + if seen == occurrence: + event_id = getattr(event, "id", None) + return event_id if isinstance(event_id, str) and event_id else None + return None + + def _legacy_boundary_event_id(self, ctx: InvocationContext) -> str | None: + """Choose a stable active-Event boundary for legacy compaction.""" + content_events = [event for event in ctx.session.events if event.content is not None] + if len(content_events) <= 1: + return None + keep_count = min( + self._auto_compact_config.keep_recent_contents, + len(content_events) - 1, + ) + boundary_index = len(content_events) - keep_count - 1 + start = self._compaction_start( + [event.content for event in content_events], + boundary_index, + ) + event_id = getattr(content_events[max(0, start - 1)], "id", None) + return event_id if isinstance(event_id, str) and event_id else None + + async def _persist_session_compaction( + self, + ctx: InvocationContext, + record: AdvancedAutoCompactRecord, + ) -> None: + """Persist the compacted active window through the original SessionService.""" + compact_events = ctx.session.compact_events + if not callable(compact_events): + # AdvancedAutoCompactSummarizer remains usable as a request-only primitive in unit + # tests and custom integrations. The standard Manager supplies + # the framework Session and persists the compacted window. + return + + boundary_event_id = record.boundary_event_id or self._resolve_boundary_event_id( + ctx, + record.boundary_signature, + record.boundary_occurrence, + ) + if boundary_event_id is None: + raise ValueError("Cannot map the AdvancedAutoCompactSummarizer boundary to an active Session Event") + + compaction_id = record.compaction_id or f"advanced_autocompact:{uuid.uuid4().hex}" + summary_event = Event( + invocation_id="summary", + author="system", + content=self._summary_content(record.summary), + custom_metadata={ + "session_compaction_source": record.source, + "session_compaction_boundary_signature": record.boundary_signature, + "session_compaction_boundary_occurrence": record.boundary_occurrence, + }, + ) + active_before = list(ctx.session.events) + historical_before = list(ctx.session.historical_events) + last_update_before = ctx.session.last_update_time + try: + changed = compact_events( + summary_event, + boundary_event_id, + compaction_id=compaction_id, + ) + if changed: + await ctx.session_service.update_session(ctx.session) + except Exception: + ctx.session.events = active_before + ctx.session.historical_events = historical_before + ctx.session.last_update_time = last_update_before + raise + + def _bounded_history(self, contents: list[Content]) -> str: + """Bound old history to the configured summary-input character limit.""" + rendered = "\n".join(f"\n{_content_text(content)}\n" for content in contents) + limit = self._auto_compact_config.summary_input_max_chars + if len(rendered) <= limit: + return rendered + marker = "\n...[middle of old history omitted due to the summary input limit]...\n" + first_size = max(1, (limit - len(marker)) // 3) + last_size = max(1, limit - len(marker) - first_size) + return rendered[:first_size] + marker + rendered[-last_size:] + + async def _legacy_summary( + self, + contents: list[Content], + ctx: InvocationContext, + ) -> str: + """Shrink old history across retries and generate a legacy summary.""" + retries = self._auto_compact_config.summary_retries_count + working = list(contents) + last_error: Exception | None = None + for attempt in range(retries): + try: + return await self._generate_summary( + self._bounded_history(working), + ctx, + ) + except Exception as exc: # noqa: BLE001 + last_error = exc + if len(working) <= 1: + break + drop_count = max(1, len(working) // (retries - attempt + 1)) + working = working[drop_count:] + raise RuntimeError("Legacy autocompact summary failed after retries") from last_error + + def _request_from_events(self, events: list[ResponseABC]) -> LlmRequest: + """Build the model-visible request view used by end-of-turn compaction.""" + contents: list[Content] = [] + for event in events: + is_model_visible = getattr(event, "is_model_visible", None) + if callable(is_model_visible) and not is_model_visible(): + continue + content = getattr(event, "content", None) + if content is not None: + contents.append(content.model_copy(deep=True)) + return LlmRequest(contents=contents) + + @override + async def should_summarize(self, session: SessionABC) -> bool: + """Check the character threshold without mutating the Session.""" + if not self._auto_compact_config.enabled: + return False + request = self._request_from_events(list(getattr(session, "events", []) or [])) + if len(request.contents) <= self._auto_compact_config.keep_recent_contents: + return False + + token_config = self._runtime.config.token_context_tracker + if token_config.enabled and token_config.model_context_window_tokens is not None: + effective_window = token_config.model_context_window_tokens - token_config.max_output_tokens + threshold = int(effective_window * token_config.auto_compact_ratio) + return TokenContextTracker(token_config).estimate_request_tokens(request) >= threshold + return estimate_request_chars(request) >= self._auto_compact_config.trigger_chars + + @override + async def create_session_summary_by_events( + self, + events: list[ResponseABC], + session_id: str, + keep_recent_count: int = 10, + ctx: InvocationContext | None = None, + historical_events: list[ResponseABC] | None = None, + store_historical_events: bool = False, + ) -> tuple[str | None, list[ResponseABC]]: + """Compact Events through the existing request-compaction algorithm.""" + del keep_recent_count + if ctx is None: + raise ValueError("Invocation context is required for advanced compaction") + if session_id != ctx.session_id: + raise ValueError("Session ID does not match the invocation context") + + request = self._request_from_events(events) + result = await self.apply(request, ctx=ctx, force=True) + if result.summary is not None: + events[:] = list(ctx.session.events) + if store_historical_events and historical_events is not None: + historical_events[:] = list(ctx.session.historical_events) + return result.summary, events + + @override + async def create_session_summary( + self, + session: SessionABC, + ctx: InvocationContext | None = None, + store_historical_events: bool = False, + ) -> str | None: + """Compact one Session and persist its active and historical Events.""" + if ctx is None: + raise ValueError("Invocation context is required for advanced compaction") + events = getattr(session, "events", None) + historical_events = getattr(session, "historical_events", None) + if not isinstance(events, list) or not isinstance(historical_events, list): + raise TypeError("Advanced compaction requires a Session with Event history") + summary, _ = await self.create_session_summary_by_events( + events, + session.id, + ctx=ctx, + historical_events=historical_events, + store_historical_events=store_historical_events, + ) + return summary + + @override + async def create_session_summary_by_request( + self, + request: RequestABC, + ctx: InvocationContext | None = None, + force: bool = False, + ) -> LlmResponse | None: + """Compact a built model request immediately before generation.""" + if ctx is None: + raise ValueError("Invocation context is required for advanced compaction") + if not isinstance(request, LlmRequest): + raise TypeError("Advanced compaction requires an LlmRequest") + + result = await self.apply(request, ctx=ctx, force=force) + if not result.blocked: + TokenContextTracker.record_request_context(request, ctx) + return None + return LlmResponse(content=Content( + role="model", + parts=[Part.from_text(text=ADVANCED_AUTOCOMPACT_BLOCKED_MESSAGE)], + )) + + @override + def get_summary_metadata(self) -> dict[str, object]: + """Return advanced compaction configuration metadata.""" + return { + "strategy": "advanced", + "auto_compact_enabled": self._auto_compact_config.enabled, + "trigger_chars": self._auto_compact_config.trigger_chars, + "keep_recent_contents": self._auto_compact_config.keep_recent_contents, + } + + async def apply( + self, + request: LlmRequest, + *, + ctx: InvocationContext, + force: bool = False, + ) -> AdvancedAutoCompactResult: + """Run compaction against the current session's tenant namespace.""" + session_id = ctx.session_id + if self._runtime.scope: + return await self._apply_scoped(request, session_id=session_id, ctx=ctx, force=force) + runtime = self._runtime.for_session(ctx.session) + processor = self._scoped_processors.get(runtime.scope) + if processor is None: + processor = copy.copy(self) + processor._runtime = runtime + processor._states = {} + processor._session_locks = {} + self._scoped_processors[runtime.scope] = processor + return await processor.apply(request, ctx=ctx, force=force) + + async def _apply_scoped( + self, + request: LlmRequest, + *, + session_id: str, + ctx: InvocationContext, + force: bool = False, + ) -> AdvancedAutoCompactResult: + """Replay old compaction and compact again when pressure is high.""" + config = self._runtime.config + auto_compact_config = config.auto_compact + tracker = TokenContextTracker(config.token_context_tracker) + if not auto_compact_config.enabled: + request_chars = estimate_request_chars(request) + return AdvancedAutoCompactResult(compacted=False, + reapplied=False, + blocked=False, + source=None, + request_chars_before=request_chars, + request_chars_after=request_chars, + consecutive_failures=0, + request_tokens_before=None, + request_tokens_after=None, + token_source=None) + async with self._session_lock(session_id): + request.contents = [content.model_copy(deep=True) for content in request.contents] + state = await self._load_state(session_id) + reapplied = False + if state.latest_compaction is not None: + reapplied = self._apply_record(request, state.latest_compaction) + + request_chars_before = estimate_request_chars(request) + token_budget_before = tracker.budget(request, ctx) + token_mode = token_budget_before.token_mode_enabled + request_tokens_before = token_budget_before.estimate.tokens + comparison_tokens_before = (tracker.estimate_request_tokens(request) if token_mode else None) + blocking_reached = (request_tokens_before >= token_budget_before.blocking_threshold_tokens + if token_mode else request_chars_before >= self._auto_compact_config.blocking_chars) + if state.consecutive_failures >= auto_compact_config.max_failures and blocking_reached: + return AdvancedAutoCompactResult( + compacted=False, + reapplied=reapplied, + blocked=True, + source=None, + request_chars_before=request_chars_before, + request_chars_after=request_chars_before, + consecutive_failures=state.consecutive_failures, + ) + if state.consecutive_failures >= auto_compact_config.max_failures: + return AdvancedAutoCompactResult( + compacted=False, + reapplied=reapplied, + blocked=False, + source=None, + request_chars_before=request_chars_before, + request_chars_after=request_chars_before, + consecutive_failures=state.consecutive_failures, + ) + auto_compact_reached = (request_tokens_before >= token_budget_before.auto_compact_threshold_tokens + if token_mode else request_chars_before >= auto_compact_config.trigger_chars) + if not force and not auto_compact_reached: + return AdvancedAutoCompactResult( + compacted=False, + reapplied=reapplied, + blocked=False, + source=state.latest_compaction.source if reapplied and state.latest_compaction else None, + request_chars_before=request_chars_before, + request_chars_after=request_chars_before, + consecutive_failures=state.consecutive_failures, + ) + + original_contents = [content.model_copy(deep=True) for content in request.contents] + try: + compact_record: AdvancedAutoCompactRecord | None = None + if self._session_memory_extractor is not None: + await self._session_memory_extractor.extract_if_needed( + ctx, + force=True, + ) + session_memory = await self._latest_session_memory_record(ctx, ) + if session_memory is not None: + memory, boundary_signature, boundary_occurrence, boundary_event_id = session_memory + boundary_index = self._find_signature_index( + request.contents, + boundary_signature, + boundary_occurrence, + ) + if boundary_index is None and reapplied: + boundary_index = self._find_last_signature_index( + request.contents, + boundary_signature, + ) + if boundary_index is not None: + compact_record = self._compact_with_summary( + request, + summary=self._summary_with_recovery_path(memory), + boundary_index=boundary_index, + source="session-memory", + strict_boundary=True, + boundary_event_id=boundary_event_id, + ) + if token_mode: + target_reached = (tracker.budget(request, ctx).estimate.tokens + <= token_budget_before.warning_threshold_tokens) + else: + target_reached = estimate_request_chars(request) <= auto_compact_config.target_chars + if not target_reached: + request.contents = [content.model_copy(deep=True) for content in original_contents] + compact_record = None + + if compact_record is None: + keep_count = min( + auto_compact_config.keep_recent_contents, + max(1, + len(request.contents) - 1), + ) + boundary_index = len(request.contents) - keep_count - 1 + if boundary_index < 0: + raise ValueError("Not enough model contents to compact") + summary = await self._legacy_summary( + request.contents[:boundary_index + 1], + ctx, + ) + compact_record = self._compact_with_summary( + request, + summary=self._summary_with_recovery_path(summary), + boundary_index=boundary_index, + source="legacy", + boundary_event_id=self._legacy_boundary_event_id(ctx), + ) + + request_chars_after = estimate_request_chars(request) + if token_mode: + comparison_tokens_after = tracker.estimate_request_tokens(request) + if (comparison_tokens_after >= comparison_tokens_before + and request_chars_after >= request_chars_before): + raise ValueError("Advanced Auto Compact did not reduce request token estimate") + elif request_chars_after >= request_chars_before: + raise ValueError("Advanced Auto Compact did not reduce request size") + await self._persist_session_compaction(ctx, compact_record) + state.latest_compaction = compact_record + state.consecutive_failures = 0 + return AdvancedAutoCompactResult( + compacted=True, + reapplied=reapplied, + blocked=False, + source=compact_record.source, + request_chars_before=request_chars_before, + request_chars_after=request_chars_after, + consecutive_failures=0, + request_tokens_before=comparison_tokens_before if token_mode else None, + request_tokens_after=comparison_tokens_after if token_mode else None, + token_source="estimated" if token_mode else None, + summary=compact_record.summary, + ) + except Exception as exc: # noqa: BLE001 + request.contents = original_contents + state.consecutive_failures += 1 + blocked = state.consecutive_failures >= auto_compact_config.max_failures and blocking_reached + return AdvancedAutoCompactResult( + compacted=False, + reapplied=reapplied, + blocked=blocked, + source=None, + request_chars_before=request_chars_before, + request_chars_after=request_chars_before, + consecutive_failures=state.consecutive_failures, + error=str(exc) if exc else None if blocked else None, + request_tokens_before=request_tokens_before if token_mode else None, + request_tokens_after=comparison_tokens_before if token_mode else None, + token_source=token_budget_before.estimate.source if token_mode else None, + ) + + +class AdvancedAutoCompactSummarizerHandler(BaseCompactSummarizerHandler): + """Advanced auto compact summarizer handler.""" + + @override + async def handle( + self, + ctx: InvocationContext, + request: LlmRequest, + force: bool = False, + ) -> LlmResponse | None: + """Compact before each request and return a local block after failures.""" + summarizer = self.get_summarizer(ctx) + if not isinstance(summarizer, AdvancedAutoCompactSummarizer): + raise ValueError("Summarizer is not an AdvancedAutoCompactSummarizer") + return await summarizer.create_session_summary_by_request( + request, + ctx=ctx, + force=force, + ) diff --git a/trpc_agent_sdk/sessions/compact/advanced/_base.py b/trpc_agent_sdk/sessions/compact/advanced/_base.py new file mode 100644 index 000000000..ec15b559a --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/_base.py @@ -0,0 +1,52 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. + +from __future__ import annotations + +from abc import ABC +from abc import abstractmethod +from typing import Any +from typing import Optional + +from trpc_agent_sdk.abc import CompactSummarizerABC +from trpc_agent_sdk.abc import CompactSummarizerManagerABC +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.models import LlmRequest + + +class BaseCompactSummarizerHandler(ABC): + """Base compact summarizer handler.""" + + def get_summarizer(self, ctx: InvocationContext) -> CompactSummarizerABC: + """Get the summarizer.""" + session_service = ctx.session_service + if session_service is None: + raise ValueError("Session service is not set") + summarizer_manager = getattr(session_service, "summarizer_manager", None) + if summarizer_manager is None or not isinstance(summarizer_manager, CompactSummarizerManagerABC): + raise ValueError("Summarizer manager is not an CompactSummarizerManagerABC") + return summarizer_manager.summarizer + + @abstractmethod + async def handle(self, ctx: InvocationContext, req: LlmRequest): + """Handle the compact summarizer.""" + pass + + +class BaseTokenEstimator(ABC): + """Define the replaceable token estimator interface.""" + + @abstractmethod + def estimate_payload_tokens(self, payload: Any) -> int: + """Estimate tokens for any JSON-compatible payload.""" + + +class BaseModelContextWindowResolver(ABC): + """Define the model-identifier context-window resolver interface.""" + + @abstractmethod + def resolve_context_window_tokens(self, model: Any) -> Optional[int]: + """Return the model context window, or None when unknown.""" diff --git a/trpc_agent_sdk/advanced_memory/_session_memory.py b/trpc_agent_sdk/sessions/compact/advanced/_compaction_memory_extractor.py similarity index 66% rename from trpc_agent_sdk/advanced_memory/_session_memory.py rename to trpc_agent_sdk/sessions/compact/advanced/_compaction_memory_extractor.py index 39f9ab454..1a9668336 100644 --- a/trpc_agent_sdk/advanced_memory/_session_memory.py +++ b/trpc_agent_sdk/sessions/compact/advanced/_compaction_memory_extractor.py @@ -3,38 +3,41 @@ # Copyright (C) 2026 Tencent. All rights reserved. # # tRPC-Agent-Python is licensed under Apache-2.0. -"""Maintain structured session memory with an isolated sub-agent.""" +"""Maintain structured session memory with direct model generation.""" from __future__ import annotations import json +import re from collections import Counter from dataclasses import dataclass from dataclasses import fields -import re +from dataclasses import field +from datetime import datetime +from datetime import timezone from typing import Any -from typing import Protocol -from typing import TYPE_CHECKING +from typing import Optional -from trpc_agent_sdk.agents import LlmAgent +from trpc_agent_sdk.context import InvocationContext from trpc_agent_sdk.log import logger -from trpc_agent_sdk.memory import InMemoryMemoryService -from trpc_agent_sdk.runners import Runner -from trpc_agent_sdk.sessions import InMemorySessionService +from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.models import LLMModel from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import Part +from ..._session import Session + from ._formats import SESSION_MEMORY_SECTION_DESCRIPTIONS from ._formats import SESSION_MEMORY_SECTIONS +from ._formats import SESSION_MEMORY_STATE_KEY from ._formats import SessionMemoryDocument -from ._runtime import AdvancedMemoryRuntime +from ._formats import build_session_memory_state +from ._formats import parse_session_memory_state +from ._runtime import AdvancedAutoCompactSummarizerRuntime from ._token_budget import TokenContextTracker +from ._utils import content_signature +from ._utils import internal_compaction_call -if TYPE_CHECKING: - from trpc_agent_sdk.abc import SessionABC - from trpc_agent_sdk.context import InvocationContext - -SESSION_MEMORY_CHECKPOINT_SCHEMA_VERSION = 1 _SESSION_MEMORY_FIELDS = tuple(field.name for field in fields(SessionMemoryDocument)) _SESSION_MEMORY_SECTION_GUIDANCE = "\n".join( @@ -144,7 +147,7 @@ class SessionMemoryExtractionInput: """Bundle old memory and the visible conversation context. ``new_events`` is retained only for callers using the older generator - interface. The built-in extractor merges any recovered transcript events + interface. The built-in extractor merges any recovered Session Events into ``context_messages`` and leaves this compatibility field empty. """ @@ -159,23 +162,12 @@ class SessionMemoryExtractionInput: class SessionMemoryExtractionResult: """Describe one incremental session-memory extraction.""" - extracted: bool - reason: str - processed_events: int = 0 - first_event_id: str | None = None - last_event_id: str | None = None - error: str | None = None - - -class SessionMemoryGenerator(Protocol): - """Define the replaceable session-memory generator interface.""" - - async def generate( - self, - extraction_input: SessionMemoryExtractionInput, - ctx: "InvocationContext", - ) -> SessionMemoryDocument: - """Generate a complete document from old memory and new context.""" + extracted: bool = field(default=False) + reason: str = field(default="") + processed_events: int = field(default=0) + first_event_id: Optional[str] = field(default=None) + last_event_id: Optional[str] = field(default=None) + error: Optional[str] = field(default=None) def has_session_memory_content(document: SessionMemoryDocument) -> bool: @@ -242,118 +234,105 @@ def limit(value: str, limit_chars: int = max_chars) -> str: return limited_document -class ForkedSessionMemoryGenerator: - """Call the isolated extraction Agent through a temporary Runner.""" +class SessionMemoryExtractor: + """Check thresholds and coordinate extraction, writes, and checkpoints.""" def __init__( self, - model: Any | None = None, - *, - section_max_chars: int = 8_000, - max_retries: int = 1, + runtime: AdvancedAutoCompactSummarizerRuntime, + model: LLMModel | None = None, ) -> None: - """Store an optional dedicated model, falling back to the parent model.""" - if max_retries < 0: - raise ValueError("max_retries must not be negative") + """Initialize extraction and per-session serialization locks.""" self._model = model - self._section_max_chars = section_max_chars - self._max_retries = max_retries + self._runtime = runtime + self._config = runtime.config.session_memory - def _resolve_model(self, ctx: "InvocationContext") -> Any: + def _resolve_model(self, ctx: InvocationContext) -> LLMModel: """Prefer the dedicated model, falling back to the parent Agent model.""" model = self._model or getattr(ctx.agent, "model", None) if not model: raise ValueError("Session memory extractor cannot resolve an LLM model") return model - async def generate( + async def _call_llm_model( self, extraction_input: SessionMemoryExtractionInput, - ctx: "InvocationContext", + ctx: InvocationContext, ) -> SessionMemoryDocument: - """Run extraction in a Runner isolated from the parent session and services.""" - config = ctx.agent.generate_content_config if isinstance(ctx.agent, LlmAgent) else None - agent = LlmAgent( - name="advanced_session_memory_extractor", - description="Update Markdown session memory in isolation.", - instruction=SESSION_MEMORY_INSTRUCTION, - model=self._resolve_model(ctx), - tools=[], - generate_content_config=config, - add_name_to_instruction=False, - ) - app_name = f"{ctx.app_name}_advanced_session_memory" - runner = Runner( - app_name=app_name, - agent=agent, - session_service=InMemorySessionService(), - memory_service=InMemoryMemoryService(), - enable_post_turn_processing=False, + """Generate session memory directly through the configured LLM model.""" + model = self._resolve_model(ctx) + prompt = build_session_memory_prompt( + extraction_input, + section_max_chars=self._config.section_max_chars, ) - try: - prompt = build_session_memory_prompt( - extraction_input, - section_max_chars=self._section_max_chars, - ) - parse_error: Exception | None = None - for attempt in range(self._max_retries + 1): - session = await runner.session_service.create_session( - app_name=app_name, - user_id="advanced-session-memory", - state={}, - ) - retry_instruction = "" - if parse_error is not None: - retry_instruction = ("\n\nThe previous response could not be parsed. " - f"Parser error: {parse_error}. Return the required short analysis " - "followed by all ten Markdown headings and their body text. " - "Do not return JSON, XML, or code fences.") - content = Content(role="user", parts=[Part.from_text(text=prompt + retry_instruction)]) - last_event = None - async for event in runner.run_async( - user_id=session.user_id, - session_id=session.id, - new_message=content, + parse_error: Exception | None = None + for attempt in range(self._config.max_retries + 1): + retry_instruction = "" + if parse_error is not None: + retry_instruction = ("\n\nThe previous response could not be parsed. " + f"Parser error: {parse_error}. Return the required short analysis " + "followed by all ten Markdown headings and their body text. " + "Do not return JSON, XML, or code fences.") + + request = LlmRequest( + contents=[Content( + role="user", + parts=[Part.from_text(text=prompt + retry_instruction)], + )], ) + request.append_instructions([SESSION_MEMORY_INSTRUCTION]) + + output = "" + with internal_compaction_call(getattr(ctx, "agent_context", None)): + async for response in model.generate_async( + request, + stream=False, + ctx=ctx, ): - if not event.partial: - last_event = event - - try: - if not last_event or not last_event.content or not last_event.content.parts: - raise ValueError("Session memory extractor returned no final content") - merged_text = "\n".join(part.text for part in last_event.content.parts if part.text) - return parse_session_memory_markdown(merged_text) - except Exception as exc: # noqa: BLE001 - parse_error = exc - if attempt >= self._max_retries: - raise - finally: - await runner.close() - - -class SessionMemoryExtractor: - """Check thresholds and coordinate extraction, writes, and checkpoints.""" - - def __init__( - self, - memory_runtime: AdvancedMemoryRuntime, - generator: SessionMemoryGenerator | None = None, - *, - model: Any | None = None, - ) -> None: - """Initialize extraction and per-session serialization locks.""" - if generator is not None and model is not None: - raise ValueError("Provide either generator or model, not both") - self._runtime = memory_runtime - self._generator = generator or ForkedSessionMemoryGenerator( - model, - section_max_chars=memory_runtime.config.session_memory_section_max_chars, - ) + if response.content and response.content.parts: + output += "\n".join(part.text for part in response.content.parts if part.text) - @property - def runtime(self) -> AdvancedMemoryRuntime: - """Return the runtime bound to this extractor.""" - return self._runtime + try: + if not output.strip(): + raise ValueError("Session memory extractor returned no final content") + return parse_session_memory_markdown(output) + except Exception as exc: # noqa: BLE001 + parse_error = exc + if attempt >= self._config.max_retries: + raise + + def _session_event_records(self, session: Session) -> list[dict[str, Any]]: + """Convert the authoritative Session Events into extraction records.""" + records: list[dict[str, Any]] = [] + seen: set[str] = set() + # Archived Events are no longer addressable in the active model + # request. Their information is already represented by the active + # summary Event included in the extraction context. + events = list(session.events or []) + for event in events: + if event.is_summary_event and event.is_summary_event(): + continue + if event.id in seen: + continue + seen.add(event.id) + timestamp = float(event.timestamp or 0.0) + records.append({ + "kind": + "event", + "event_id": + event.id, + "recorded_at": + datetime.fromtimestamp( + timestamp, + tz=timezone.utc, + ).isoformat(), + "event": + event.model_copy(deep=True).model_dump( + mode="json", + by_alias=True, + exclude_none=True, + ), + }) + return records def _event_records_after_checkpoint( self, @@ -361,7 +340,7 @@ def _event_records_after_checkpoint( checkpoint_event_id: str | None, checkpoint_recorded_at: str | None = None, ) -> list[dict[str, Any]]: - """Return Event transcript records after the checkpoint in order.""" + """Return Session Event records after the checkpoint in order.""" event_records = [record for record in records if record.get("kind") == "event"] if checkpoint_event_id is None: return event_records @@ -381,24 +360,14 @@ def _event_records_after_checkpoint( ) return recovered logger.warning( - "Session memory checkpoint %s is missing from transcript; " + "Session memory checkpoint %s is missing from active events; " "skipping extraction to avoid replaying the full history", checkpoint_event_id, ) return [] - def _last_checkpoint( - self, - records: list[dict[str, Any]], - ) -> dict[str, Any] | None: - """Restore the latest successful session-memory checkpoint.""" - for record in reversed(records): - if record.get("kind") == "session-memory-checkpoint" and isinstance(record.get("last_event_id"), str): - return record - return None - def _serialized_record(self, record: dict[str, Any]) -> str: - """Serialize one transcript record as stable extraction text.""" + """Serialize one Session Event record as stable extraction text.""" return json.dumps( record, ensure_ascii=False, @@ -407,27 +376,25 @@ def _serialized_record(self, record: dict[str, Any]) -> str: default=str, ) - def _context_contents(self, ctx: "InvocationContext") -> list[Any]: + def _context_contents(self, ctx: InvocationContext) -> list[Content]: """Extract model-context Content without Event metadata.""" - override_messages = getattr(ctx, "override_messages", None) + override_messages = ctx.override_messages if isinstance(override_messages, list): return [content for content in override_messages if content is not None] - contents: list[Any] = [] - session = getattr(ctx, "session", None) - for event in getattr(session, "events", []) or []: - is_model_visible = getattr(event, "is_model_visible", None) + contents: list[Content] = [] + session = ctx.session + for event in session.events or []: + is_model_visible = event.is_model_visible if callable(is_model_visible) and not is_model_visible(): continue - content = getattr(event, "content", None) + content = event.content if content is not None: contents.append(content) return contents - def _serialized_context_content(self, content: Any) -> str | None: + def _serialized_context_content(self, content: Content) -> str | None: """Serialize visible message content while excluding hidden thoughts.""" - if not hasattr(content, "model_dump"): - return None payload = content.model_dump( mode="json", by_alias=True, @@ -458,7 +425,7 @@ def _excerpt_text(self, serialized: str, limit: int) -> str: side = max(1, (limit - len(marker)) // 2) return serialized[:side] + marker + serialized[-(limit - len(marker) - side):] - def _context_messages(self, ctx: "InvocationContext") -> list[str]: + def _context_messages(self, ctx: InvocationContext) -> list[str]: """Render the complete visible conversation context in order.""" messages: list[str] = [] for content in self._context_contents(ctx): @@ -487,7 +454,7 @@ def _record_chars(self, records: list[dict[str, Any]]) -> int: return sum(len(self._serialized_record(record)) for record in records) def _count_tool_calls(self, records: list[dict[str, Any]]) -> int: - """Count model-initiated function calls in a transcript increment.""" + """Count model-initiated function calls in a Session Event increment.""" count = 0 for record in records: parts = record.get("event", {}).get("content", {}).get("parts", []) @@ -502,32 +469,39 @@ def _last_event_has_tool_call(self, records: list[dict[str, Any]]) -> bool: return self._event_has_tool_call(records[-1]) def _event_has_tool_call(self, record: dict[str, Any]) -> bool: - """Return whether one transcript Event contains a function call.""" + """Return whether one Session Event contains a function call.""" parts = record.get("event", {}).get("content", {}).get("parts", []) return any(isinstance(part, dict) and (part.get("function_call") or part.get("functionCall")) for part in parts) + def _event_has_tool_response(self, record: dict[str, Any]) -> bool: + """Return whether one Session Event contains a function response.""" + parts = record.get("event", {}).get("content", {}).get("parts", []) + return any( + isinstance(part, dict) and (part.get("function_response") or part.get("functionResponse")) + for part in parts) + def _fits_prompt_budget( self, extraction_input: SessionMemoryExtractionInput, tracker: TokenContextTracker, - ctx: "InvocationContext", + ctx: InvocationContext, ) -> bool: """Return whether the complete sub-agent prompt fits the input budget.""" prompt = build_session_memory_prompt( extraction_input, - section_max_chars=self._runtime.config.session_memory_section_max_chars, + section_max_chars=self._config.section_max_chars, ) effective_window = tracker.effective_context_window_tokens(ctx) if effective_window is not None: - limit = effective_window - self._runtime.config.session_memory_request_overhead_tokens + limit = effective_window - self._config.request_overhead_tokens return limit > 0 and tracker.estimate_payload_tokens(prompt) <= limit - return len(prompt) <= self._runtime.config.session_memory_prompt_max_chars + return len(prompt) <= self._config.prompt_max_chars def _build_extraction_input( self, current_memory: str, pending: list[dict[str, Any]], - ctx: "InvocationContext", + ctx: InvocationContext, tracker: TokenContextTracker, ) -> tuple[list[dict[str, Any]], SessionMemoryExtractionInput | None]: """Build the largest safe input that fits the extraction budget. @@ -607,19 +581,60 @@ def missing_context(end: int) -> list[str]: return [], None - async def _read_current_memory(self, session_id: str) -> str: - """Read old session memory or return the complete empty template.""" - current = await self._runtime.session_memory.read(session_id) - return current if current is not None else SessionMemoryDocument().to_markdown() + async def _read_current_memory(self, session: Session) -> str: + """Read Session Memory from the SessionService-owned state.""" + parsed = parse_session_memory_state(session.state.get(SESSION_MEMORY_STATE_KEY)) + return parsed[0].to_markdown() if parsed is not None else SessionMemoryDocument().to_markdown() + + def _state_checkpoint( + self, + session: Session, + ) -> tuple[dict[str, Any] | None, int | None]: + """Read the checkpoint and token metric from Session.state.""" + parsed = parse_session_memory_state(session.state.get(SESSION_MEMORY_STATE_KEY)) + if parsed is None: + return None, None + _, checkpoint, metrics = parsed + context_tokens = metrics.get("context_tokens") + return ( + checkpoint, + context_tokens if isinstance(context_tokens, int) else None, + ) + + def _boundary_for_event( + self, + session: Session, + event_id: str, + ) -> tuple[str, int] | None: + """Return a model-content signature and occurrence for one Event.""" + + signatures: list[str] = [] + # AutoCompact matches against the active model request, so occurrence + # counts must not include archived Events. + events = list(session.events or []) + seen_ids: set[str] = set() + for event in events: + if event.id in seen_ids: + continue + seen_ids.add(event.id) + content = event.content + if content is None: + continue + signature = content_signature(content) + signatures.append(signature) + if event.id == event_id: + return signature, signatures.count(signature) + return None async def _persist_checkpoint( self, - session_id: str, + ctx: InvocationContext, included_records: list[dict[str, Any]], document: SessionMemoryDocument, context_tokens: int | None, ) -> None: """Persist the processed increment boundary after a successful write.""" + session = ctx.session first_event_id = included_records[0]["event_id"] last_event_id = included_records[-1]["event_id"] values = ( @@ -634,39 +649,62 @@ async def _persist_checkpoint( document.key_results, document.worklog, ) - await self._runtime.transcripts.append_unique( - session_id, - { - "schema_version": SESSION_MEMORY_CHECKPOINT_SCHEMA_VERSION, - "kind": "session-memory-checkpoint", - "checkpoint_id": f"session-memory:{last_event_id}", - "first_event_id": first_event_id, - "last_event_id": last_event_id, - "processed_events": len(included_records), - "non_empty_sections": sum(1 for value in values if value.strip()), - "session_memory_chars": len(document.to_markdown()), - "context_tokens": context_tokens, - }, - unique_key="checkpoint_id", + boundary = self._boundary_for_event(session, last_event_id) + if boundary is None: + raise ValueError(f"Session Memory boundary Event {last_event_id} has no visible content") + signature, occurrence = boundary + checkpoint = { + "first_event_id": first_event_id, + "last_event_id": last_event_id, + "recorded_at": included_records[-1].get("recorded_at"), + "last_event_timestamp": included_records[-1].get("event", {}).get("timestamp"), + "boundary_signature": signature, + "boundary_occurrence": occurrence, + "processed_events": len(included_records), + "non_empty_sections": sum(1 for value in values if value.strip()), + "updated_at": datetime.now(timezone.utc).isoformat(), + } + payload = build_session_memory_state( + document, + checkpoint=checkpoint, + context_tokens=context_tokens, ) + state_delta = {SESSION_MEMORY_STATE_KEY: payload} + session_service = getattr(ctx, "session_service", None) + if session_service is None: + session.state.update(state_delta) + return + + update_state = getattr(session_service, "update_session_state", None) + if callable(update_state): + await update_state(session, state_delta) + return + + # Compatibility fallback for duck-typed SessionService implementations + # that do not inherit the latest SessionServiceABC. + session.state.update(state_delta) + await session_service.update_session(session) async def extract_if_needed( self, - session: "SessionABC", - ctx: "InvocationContext", - *, + ctx: InvocationContext, force: bool = False, ) -> SessionMemoryExtractionResult: """Update memory when the threshold or force flag is reached.""" - config = self._runtime.config - if not config.enabled or not config.session_memory_enabled: - return SessionMemoryExtractionResult(False, "disabled") - await self._runtime.initialize() - async with self._runtime.coordination.guard(session.id) as acquired: + config = self._config + session = ctx.session + if not config.enabled: + return SessionMemoryExtractionResult(reason="disabled") + runtime = self._runtime.for_session(session) + session_key = runtime.session_key(session.id) + async with self._runtime.coordination.guard( + session_key, + timeout=config.wait_timeout_seconds, + ) as acquired: if not acquired: - return SessionMemoryExtractionResult(False, "coordination-timeout") - records = await self._runtime.transcripts.read_all(session.id) - checkpoint = self._last_checkpoint(records) + return SessionMemoryExtractionResult(reason="coordination-timeout") + records = self._session_event_records(session) + checkpoint, checkpoint_context_tokens = self._state_checkpoint(session) checkpoint_event_id = checkpoint["last_event_id"] if checkpoint is not None else None checkpoint_recorded_at = checkpoint.get("recorded_at") if checkpoint is not None else None pending = self._event_records_after_checkpoint( @@ -675,49 +713,45 @@ async def extract_if_needed( checkpoint_recorded_at if isinstance(checkpoint_recorded_at, str) else None, ) if not pending: - return SessionMemoryExtractionResult(False, "no-new-events") + return SessionMemoryExtractionResult(reason="no-new-events") pending_chars = self._record_chars(pending) - tracker = TokenContextTracker(config) + tracker = TokenContextTracker(self._runtime.config.token_context_tracker) token_mode = tracker.token_mode_enabled(ctx) context_tokens = tracker.estimate_payload_tokens(self._context_contents(ctx)) - checkpoint_context_tokens = (checkpoint.get("context_tokens") if checkpoint is not None - and isinstance(checkpoint.get("context_tokens"), int) else None) - threshold = (config.session_memory_update_tokens if checkpoint_event_id is not None and token_mode else - (config.session_memory_initial_tokens if token_mode else - (config.session_memory_update_chars - if checkpoint_event_id is not None else config.session_memory_initial_chars))) + threshold = (config.update_tokens if checkpoint_event_id is not None and token_mode else + (config.initial_tokens if token_mode else + (config.update_chars if checkpoint_event_id is not None else config.initial_chars))) tool_calls = self._count_tool_calls(pending) - natural_break = not self._last_event_has_tool_call(pending) - if not natural_break: + if self._last_event_has_tool_call(pending): return SessionMemoryExtractionResult(False, "unsafe-boundary") + natural_break = not self._event_has_tool_response(pending[-1]) threshold_met = ((context_tokens >= threshold if checkpoint_context_tokens is None else (context_tokens < checkpoint_context_tokens or context_tokens - checkpoint_context_tokens >= threshold)) if token_mode else pending_chars >= threshold) - tool_condition_met = tool_calls >= config.session_memory_tool_calls_between_updates or natural_break + tool_condition_met = tool_calls >= config.tool_calls_between_updates or natural_break if not force and (not threshold_met or not tool_condition_met): - return SessionMemoryExtractionResult(False, "threshold-not-met") + return SessionMemoryExtractionResult(reason="threshold-not-met") included, extraction_input = self._build_extraction_input( - await self._read_current_memory(session.id), + await self._read_current_memory(session), pending, ctx, tracker, ) if extraction_input is None: - return SessionMemoryExtractionResult(False, "context-unavailable") + return SessionMemoryExtractionResult(reason="context-unavailable") try: - document = await self._generator.generate(extraction_input, ctx) + document = await self._call_llm_model(extraction_input, ctx) if not has_session_memory_content(document): raise ValueError("Session memory generator returned an all-empty document") document = limit_session_memory_document( document, - max_chars=config.session_memory_section_max_chars, - total_max_chars=config.session_memory_total_max_chars, + max_chars=config.section_max_chars, + total_max_chars=config.total_max_chars, ) - await self._runtime.session_memory.write(session.id, document) await self._persist_checkpoint( - session.id, + ctx, included, document, context_tokens, @@ -730,8 +764,7 @@ async def extract_if_needed( exc_info=True, ) return SessionMemoryExtractionResult( - False, - "extraction-failed", + reason="extraction-failed", processed_events=0, first_event_id=included[0]["event_id"], last_event_id=included[-1]["event_id"], @@ -739,8 +772,8 @@ async def extract_if_needed( ) return SessionMemoryExtractionResult( - True, - "forced" if force else "threshold-met", + extracted=True, + reason="forced" if force else "threshold-met", processed_events=len(included), first_event_id=included[0]["event_id"], last_event_id=included[-1]["event_id"], diff --git a/trpc_agent_sdk/sessions/compact/advanced/_config.py b/trpc_agent_sdk/sessions/compact/advanced/_config.py new file mode 100644 index 000000000..35f16b0d5 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/_config.py @@ -0,0 +1,174 @@ +# Tencent is pleased to support the open source ecosystem. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# Licensed under Apache-2.0. +"""Configuration for Session Compact.""" + +from __future__ import annotations + +from dataclasses import dataclass +from dataclasses import field + +from ._base import BaseTokenEstimator +from ._base import BaseModelContextWindowResolver + +DEFAULT_COMPACTABLE_TOOL_NAMES = ( + "Read", + "Bash", + "Grep", + "Glob", + "Search", + "CodeSearch", +) + + +def _require_positive(**values: int | float) -> None: + for name, value in values.items(): + if value <= 0: + raise ValueError(f"{name} must be greater than zero") + + +def _require_non_negative(**values: int | float) -> None: + for name, value in values.items(): + if value < 0: + raise ValueError(f"{name} must be non-negative") + + +def _require_non_empty_names(name: str, values: tuple[str, ...]) -> None: + if not values or any(not value.strip() for value in values): + raise ValueError(f"{name} must contain non-empty names") + + +@dataclass(frozen=True) +class AutoCompactSummarizerConfig: + """Configure auto compact summarizer.""" + enabled: bool = field(default=True) + trigger_chars: int = field(default=700_000) + target_chars: int = field(default=350_000) + blocking_chars: int = field(default=780_000) + keep_recent_contents: int = field(default=8) + max_failures: int = field(default=3) + summary_input_max_chars: int = field(default=600_000) + summary_retries_count: int = field(default=3) + + +@dataclass(frozen=True) +class HistorySnipConfig: + """Configure history snip.""" + enabled: bool = field(default=True) + trigger_chars: int = field(default=600_000) + target_chars: int = field(default=400_000) + keep_recent: int = field(default=5) + tool_names: tuple[str, ...] = field(default=DEFAULT_COMPACTABLE_TOOL_NAMES) + + +@dataclass(frozen=True) +class TokenContextTrackerConfig: + """Configure token context tracker.""" + enabled: bool = field(default=True) + warning_ratio: float = field(default=0.85) + auto_compact_ratio: float = field(default=0.90) + blocking_ratio: float = field(default=0.95) + model_context_window_tokens: int | None = field(default=None) + max_output_tokens: int = field(default=0) + estimator: BaseTokenEstimator | None = field(default=None) + context_window_resolver: BaseModelContextWindowResolver | None = field(default=None) + + +@dataclass(frozen=True) +class MicroCompactConfig: + """Configure micro compact.""" + enabled: bool = field(default=True) + gap_seconds: float = field(default=3_600.0) + trigger_count: int = field(default=20) + keep_recent: int = field(default=5) + tool_names: tuple[str, ...] = field(default=DEFAULT_COMPACTABLE_TOOL_NAMES) + + +@dataclass(frozen=True) +class ToolResultBudgetConfig: + """Configure tool result budget.""" + enabled: bool = field(default=True) + max_chars: int = field(default=50_000) + per_message_max_chars: int = field(default=200_000) + preview_chars: int = field(default=2_000) + + +@dataclass(frozen=True) +class SessionMemoryExtractorConfig: + """Configure session memory.""" + enabled: bool = field(default=True) + initial_chars: int = field(default=40_000) + update_chars: int = field(default=20_000) + initial_tokens: int = field(default=10_000) + update_tokens: int = field(default=5_000) + tool_calls_between_updates: int = field(default=3) + prompt_max_chars: int = field(default=200_000) + request_overhead_tokens: int = field(default=2_048) + section_max_chars: int = field(default=8_000) + total_max_chars: int = field(default=54_000) + wait_timeout_seconds: float = field(default=15.0) + max_retries: int = field(default=1) + + +@dataclass(frozen=True) +class AdvancedAutoCompactSummarizerConfig: + """Configure advanced compact.""" + history_snip: HistorySnipConfig = field(default_factory=HistorySnipConfig) + token_context_tracker: TokenContextTrackerConfig = field(default_factory=TokenContextTrackerConfig) + session_memory: SessionMemoryExtractorConfig = field(default_factory=SessionMemoryExtractorConfig) + tool_result_budget: ToolResultBudgetConfig = field(default_factory=ToolResultBudgetConfig) + micro_compact: MicroCompactConfig = field(default_factory=MicroCompactConfig) + auto_compact: AutoCompactSummarizerConfig = field(default_factory=AutoCompactSummarizerConfig) + + def __post_init__(self) -> None: + """Validate compression limits and token thresholds.""" + _require_positive( + tool_result_budget_max_chars=self.tool_result_budget.max_chars, + tool_result_budget_per_message_max_chars=self.tool_result_budget.per_message_max_chars, + tool_result_budget_preview_chars=self.tool_result_budget.preview_chars, + history_snip_trigger_chars=self.history_snip.trigger_chars, + history_snip_target_chars=self.history_snip.target_chars, + history_snip_keep_recent=self.history_snip.keep_recent, + session_memory_initial_chars=self.session_memory.initial_chars, + session_memory_update_chars=self.session_memory.update_chars, + session_memory_initial_tokens=self.session_memory.initial_tokens, + session_memory_update_tokens=self.session_memory.update_tokens, + session_memory_tool_calls_between_updates=self.session_memory.tool_calls_between_updates, + session_memory_prompt_max_chars=self.session_memory.prompt_max_chars, + session_memory_section_max_chars=self.session_memory.section_max_chars, + session_memory_total_max_chars=self.session_memory.total_max_chars, + session_memory_wait_timeout_seconds=self.session_memory.wait_timeout_seconds, + auto_compact_trigger_chars=self.auto_compact.trigger_chars, + auto_compact_target_chars=self.auto_compact.target_chars, + auto_compact_blocking_chars=self.auto_compact.blocking_chars, + auto_compact_keep_recent_contents=self.auto_compact.keep_recent_contents, + auto_compact_max_failures=self.auto_compact.max_failures, + auto_compact_summary_input_max_chars=self.auto_compact.summary_input_max_chars, + auto_compact_summary_retries_count=self.auto_compact.summary_retries_count, + micro_compact_gap_seconds=self.micro_compact.gap_seconds, + micro_compact_trigger_count=self.micro_compact.trigger_count, + micro_compact_keep_recent=self.micro_compact.keep_recent, + ) + _require_non_negative( + max_output_tokens=self.token_context_tracker.max_output_tokens, + session_memory_request_overhead_tokens=self.session_memory.request_overhead_tokens, + ) + context_window_tokens = self.token_context_tracker.model_context_window_tokens + if context_window_tokens is not None: + _require_positive(model_context_window_tokens=context_window_tokens) + if self.token_context_tracker.max_output_tokens >= context_window_tokens: + raise ValueError("max_output_tokens must be smaller than model_context_window_tokens") + if not (0 < self.token_context_tracker.warning_ratio < self.token_context_tracker.auto_compact_ratio < + self.token_context_tracker.blocking_ratio < 1): + raise ValueError("token context tracker ratios must satisfy 0 < warning < auto compact < blocking < 1") + if self.tool_result_budget.preview_chars >= self.tool_result_budget.max_chars: + raise ValueError("tool_result_budget.preview_chars must be smaller than tool_result_budget.max_chars") + if self.history_snip.target_chars >= self.history_snip.trigger_chars: + raise ValueError("history_snip.target_chars must be smaller than history_snip.trigger_chars") + if self.auto_compact.trigger_chars <= self.auto_compact.target_chars: + raise ValueError("auto_compact.trigger_chars must be greater than auto_compact.target_chars") + if self.auto_compact.blocking_chars <= self.auto_compact.trigger_chars: + raise ValueError("auto_compact.blocking_chars must be greater than auto_compact.trigger_chars") + _require_non_empty_names("history_snip.tool_names", self.history_snip.tool_names) + _require_non_empty_names("micro_compact.tool_names", self.micro_compact.tool_names) diff --git a/trpc_agent_sdk/advanced_memory/_coordination.py b/trpc_agent_sdk/sessions/compact/advanced/_coordination.py similarity index 100% rename from trpc_agent_sdk/advanced_memory/_coordination.py rename to trpc_agent_sdk/sessions/compact/advanced/_coordination.py diff --git a/trpc_agent_sdk/sessions/compact/advanced/_filters.py b/trpc_agent_sdk/sessions/compact/advanced/_filters.py new file mode 100644 index 000000000..fb7f1a0cc --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/_filters.py @@ -0,0 +1,80 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. + +from typing import Any +from typing_extensions import override + +from trpc_agent_sdk.abc import CompactTrigger +from trpc_agent_sdk.context import AgentContext +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.context import get_invocation_ctx +from trpc_agent_sdk.filter import BaseFilter +from trpc_agent_sdk.filter import FilterResult +from trpc_agent_sdk.filter import FilterType + +from ._manager import AdvancedAutoCompactSummarizerManager +from ._utils import INTERNAL_COMPACTION_METADATA_KEY + + +class AdvancedAutoCompactSummarizerFilter(BaseFilter): + """Advanced auto compact summarizer filter.""" + + def __init__(self) -> None: + """Initialize the advanced auto compact summarizer filter.""" + super().__init__() + self.name = "advanced_auto_compact_summarizer_filter" + self.type = FilterType.MODEL + + def get_summarizer_manager(self, ctx: InvocationContext) -> AdvancedAutoCompactSummarizerManager: + """Get the summarizer.""" + session_service = ctx.session_service + if session_service is None: + raise ValueError("Session service is not set") + summarizer_manager = getattr(session_service, "summarizer_manager", None) + if summarizer_manager is None or not isinstance(summarizer_manager, AdvancedAutoCompactSummarizerManager): + raise ValueError("Summarizer manager is not an AdvancedAutoCompactSummarizerManager") + return summarizer_manager + + @override + async def _before(self, ctx: AgentContext, req: Any, rsp: FilterResult): + """Run the advanced auto compact summarizer filter.""" + if ctx.get_metadata(INTERNAL_COMPACTION_METADATA_KEY, False): + return None + invocation_ctx: InvocationContext = get_invocation_ctx() + summarizer_manager = self.get_summarizer_manager(invocation_ctx) + result = await summarizer_manager.create_session_summary_before_model( + req, + invocation_ctx, + ) + if not result: + return None + invocation_ctx.end_invocation = True + rsp.rsp = result + rsp.is_continue = False + rsp.error = None + session_memory_extractor = summarizer_manager.get_session_memory_extractor() + if session_memory_extractor is not None: + await session_memory_extractor.extract_if_needed( + invocation_ctx, + force=False, + ) + return + + @override + async def _after(self, ctx: AgentContext, req: Any, rsp: FilterResult): + """Run the advanced auto compact summarizer filter.""" + if ctx.get_metadata(INTERNAL_COMPACTION_METADATA_KEY, False): + return None + invocation_ctx: InvocationContext = get_invocation_ctx() + summarizer_manager = self.get_summarizer_manager(invocation_ctx) + if summarizer_manager.compact_trigger != CompactTrigger.BEFORE_MODEL: + return None + session_memory_extractor = summarizer_manager.get_session_memory_extractor() + if session_memory_extractor is not None: + await session_memory_extractor.extract_if_needed( + invocation_ctx, + force=False, + ) diff --git a/trpc_agent_sdk/sessions/compact/advanced/_formats.py b/trpc_agent_sdk/sessions/compact/advanced/_formats.py new file mode 100644 index 000000000..05a19c965 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/_formats.py @@ -0,0 +1,131 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Define shared formats for long-term and session memory.""" + +# flake8: noqa: E125 + +from __future__ import annotations + +from dataclasses import asdict +from dataclasses import dataclass +from dataclasses import fields + +SESSION_MEMORY_SECTIONS = ( + "Session Title", + "Current State", + "Task specification", + "Files and Functions", + "Workflow", + "Errors & Corrections", + "Codebase and System Documentation", + "Learnings", + "Key results", + "Worklog", +) +SESSION_MEMORY_STATE_KEY = "_trpc_agent:summary" +SESSION_MEMORY_STATE_SCHEMA_VERSION = 1 + +SESSION_MEMORY_SECTION_DESCRIPTIONS = ( + "A short and distinctive 5-10 word descriptive title for the session", + "What is actively being worked on right now? Pending tasks not yet completed.", + "What did the user ask to build? Any design decisions or other explanatory context", + "What are the important files? In short, what do they contain?", + "What bash commands are usually run and in what order?", + "Errors encountered and how they were fixed. What approaches failed?", + "What are the important system components? How do they work/fit together?", + "What has worked well? What has not? What to avoid?", + "If the user asked a specific output, repeat the exact result here", + "Step by step, what was attempted, done? Very terse summary", +) + + +@dataclass(frozen=True) +class SessionMemoryDocument: + """Represent structured session memory with ten fixed sections.""" + + session_title: str = "" + current_state: str = "" + task_specification: str = "" + files_and_functions: str = "" + workflow: str = "" + errors_and_corrections: str = "" + codebase_and_system_documentation: str = "" + learnings: str = "" + key_results: str = "" + worklog: str = "" + + def to_markdown(self) -> str: + """Render all sections in fixed order, including empty sections.""" + values = ( + self.session_title, + self.current_state, + self.task_specification, + self.files_and_functions, + self.workflow, + self.errors_and_corrections, + self.codebase_and_system_documentation, + self.learnings, + self.key_results, + self.worklog, + ) + sections = [ + f"# {section}\n_{description}_\n\n{value.strip()}" for section, description, value in zip( + SESSION_MEMORY_SECTIONS, + SESSION_MEMORY_SECTION_DESCRIPTIONS, + values, + ) + ] + return "\n\n".join(sections).rstrip() + "\n" + + +def build_session_memory_state( + document: SessionMemoryDocument, + *, + checkpoint: dict[str, object], + context_tokens: int | None, +) -> dict[str, object]: + """Build the versioned Session.state payload used by Redis and SQL.""" + return { + "schema_version": SESSION_MEMORY_STATE_SCHEMA_VERSION, + "document": asdict(document), + "checkpoint": checkpoint, + "metrics": { + "session_memory_chars": len(document.to_markdown()), + "context_tokens": context_tokens, + }, + } + + +def parse_session_memory_state( + value: object, ) -> tuple[SessionMemoryDocument, dict[str, object], dict[str, object]] | None: + """Parse a persisted Session Memory state value.""" + if not isinstance(value, dict): + return None + if value.get("schema_version") != SESSION_MEMORY_STATE_SCHEMA_VERSION: + return None + raw_document = value.get("document") + raw_checkpoint = value.get("checkpoint") + raw_metrics = value.get("metrics", {}) + if not isinstance(raw_document, dict) or not isinstance(raw_checkpoint, dict): + return None + if (not isinstance(raw_checkpoint.get("last_event_id"), str) + or not isinstance(raw_checkpoint.get("boundary_signature"), str) + or not isinstance(raw_checkpoint.get("boundary_occurrence"), int)): + return None + if not isinstance(raw_metrics, dict): + raw_metrics = {} + allowed = {field.name for field in fields(SessionMemoryDocument)} + if (any(key not in allowed for key in raw_document) + or any(not isinstance(item, str) for item in raw_document.values())): + return None + try: + document = SessionMemoryDocument(**{ + key: item + for key, item in raw_document.items() if key in allowed and isinstance(item, str) + }) + except TypeError: + return None + return document, dict(raw_checkpoint), dict(raw_metrics) diff --git a/trpc_agent_sdk/advanced_memory/_history_snip.py b/trpc_agent_sdk/sessions/compact/advanced/_history_snip.py similarity index 69% rename from trpc_agent_sdk/advanced_memory/_history_snip.py rename to trpc_agent_sdk/sessions/compact/advanced/_history_snip.py index 67d8ff1a6..214f17bbb 100644 --- a/trpc_agent_sdk/advanced_memory/_history_snip.py +++ b/trpc_agent_sdk/sessions/compact/advanced/_history_snip.py @@ -8,25 +8,23 @@ from __future__ import annotations import asyncio +import copy import json from dataclasses import dataclass from typing import Any -from typing import TYPE_CHECKING +from typing_extensions import override -from ._callbacks import install_staged_callback -from ._runtime import AdvancedMemoryRuntime +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.models import LlmRequest + +from ._base import BaseCompactSummarizerHandler +from ._runtime import AdvancedAutoCompactSummarizerRuntime from ._tool_result_budget import is_budget_replacement_response from ._tool_result_budget import serialize_tool_response from ._tool_result_budget import stable_tool_result_id from ._tool_result_budget import tool_result_sha256 from ._token_budget import TokenContextTracker -if TYPE_CHECKING: - from trpc_agent_sdk.agents import LlmAgent - from trpc_agent_sdk.context import InvocationContext - from trpc_agent_sdk.models import LlmRequest - -HISTORY_SNIP_SCHEMA_VERSION = 1 HISTORY_SNIP_CLEARED_MESSAGE = "[Older tool result removed by history snip]" @@ -64,7 +62,7 @@ class HistorySnipResult: token_source: str | None = None -def estimate_request_chars(request: "LlmRequest") -> int: +def estimate_request_chars(request: LlmRequest) -> int: """Estimate the full model request using stable JSON serialization.""" payload = request.model_dump( mode="python", @@ -83,51 +81,44 @@ def estimate_request_chars(request: "LlmRequest") -> int: class HistorySnip: """Mechanically remove the oldest tool results when the request is too large.""" - def __init__(self, memory_runtime: AdvancedMemoryRuntime) -> None: + def __init__(self, runtime: AdvancedAutoCompactSummarizerRuntime) -> None: """Initialize history-snip state and per-session async locks.""" - self._runtime = memory_runtime + self._runtime = runtime + self._config = runtime.config.history_snip self._states: dict[str, HistorySnipState] = {} self._session_locks: dict[str, asyncio.Lock] = {} + self._scoped_processors: dict[object, "HistorySnip"] = {} @property - def runtime(self) -> AdvancedMemoryRuntime: + def runtime(self) -> AdvancedAutoCompactSummarizerRuntime: """Return the runtime bound to this history snipper.""" return self._runtime def _session_lock(self, session_id: str) -> asyncio.Lock: """Return the unique history-snip lock for a session.""" - lock = self._session_locks.get(session_id) + key = self._runtime.session_key(session_id) if hasattr(self._runtime, "session_key") else session_id + lock = self._session_locks.get(key) if lock is None: lock = asyncio.Lock() - self._session_locks[session_id] = lock + self._session_locks[key] = lock return lock async def _load_state(self, session_id: str) -> HistorySnipState: - """Restore prior history-snip decisions from the transcript.""" - state = self._states.get(session_id) + """Return process-local history-snip state.""" + state_key = self._runtime.session_key(session_id) if hasattr(self._runtime, "session_key") else session_id + state = self._states.get(state_key) if state is not None: return state - records = await self._runtime.transcripts.read_all(session_id) - snipped_ids: set[str] = set() - result_hashes: dict[str, str] = {} - for record in records: - result_id = record.get("result_id") - if record.get("kind") != "history-snip" or not isinstance(result_id, str): - continue - snipped_ids.add(result_id) - original_sha256 = record.get("original_sha256") - if isinstance(original_sha256, str): - result_hashes[result_id] = original_sha256 state = HistorySnipState( - snipped_ids=snipped_ids, - result_hashes=result_hashes, + snipped_ids=set(), + result_hashes={}, ) - self._states[session_id] = state + self._states[state_key] = state return state def _collect_candidates(self, request: "LlmRequest") -> list[HistorySnipCandidate]: """Collect eligible function results in request order.""" - allowed_tools = set(self._runtime.config.history_snip_tool_names) + allowed_tools = set(self._config.tool_names) candidates: list[HistorySnipCandidate] = [] serialized_by_result_id: dict[str, str] = {} for content in request.contents: @@ -158,45 +149,42 @@ def _snipped_response(self) -> dict[str, str]: """Return the stable placeholder used by history snip.""" return {"output": HISTORY_SNIP_CLEARED_MESSAGE} - async def _persist_snip( - self, - session_id: str, - candidate: HistorySnipCandidate, - trigger: str, - ) -> None: - """Persist the history-snip decision to the transcript.""" - await self._runtime.transcripts.append_unique( - session_id, - { - "schema_version": HISTORY_SNIP_SCHEMA_VERSION, - "kind": "history-snip", - "snip_id": f"history-snip:{candidate.result_id}", - "result_id": candidate.result_id, - "tool_name": candidate.tool_name, - "original_chars": candidate.original_size, - "original_sha256": candidate.original_sha256, - "trigger": trigger, - "snipped_response": self._snipped_response(), - }, - unique_key="snip_id", - ) - async def apply( self, - request: "LlmRequest", + request: LlmRequest, *, - session_id: str, - ctx: "InvocationContext | None" = None, + ctx: InvocationContext, force: bool = False, ) -> HistorySnipResult: """Clean old tool results when over budget or explicitly forced.""" - config = self._runtime.config - tracker = TokenContextTracker(config) - if not config.enabled or not config.history_snip_enabled: + if not self._config.enabled: request_chars = estimate_request_chars(request) return HistorySnipResult(None, 0, 0, 0, request_chars, request_chars) - - await self._runtime.initialize() + session_id = ctx.session_id + if self._runtime.scope: + return await self._apply_scoped(request, session_id=session_id, ctx=ctx, force=force) + runtime = self._runtime.for_session(ctx.session) + processor = self._scoped_processors.get(runtime.scope) + if processor is None: + processor = copy.copy(self) + processor._runtime = runtime + processor._states = {} + processor._session_locks = {} + self._scoped_processors[runtime.scope] = processor + return await processor.apply(request, ctx=ctx, force=force) + + async def _apply_scoped( + self, + request: LlmRequest, + *, + session_id: str, + ctx: InvocationContext, + force: bool, + ) -> HistorySnipResult: + """Apply one tenant-bound history-snipping operation.""" + token_context_tracker_config = self._runtime.config.token_context_tracker + history_snip_config = self._config + tracker = TokenContextTracker(token_context_tracker_config) async with self._session_lock(session_id): request.contents = [content.model_copy(deep=True) for content in request.contents] state = await self._load_state(session_id) @@ -217,7 +205,7 @@ async def apply( token_mode = token_budget_before.token_mode_enabled current_tokens = token_budget_before.estimate.tokens if not force and (current_tokens <= token_budget_before.warning_threshold_tokens - if token_mode else request_chars_before <= config.history_snip_trigger_chars): + if token_mode else request_chars_before <= history_snip_config.trigger_chars): return HistorySnipResult( None, 0, @@ -231,7 +219,7 @@ async def apply( ) trigger = "force" if force else "pressure" - protected_ids = {candidate.result_id for candidate in candidates[-config.history_snip_keep_recent:]} + protected_ids = {candidate.result_id for candidate in candidates[-history_snip_config.keep_recent:]} eligible = [ candidate for candidate in candidates if candidate.result_id not in state.snipped_ids and candidate.result_id not in protected_ids @@ -241,12 +229,11 @@ async def apply( snipped_count = 0 for candidate in eligible: if not force and (current_tokens <= token_budget_before.warning_threshold_tokens - if token_mode else current_chars <= config.history_snip_target_chars): + if token_mode else current_chars <= history_snip_config.target_chars): break candidate_saving = max(0, candidate.original_size - replacement_size) if candidate_saving == 0: continue - await self._persist_snip(session_id, candidate, trigger) candidate.part.function_response.response = self._snipped_response() state.snipped_ids.add(candidate.result_id) state.result_hashes[candidate.result_id] = candidate.original_sha256 @@ -270,39 +257,20 @@ async def apply( ) -class HistorySnipCallback: +class HistorySnipHandler(BaseCompactSummarizerHandler): """Adapt history snip to before_model_callback.""" - advanced_memory_stage = 20 - - def __init__(self, history_snip: HistorySnip) -> None: + def __init__(self) -> None: """Store the history-snip processor run before model requests.""" - self._history_snip = history_snip - - @property - def history_snip(self) -> HistorySnip: - """Return the history-snip processor used by this callback.""" - return self._history_snip + self._history_snip: HistorySnip | None = None - async def __call__(self, ctx: "InvocationContext", request: "LlmRequest") -> None: + @override + async def handle(self, ctx: InvocationContext, request: LlmRequest) -> None: """Run history snip before a request based on request size.""" - await self._history_snip.apply(request, session_id=ctx.session_id, ctx=ctx) - return None - - -def setup_history_snip( - agent: "LlmAgent", - memory_runtime: AdvancedMemoryRuntime, -) -> HistorySnip: - """Install history snip while preserving context stage order.""" - history_snip = HistorySnip(memory_runtime) - callback = HistorySnipCallback(history_snip) - existing_snip = install_staged_callback( - agent, - callback, - callback_type=HistorySnipCallback, - component_attribute="history_snip", - memory_runtime=memory_runtime, - conflict_message="History snip is already configured with another runtime", - ) - return existing_snip or history_snip + summarizer = self.get_summarizer(ctx) + from ._auto_compact import AdvancedAutoCompactSummarizer + if not isinstance(summarizer, AdvancedAutoCompactSummarizer): + raise ValueError("Summarizer is not an AdvancedAutoCompactSummarizer") + if self._history_snip is None: + self._history_snip = HistorySnip(summarizer.runtime) + await self._history_snip.apply(request, ctx=ctx) diff --git a/trpc_agent_sdk/sessions/compact/advanced/_manager.py b/trpc_agent_sdk/sessions/compact/advanced/_manager.py new file mode 100644 index 000000000..8759a5e95 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/_manager.py @@ -0,0 +1,101 @@ +# Tencent is pleased to support the open source community by making +# contributions to the open source ecosystem. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Integrate Session Compact with the native SessionService lifecycle.""" + +from __future__ import annotations + +from typing_extensions import override + +from trpc_agent_sdk.abc import CompactSummarizerManagerABC +from trpc_agent_sdk.abc import CompactTrigger +from trpc_agent_sdk.abc import RequestABC +from trpc_agent_sdk.abc import ResponseABC +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.models import LlmRequest + +from ..._session import Session +from ._auto_compact import AdvancedAutoCompactSummarizer +from ._formats import parse_session_memory_state +from ._formats import SESSION_MEMORY_STATE_KEY +from ._compaction_memory_extractor import SessionMemoryExtractor +from ._history_snip import HistorySnipHandler +from ._micro_compact import MicroCompactHandler +from ._tool_result_budget import ToolResultBudgetHandler +from ._auto_compact import AdvancedAutoCompactSummarizerHandler + + +class AdvancedAutoCompactSummarizerManager(CompactSummarizerManagerABC): + """Coordinate Advanced Compact state without wrapping a SessionService.""" + + def __init__( + self, + summarizer: AdvancedAutoCompactSummarizer, + compact_trigger: CompactTrigger = CompactTrigger.BEFORE_MODEL, + ) -> None: + """Store configuration until Runner supplies the Agent.""" + super().__init__(summarizer, compact_trigger=compact_trigger) + self._tool_result_budget_handler = ToolResultBudgetHandler() + self._history_snip_handler = HistorySnipHandler() + self._micro_compact_handler = MicroCompactHandler() + self._advanced_auto_compact_summarizer_handler = AdvancedAutoCompactSummarizerHandler() + + def get_session_memory_extractor(self) -> SessionMemoryExtractor: + """Get the session memory extractor.""" + return self.summarizer.session_memory_extractor + + @override + async def create_session_summary( + self, + session: Session, + force: bool = False, + ctx: InvocationContext | None = None, + ) -> None: + """Compact persisted Events when configured for end-of-turn execution.""" + if self.compact_trigger != CompactTrigger.AFTER_TURN: + return + if ctx is None: + raise ValueError("Invocation context is required for advanced compaction") + session_memory_extractor = self.get_session_memory_extractor() + if session_memory_extractor is not None: + await session_memory_extractor.extract_if_needed( + ctx, + force=False, + ) + if force or await self.summarizer.should_summarize(session): + await self.summarizer.create_session_summary( + session, + ctx=ctx, + store_historical_events=True, + ) + + @override + async def create_session_summary_before_model( + self, + request: RequestABC, + ctx: InvocationContext, + force: bool = False, + ) -> ResponseABC | None: + """Run the advanced request pipeline immediately before model generation.""" + if self.compact_trigger != CompactTrigger.BEFORE_MODEL: + return None + if not isinstance(request, LlmRequest): + raise TypeError("Advanced compaction requires an LlmRequest") + await self._tool_result_budget_handler.handle(ctx, request) + await self._history_snip_handler.handle(ctx, request) + await self._micro_compact_handler.handle(ctx, request) + return await self._advanced_auto_compact_summarizer_handler.handle( + ctx, + request, + force=force, + ) + + async def get_session_summary(self, session: Session) -> str | None: + """Read compact Session Memory through the existing summary API.""" + parsed = parse_session_memory_state(session.state.get(SESSION_MEMORY_STATE_KEY)) + if parsed is not None: + return parsed[0].to_markdown() + return None diff --git a/trpc_agent_sdk/advanced_memory/_microcompact.py b/trpc_agent_sdk/sessions/compact/advanced/_micro_compact.py similarity index 55% rename from trpc_agent_sdk/advanced_memory/_microcompact.py rename to trpc_agent_sdk/sessions/compact/advanced/_micro_compact.py index eeaabdd36..da26d9295 100644 --- a/trpc_agent_sdk/advanced_memory/_microcompact.py +++ b/trpc_agent_sdk/sessions/compact/advanced/_micro_compact.py @@ -8,29 +8,29 @@ from __future__ import annotations import asyncio +import copy import time from dataclasses import dataclass -from typing import Any -from typing import TYPE_CHECKING +from dataclasses import field +from typing import Optional +from typing_extensions import override -from ._callbacks import install_staged_callback -from ._runtime import AdvancedMemoryRuntime +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.types import Part + +from ._base import BaseCompactSummarizerHandler +from ._runtime import AdvancedAutoCompactSummarizerRuntime from ._tool_result_budget import is_budget_replacement_response from ._tool_result_budget import serialize_tool_response from ._tool_result_budget import stable_tool_result_id from ._tool_result_budget import tool_result_sha256 -if TYPE_CHECKING: - from trpc_agent_sdk.agents import LlmAgent - from trpc_agent_sdk.context import InvocationContext - from trpc_agent_sdk.models import LlmRequest - -MICROCOMPACT_SCHEMA_VERSION = 1 MICROCOMPACT_CLEARED_MESSAGE = "[Old tool result content cleared]" @dataclass -class MicrocompactState: +class MicroCompactState: """Store identifiers for mechanically cleaned tool results.""" cleared_ids: set[str] @@ -38,24 +38,24 @@ class MicrocompactState: @dataclass(frozen=True) -class MicrocompactCandidate: +class MicroCompactCandidate: """Describe a function response eligible for mechanical cleanup.""" result_id: str tool_name: str original_size: int original_sha256: str - part: Any + part: Part @dataclass(frozen=True) -class MicrocompactResult: +class MicroCompactResult: """Summarize new and repeated mechanical cleanup operations.""" - trigger: str | None - cleared_count: int - reapplied_count: int - chars_saved: int + trigger: Optional[str] = field(default=None) + cleared_count: int = field(default=0) + reapplied_count: int = field(default=0) + chars_saved: int = field(default=0) def find_last_assistant_timestamp(ctx: "InvocationContext") -> float | None: @@ -71,55 +71,43 @@ def find_last_assistant_timestamp(ctx: "InvocationContext") -> float | None: return None -class Microcompact: +class MicroCompact: """Local compressor that cleans old tool results by time or count.""" - def __init__(self, memory_runtime: AdvancedMemoryRuntime) -> None: + def __init__(self, runtime: AdvancedAutoCompactSummarizerRuntime) -> None: """Initialize mechanical-compaction state and per-session locks.""" - self._runtime = memory_runtime - self._states: dict[str, MicrocompactState] = {} + self._runtime = runtime + self._config = runtime.config.micro_compact + self._states: dict[str, MicroCompactState] = {} self._session_locks: dict[str, asyncio.Lock] = {} - - @property - def runtime(self) -> AdvancedMemoryRuntime: - """Return the runtime bound to this mechanical compressor.""" - return self._runtime + self._scoped_processors: dict[object, MicroCompact] = {} def _session_lock(self, session_id: str) -> asyncio.Lock: """Return the unique async compaction lock for a session.""" - lock = self._session_locks.get(session_id) + key = self._runtime.session_key(session_id) if hasattr(self._runtime, "session_key") else session_id + lock = self._session_locks.get(key) if lock is None: lock = asyncio.Lock() - self._session_locks[session_id] = lock + self._session_locks[key] = lock return lock - async def _load_state(self, session_id: str) -> MicrocompactState: - """Restore cleaned tool-result identifiers from the transcript.""" - state = self._states.get(session_id) + async def _load_state(self, session_id: str) -> MicroCompactState: + """Return process-local micro-compact state.""" + state_key = self._runtime.session_key(session_id) if hasattr(self._runtime, "session_key") else session_id + state = self._states.get(state_key) if state is not None: return state - records = await self._runtime.transcripts.read_all(session_id) - cleared_ids: set[str] = set() - result_hashes: dict[str, str] = {} - for record in records: - result_id = record.get("result_id") - if record.get("kind") != "microcompact-clear" or not isinstance(result_id, str): - continue - cleared_ids.add(result_id) - original_sha256 = record.get("original_sha256") - if isinstance(original_sha256, str): - result_hashes[result_id] = original_sha256 - state = MicrocompactState( - cleared_ids=cleared_ids, - result_hashes=result_hashes, + state = MicroCompactState( + cleared_ids=set(), + result_hashes={}, ) - self._states[session_id] = state + self._states[state_key] = state return state - def _collect_candidates(self, request: "LlmRequest") -> list[MicrocompactCandidate]: + def _collect_candidates(self, request: LlmRequest) -> list[MicroCompactCandidate]: """Collect eligible tool results in request order.""" - allowed_tools = set(self._runtime.config.microcompact_tool_names) - candidates: list[MicrocompactCandidate] = [] + allowed_tools = set(self._config.tool_names) + candidates: list[MicroCompactCandidate] = [] serialized_by_result_id: dict[str, str] = {} for content in request.contents: for part in content.parts or []: @@ -135,7 +123,7 @@ def _collect_candidates(self, request: "LlmRequest") -> list[MicrocompactCandida raise ValueError(f"Tool result id {result_id!r} is reused with different content") serialized_by_result_id[result_id] = serialized candidates.append( - MicrocompactCandidate( + MicroCompactCandidate( result_id=result_id, tool_name=function_response.name, original_size=len(serialized), @@ -148,42 +136,48 @@ def _cleared_response(self) -> dict[str, str]: """Return the minimal placeholder shared by cleanups.""" return {"output": MICROCOMPACT_CLEARED_MESSAGE} - async def _persist_clear( + async def apply( self, - session_id: str, - candidate: MicrocompactCandidate, - trigger: str, - ) -> None: - """Persist the cleanup decision for restart recovery.""" - await self._runtime.transcripts.append_unique( - session_id, - { - "schema_version": MICROCOMPACT_SCHEMA_VERSION, - "kind": "microcompact-clear", - "clear_id": f"microcompact:{candidate.result_id}", - "result_id": candidate.result_id, - "tool_name": candidate.tool_name, - "original_chars": candidate.original_size, - "original_sha256": candidate.original_sha256, - "trigger": trigger, - "cleared_response": self._cleared_response(), - }, - unique_key="clear_id", + request: LlmRequest, + *, + ctx: InvocationContext, + last_assistant_timestamp: float | None, + now: float | None = None, + ) -> MicroCompactResult: + """Clean a request copy by age first and count second.""" + if not self._config.enabled: + return MicroCompactResult() + if self._runtime.scope: + return await self._apply_scoped( + request, + session_id=ctx.session_id, + last_assistant_timestamp=last_assistant_timestamp, + now=now, + ) + runtime = self._runtime.for_session(ctx.session) + processor = self._scoped_processors.get(runtime.scope) + if processor is None: + processor = copy.copy(self) + processor._runtime = runtime + processor._states = {} + processor._session_locks = {} + self._scoped_processors[runtime.scope] = processor + return await processor.apply( + request, + last_assistant_timestamp=last_assistant_timestamp, + ctx=ctx, + now=now, ) - async def apply( + async def _apply_scoped( self, - request: "LlmRequest", + request: LlmRequest, *, session_id: str, last_assistant_timestamp: float | None, - now: float | None = None, - ) -> MicrocompactResult: - """Clean a request copy by age first and count second.""" - config = self._runtime.config - if not config.enabled or not config.microcompact_enabled: - return MicrocompactResult(None, 0, 0, 0) - await self._runtime.initialize() + now: float | None, + ) -> MicroCompactResult: + """Apply one tenant-bound mechanical compaction.""" async with self._session_lock(session_id): request.contents = [content.model_copy(deep=True) for content in request.contents] state = await self._load_state(session_id) @@ -194,7 +188,7 @@ async def apply( raise ValueError(f"Tool result id {candidate.result_id!r} is reused with different content") reapplied_count = 0 - active_candidates: list[MicrocompactCandidate] = [] + active_candidates: list[MicroCompactCandidate] = [] for candidate in candidates: if candidate.result_id in state.cleared_ids: candidate.part.function_response.response = self._cleared_response() @@ -204,30 +198,29 @@ async def apply( current_time = time.time() if now is None else now gap_seconds = current_time - last_assistant_timestamp if last_assistant_timestamp is not None else None - if gap_seconds is not None and gap_seconds >= config.microcompact_gap_seconds: + if gap_seconds is not None and gap_seconds >= self._config.gap_seconds: trigger = "time" - elif len(active_candidates) > config.microcompact_trigger_count: + elif len(active_candidates) > self._config.trigger_count: trigger = "count" else: trigger = None if trigger is None: - return MicrocompactResult(None, 0, reapplied_count, 0) + return MicroCompactResult(reapplied_count=reapplied_count) - clear_candidates = active_candidates[:-config.microcompact_keep_recent] + clear_candidates = active_candidates[:-self._config.keep_recent] if not clear_candidates: - return MicrocompactResult(None, 0, reapplied_count, 0) + return MicroCompactResult(reapplied_count=reapplied_count) cleared_size = len(serialize_tool_response(self._cleared_response())) chars_saved = 0 for candidate in clear_candidates: - await self._persist_clear(session_id, candidate, trigger) candidate.part.function_response.response = self._cleared_response() state.cleared_ids.add(candidate.result_id) state.result_hashes[candidate.result_id] = candidate.original_sha256 chars_saved += max(0, candidate.original_size - cleared_size) - return MicrocompactResult( + return MicroCompactResult( trigger=trigger, cleared_count=len(clear_candidates), reapplied_count=reapplied_count, @@ -235,43 +228,20 @@ async def apply( ) -class MicrocompactCallback: +class MicroCompactHandler(BaseCompactSummarizerHandler): """Adapt the mechanical compressor to before_model_callback.""" - advanced_memory_stage = 30 - - def __init__(self, microcompact: Microcompact) -> None: + def __init__(self) -> None: """Store the compressor executed before model requests.""" - self._microcompact = microcompact - - @property - def microcompact(self) -> Microcompact: - """Return the compressor used by this callback.""" - return self._microcompact + self._micro_compact: MicroCompact | None = None - async def __call__(self, ctx: "InvocationContext", request: "LlmRequest") -> None: + @override + async def handle(self, ctx: InvocationContext, request: LlmRequest) -> None: """Calculate the time gap and run mechanical cleanup before a request.""" - await self._microcompact.apply( - request, - session_id=ctx.session_id, - last_assistant_timestamp=find_last_assistant_timestamp(ctx), - ) - return None - - -def setup_microcompact( - agent: "LlmAgent", - memory_runtime: AdvancedMemoryRuntime, -) -> Microcompact: - """Install the mechanical callback while preserving existing order.""" - microcompact = Microcompact(memory_runtime) - callback = MicrocompactCallback(microcompact) - existing_microcompact = install_staged_callback( - agent, - callback, - callback_type=MicrocompactCallback, - component_attribute="microcompact", - memory_runtime=memory_runtime, - conflict_message="Microcompact is already configured with another runtime", - ) - return existing_microcompact or microcompact + summarizer = self.get_summarizer(ctx) + from ._auto_compact import AdvancedAutoCompactSummarizer + if not isinstance(summarizer, AdvancedAutoCompactSummarizer): + raise ValueError("Summarizer is not an AdvancedAutoCompactSummarizer") + if self._micro_compact is None: + self._micro_compact = MicroCompact(summarizer.runtime) + await self._micro_compact.apply(request, ctx=ctx, last_assistant_timestamp=find_last_assistant_timestamp(ctx)) diff --git a/trpc_agent_sdk/sessions/compact/advanced/_runtime.py b/trpc_agent_sdk/sessions/compact/advanced/_runtime.py new file mode 100644 index 000000000..1369e88e4 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/_runtime.py @@ -0,0 +1,30 @@ +# Tencent is pleased to support the open source ecosystem. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# Licensed under Apache-2.0. +"""Runtime coordination for Session Compact.""" + +from __future__ import annotations + +from dataclasses import dataclass, field + +from ..._session import Session +from ._config import AdvancedAutoCompactSummarizerConfig +from ._coordination import SessionOperationCoordinator + + +@dataclass +class AdvancedAutoCompactSummarizerRuntime: + """Hold advanced auto compact summarizer configuration and per-session coordination only.""" + + config: AdvancedAutoCompactSummarizerConfig = field(default_factory=AdvancedAutoCompactSummarizerConfig) + coordination: SessionOperationCoordinator = field(default_factory=SessionOperationCoordinator) + scope: str = field(default="") + + def for_session(self, session: Session) -> AdvancedAutoCompactSummarizerRuntime: + return AdvancedAutoCompactSummarizerRuntime(config=self.config, + coordination=self.coordination, + scope=f"{session.app_name}\0{session.user_id}") + + def session_key(self, session_id: str) -> str: + return f"{self.scope}\0{session_id}" diff --git a/trpc_agent_sdk/advanced_memory/_token_budget.py b/trpc_agent_sdk/sessions/compact/advanced/_token_budget.py similarity index 68% rename from trpc_agent_sdk/advanced_memory/_token_budget.py rename to trpc_agent_sdk/sessions/compact/advanced/_token_budget.py index 544fefd2a..22a77cd36 100644 --- a/trpc_agent_sdk/advanced_memory/_token_budget.py +++ b/trpc_agent_sdk/sessions/compact/advanced/_token_budget.py @@ -11,13 +11,17 @@ import math import re from dataclasses import dataclass +from dataclasses import field from typing import Any -from typing import Protocol -from typing import TYPE_CHECKING +from typing import Optional +from typing_extensions import override -if TYPE_CHECKING: - from trpc_agent_sdk.context import InvocationContext - from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.types import Content + +from ._config import TokenContextTrackerConfig +from ._base import BaseTokenEstimator _CJK_CHARACTER = re.compile(r"[\u3400-\u4dbf\u4e00-\u9fff\uf900-\ufaff]") @@ -28,7 +32,7 @@ class ContextTokenEstimate: tokens: int source: str - usage_event_id: str | None = None + usage_event_id: Optional[str] = field(default=None) @dataclass(frozen=True) @@ -36,11 +40,11 @@ class ContextBudget: """Describe the model window and the request's position within it.""" estimate: ContextTokenEstimate - context_window_tokens: int | None - effective_window_tokens: int | None - warning_threshold_tokens: int | None - autocompact_threshold_tokens: int | None - blocking_threshold_tokens: int | None + context_window_tokens: Optional[int] = field(default=None) + effective_window_tokens: Optional[int] = field(default=None) + warning_threshold_tokens: Optional[int] = field(default=None) + auto_compact_threshold_tokens: Optional[int] = field(default=None) + blocking_threshold_tokens: Optional[int] = field(default=None) @property def token_mode_enabled(self) -> bool: @@ -48,23 +52,10 @@ def token_mode_enabled(self) -> bool: return self.effective_window_tokens is not None -class TokenEstimator(Protocol): - """Define the replaceable token estimator interface.""" - - def estimate_payload_tokens(self, payload: Any) -> int: - """Estimate tokens for any JSON-compatible payload.""" - - -class ModelContextWindowResolver(Protocol): - """Define the model-identifier context-window resolver interface.""" - - def resolve_context_window_tokens(self, model: Any) -> int | None: - """Return the model context window, or None when unknown.""" - - -class HeuristicTokenEstimator: +class HeuristicTokenEstimator(BaseTokenEstimator): """Estimate JSON request tokens with a mixed-language heuristic.""" + @override def estimate_payload_tokens(self, payload: Any) -> int: """Estimate CJK at one token per character and other text at four characters per token.""" rendered = json.dumps( @@ -111,17 +102,17 @@ def _usage_context_tokens(usage: Any) -> int | None: return sum(value for value in values if isinstance(value, int)) -def _request_static_fingerprint(request: "LlmRequest") -> str: +def _request_static_fingerprint(request: LlmRequest) -> str: """Extract fingerprints for model, instructions, and tool configuration.""" - config = getattr(request, "config", None) - if hasattr(config, "model_dump"): + config = request.config + if config is not None: config_payload = config.model_dump( mode="python", by_alias=True, exclude_none=True, ) else: - config_payload = config + config_payload = {} return json.dumps( { "model": request.model, @@ -134,45 +125,46 @@ def _request_static_fingerprint(request: "LlmRequest") -> str: ) -class TokenContextTracker: +class TokenContextTracker(BaseTokenEstimator): """Estimate request context tokens from usage and new content.""" - def __init__(self, config: Any) -> None: + def __init__(self, config: TokenContextTrackerConfig): """Store configuration and choose the default or injected estimator.""" self._config = config - estimator = getattr(config, "token_estimator", None) - self._estimator: TokenEstimator = estimator or HeuristicTokenEstimator() + self._estimator = config.estimator or HeuristicTokenEstimator() - def _resolve_window_tokens(self, ctx: "InvocationContext | None") -> int | None: + def _resolve_window_tokens(self, ctx: InvocationContext) -> Optional[int]: """Resolve the model context window from config or an application resolver.""" - explicit = getattr(self._config, "model_context_window_tokens", None) - if isinstance(explicit, int) and explicit > 0: + if not self._config.enabled: + return None + explicit = self._config.model_context_window_tokens + if explicit is not None and explicit > 0: return explicit - resolver = getattr(self._config, "context_window_resolver", None) + resolver = self._config.context_window_resolver if resolver is None or ctx is None: return None - model = getattr(getattr(ctx, "agent", None), "model", None) + model = ctx.agent.model if ctx.agent else None resolved = resolver.resolve_context_window_tokens(model) - return resolved if isinstance(resolved, int) and resolved > 0 else None + return resolved if resolved is not None and resolved > 0 else None - def _estimate_request(self, request: "LlmRequest") -> int: + def _estimate_request(self, request: LlmRequest) -> int: """Estimate tokens for a complete LlmRequest.""" payload = request.model_dump(mode="python", by_alias=True, exclude_none=True) return self._estimator.estimate_payload_tokens(payload) - def _estimate_new_contents(self, contents: list[Any]) -> int: + def _estimate_new_contents(self, contents: list[Content]) -> int: """Estimate content tokens added after a usage baseline.""" payload = [content.model_dump(mode="python", by_alias=True, exclude_none=True) for content in contents] return self._estimator.estimate_payload_tokens(payload) if payload else 0 def _latest_usage_baseline( self, - request: "LlmRequest", - ctx: "InvocationContext | None", + request: LlmRequest, + ctx: InvocationContext, ) -> ContextTokenEstimate | None: """Match the latest usage event and estimate subsequent context.""" - session = getattr(ctx, "session", None) if ctx is not None else None - events = getattr(session, "events", None) + session = ctx.session + events = session.events if not isinstance(events, list): return None fingerprints = [_content_fingerprint(content) for content in request.contents] @@ -200,10 +192,10 @@ def _latest_usage_baseline( ) return None - def estimate( + def _estimate( self, - request: "LlmRequest", - ctx: "InvocationContext | None" = None, + request: LlmRequest, + ctx: InvocationContext, ) -> ContextTokenEstimate: """Prefer recent model usage, falling back to a full request estimate.""" baseline = self._latest_usage_baseline(request, ctx) @@ -214,17 +206,30 @@ def estimate( source="estimated", ) + def estimate( + self, + request: LlmRequest, + ctx: InvocationContext, + ) -> ContextTokenEstimate: + """Return the best available token estimate for a request.""" + return self._estimate(request, ctx) + + @override def estimate_payload_tokens(self, payload: Any) -> int: """Reuse the same estimator for non-request inputs such as session memory.""" return self._estimator.estimate_payload_tokens(payload) - def token_mode_enabled(self, ctx: "InvocationContext | None" = None) -> bool: + def estimate_request_tokens(self, request: LlmRequest) -> int: + """Estimate a complete request without applying a usage baseline.""" + return self._estimate_request(request) + + def token_mode_enabled(self, ctx: InvocationContext) -> bool: """Return whether the configuration resolves a model context window.""" return self._resolve_window_tokens(ctx) is not None def effective_context_window_tokens( self, - ctx: "InvocationContext | None" = None, + ctx: InvocationContext, ) -> int | None: """Return the input window after max output, or None when unknown.""" context_window = self._resolve_window_tokens(ctx) @@ -233,39 +238,40 @@ def effective_context_window_tokens( effective = context_window - getattr(self._config, "max_output_tokens", 0) return effective if effective > 0 else None + @classmethod def record_request_context( - self, - request: "LlmRequest", - ctx: "InvocationContext | None", + cls, + request: LlmRequest, + ctx: InvocationContext, ) -> None: """Stage the final request fingerprint for persistence on the response Event.""" - session = getattr(ctx, "session", None) if ctx is not None else None - state = getattr(session, "state", None) + session = ctx.session + state = session.state if isinstance(state, dict): state["advanced_memory_pending_request_context_fingerprint"] = _request_static_fingerprint(request) def budget( self, - request: "LlmRequest", - ctx: "InvocationContext | None" = None, + request: LlmRequest, + ctx: InvocationContext, ) -> ContextBudget: """Calculate the effective window, thresholds, and token estimate.""" - estimate = self.estimate(request, ctx) + estimate = self._estimate(request, ctx) context_window = self._resolve_window_tokens(ctx) if context_window is None: - return ContextBudget(estimate, None, None, None, None, None) - max_output_tokens = getattr(self._config, "max_output_tokens", 0) + return ContextBudget(estimate=estimate) + max_output_tokens = self._config.max_output_tokens effective = context_window - max_output_tokens if effective <= 0: - return ContextBudget(estimate, context_window, None, None, None, None) - warning = math.floor(effective * getattr(self._config, "token_warning_ratio", 0.85)) - autocompact = math.floor(effective * getattr(self._config, "token_autocompact_ratio", 0.90)) - blocking = math.floor(effective * getattr(self._config, "token_blocking_ratio", 0.95)) + return ContextBudget(estimate=estimate, context_window_tokens=context_window) + warning = math.floor(effective * self._config.warning_ratio) + auto_compact = math.floor(effective * self._config.auto_compact_ratio) + blocking = math.floor(effective * self._config.blocking_ratio) return ContextBudget( - estimate, - context_window, - effective, - warning, - autocompact, - blocking, + estimate=estimate, + context_window_tokens=context_window, + effective_window_tokens=effective, + warning_threshold_tokens=warning, + auto_compact_threshold_tokens=auto_compact, + blocking_threshold_tokens=blocking, ) diff --git a/trpc_agent_sdk/advanced_memory/_tool_result_budget.py b/trpc_agent_sdk/sessions/compact/advanced/_tool_result_budget.py similarity index 58% rename from trpc_agent_sdk/advanced_memory/_tool_result_budget.py rename to trpc_agent_sdk/sessions/compact/advanced/_tool_result_budget.py index 73a710aa2..0f9fd51ec 100644 --- a/trpc_agent_sdk/advanced_memory/_tool_result_budget.py +++ b/trpc_agent_sdk/sessions/compact/advanced/_tool_result_budget.py @@ -8,20 +8,19 @@ from __future__ import annotations import asyncio +import copy import hashlib import json -from dataclasses import dataclass -from pathlib import Path -from typing import Any -from typing import TYPE_CHECKING +from dataclasses import dataclass, field +from typing import Any, Optional +from typing_extensions import override -from ._callbacks import install_staged_callback -from ._runtime import AdvancedMemoryRuntime +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.models import LlmRequest -if TYPE_CHECKING: - from trpc_agent_sdk.agents import LlmAgent - from trpc_agent_sdk.context import InvocationContext - from trpc_agent_sdk.models import LlmRequest +from ..._session import Session +from ._runtime import AdvancedAutoCompactSummarizerRuntime +from ._base import BaseCompactSummarizerHandler TOOL_RESULT_REPLACEMENT_SCHEMA_VERSION = 1 @@ -44,14 +43,14 @@ class ToolResultCandidate: serialized_result: str original_size: int part: Any + event_id: Optional[str] = field(default=None) @dataclass(frozen=True) class ToolResultReplacement: - """Describe a tool result about to be persisted and replaced by a preview.""" + """Describe a tool result replaced by an Event reference and preview.""" candidate: ToolResultCandidate - persisted_path: Path replacement_response: dict[str, Any] replacement_size: int @@ -60,13 +59,13 @@ class ToolResultReplacement: class ToolResultBudgetResult: """Summarize replacements and character savings from budget processing.""" - replaced_count: int - original_chars: int - replacement_chars: int + replaced_count: int = field(default=0) + original_chars: int = field(default=0) + replacement_chars: int = field(default=0) def serialize_tool_response(response: Any) -> str: - """Serialize a tool result as stable JSON for counting and storage.""" + """Serialize a tool result as stable JSON for character counting.""" return json.dumps( response, ensure_ascii=False, @@ -91,14 +90,11 @@ def tool_result_sha256(serialized_result: str) -> str: def is_budget_replacement_response(response: Any) -> bool: - """Return whether a response is already an immutable storage pointer.""" + """Return whether a response is already a budget replacement.""" if not isinstance(response, dict): return False marker = response.get("_advanced_memory") - if isinstance(marker, dict) and marker.get("kind") == "tool-result-budget": - return True - persisted = response.get("persisted_output") - return isinstance(persisted, dict) and isinstance(persisted.get("path"), str) + return isinstance(marker, dict) and marker.get("kind") == "tool-result-budget" def _preview_text(serialized_result: str, limit: int) -> tuple[str, bool]: @@ -115,59 +111,56 @@ def _preview_text(serialized_result: str, limit: int) -> tuple[str, bool]: class ToolResultBudget: """Apply stable, recoverable tool-result budgeting to each request.""" - def __init__(self, memory_runtime: AdvancedMemoryRuntime) -> None: + def __init__(self, runtime: AdvancedAutoCompactSummarizerRuntime) -> None: """Initialize the budget processor and per-session state locks.""" - self._runtime = memory_runtime + self._runtime: AdvancedAutoCompactSummarizerRuntime = runtime + self._tool_result_budget_config = runtime.config.tool_result_budget self._states: dict[str, ToolResultBudgetState] = {} self._session_locks: dict[str, asyncio.Lock] = {} - - @property - def runtime(self) -> AdvancedMemoryRuntime: - """Return the runtime bound to this budget processor.""" - return self._runtime + self._scoped_processors: dict[object, ToolResultBudget] = {} def _session_lock(self, session_id: str) -> asyncio.Lock: """Return the unique async budget lock for a session.""" - lock = self._session_locks.get(session_id) + key = self._runtime.session_key(session_id) + lock = self._session_locks.get(key) if lock is None: lock = asyncio.Lock() - self._session_locks[session_id] = lock + self._session_locks[key] = lock return lock async def _load_state(self, session_id: str) -> ToolResultBudgetState: - """Restore frozen results and historical replacements from the transcript.""" - state = self._states.get(session_id) + """Return process-local state for the current Session.""" + state_key = self._runtime.session_key(session_id) + state = self._states.get(state_key) if state is not None: return state - records = await self._runtime.transcripts.read_all(session_id) - seen_ids: set[str] = set() - replacements: dict[str, dict[str, Any]] = {} - result_hashes: dict[str, str] = {} - for record in records: - if record.get("kind") not in { - "content-replacement", - "content-replacement-decision", - }: - continue - result_id = record.get("result_id") - replacement = record.get("replacement_response") - original_sha256 = record.get("original_sha256") - if isinstance(result_id, str): - seen_ids.add(result_id) - if isinstance(original_sha256, str): - result_hashes[result_id] = original_sha256 - if record.get("kind") == "content-replacement" and isinstance(replacement, dict): - replacements[result_id] = replacement state = ToolResultBudgetState( - seen_ids=seen_ids, - replacements=replacements, - result_hashes=result_hashes, + seen_ids=set(), + replacements={}, + result_hashes={}, ) - self._states[session_id] = state + self._states[state_key] = state return state - def _collect_candidates(self, request: "LlmRequest") -> list[list[ToolResultCandidate]]: + def _collect_candidates( + self, + request: LlmRequest, + session: Session, + ) -> list[list[ToolResultCandidate]]: """Group function responses from consecutive user contents.""" + event_ids: dict[str, str] = {} + for event in session.events: + event_id = getattr(event, "id", None) + if event_id is None: + continue + event_content = event.content + if event_content is None: + continue + for event_part in event_content.parts: + response = getattr(event_part, "function_response", None) + response_id = response.id if response is not None else None + if response_id is not None: + event_ids[response_id] = event_id candidate_groups: list[list[ToolResultCandidate]] = [] serialized_by_result_id: dict[str, str] = {} current_group: list[ToolResultCandidate] = [] @@ -191,6 +184,7 @@ def _collect_candidates(self, request: "LlmRequest") -> list[list[ToolResultCand current_group.append( ToolResultCandidate( result_id=result_id, + event_id=event_ids.get(result_id), tool_name=tool_name, serialized_result=serialized_result, original_size=len(serialized_result), @@ -202,44 +196,38 @@ def _collect_candidates(self, request: "LlmRequest") -> list[list[ToolResultCand def _build_replacement( self, - session_id: str, candidate: ToolResultCandidate, ) -> ToolResultReplacement: - """Build a deterministic storage path and model-visible preview.""" - persisted_path = self._runtime.paths.tool_result_path(session_id, candidate.result_id) + """Build an event reference and model-visible preview.""" preview, truncated = _preview_text( candidate.serialized_result, - self._runtime.config.tool_result_preview_chars, + self._tool_result_budget_config.preview_chars, ) replacement_response = { "_advanced_memory": { "kind": "tool-result-budget", "schema_version": TOOL_RESULT_REPLACEMENT_SCHEMA_VERSION, }, - "persisted_output": { - "message": "The tool result exceeded the context budget; the complete content was saved to disk.", - "path": str(persisted_path), - "original_chars": candidate.original_size, - "preview": preview, - "truncated": truncated, - }, + "message": ("The tool result exceeded the context budget; read the referenced " + "Session Event when needed."), + "session_event_id": candidate.event_id or candidate.result_id, + "original_chars": candidate.original_size, + "preview": preview, + "truncated": truncated, } return ToolResultReplacement( candidate=candidate, - persisted_path=persisted_path, replacement_response=replacement_response, replacement_size=len(serialize_tool_response(replacement_response)), ) def _select_replacements( self, - session_id: str, groups: list[list[ToolResultCandidate]], state: ToolResultBudgetState, ) -> list[ToolResultReplacement]: """Apply per-result limits, then select results under the aggregate limit.""" selected: dict[str, ToolResultReplacement] = {} - config = self._runtime.config for group in groups: fresh = [ candidate for candidate in group @@ -247,8 +235,8 @@ def _select_replacements( ] fresh_ids = {candidate.result_id for candidate in fresh} for candidate in fresh: - if candidate.original_size > config.tool_result_max_chars: - selected[candidate.result_id] = self._build_replacement(session_id, candidate) + if candidate.original_size > self._tool_result_budget_config.max_chars: + selected[candidate.result_id] = self._build_replacement(candidate) visible_size = 0 remaining_fresh: list[ToolResultCandidate] = [] @@ -265,82 +253,46 @@ def _select_replacements( remaining_fresh.append(candidate) for candidate in sorted(remaining_fresh, key=lambda item: item.original_size, reverse=True): - if visible_size <= config.tool_results_per_message_max_chars: + if visible_size <= self._tool_result_budget_config.per_message_max_chars: break - replacement = self._build_replacement(session_id, candidate) + replacement = self._build_replacement(candidate) if replacement.replacement_size >= candidate.original_size: continue selected[candidate.result_id] = replacement visible_size -= candidate.original_size - replacement.replacement_size return list(selected.values()) - async def _persist_replacement( - self, - session_id: str, - replacement: ToolResultReplacement, - ) -> None: - """Persist the full result before appending its replacement record.""" - candidate = replacement.candidate - await self._runtime.tool_results.write( - session_id, - candidate.result_id, - candidate.serialized_result, - ) - await self._runtime.transcripts.append_unique( - session_id, - { - "schema_version": TOOL_RESULT_REPLACEMENT_SCHEMA_VERSION, - "kind": "content-replacement", - "decision_id": f"budget:{candidate.result_id}", - "result_id": candidate.result_id, - "tool_name": candidate.tool_name, - "original_chars": candidate.original_size, - "original_sha256": tool_result_sha256(candidate.serialized_result), - "persisted_path": str(replacement.persisted_path), - "replacement_response": replacement.replacement_response, - }, - unique_key="decision_id", - ) - - async def _persist_seen_decision( - self, - session_id: str, - candidate: ToolResultCandidate, - ) -> None: - """Record a no-replacement decision to preserve sent prompt prefixes.""" - await self._runtime.transcripts.append_unique( - session_id, - { - "schema_version": TOOL_RESULT_REPLACEMENT_SCHEMA_VERSION, - "kind": "content-replacement-decision", - "decision_id": f"budget:{candidate.result_id}", - "result_id": candidate.result_id, - "tool_name": candidate.tool_name, - "replaced": False, - "original_sha256": tool_result_sha256(candidate.serialized_result), - }, - unique_key="decision_id", - ) - - async def apply(self, request: "LlmRequest", *, session_id: str) -> ToolResultBudgetResult: + async def apply(self, request: LlmRequest, ctx: InvocationContext) -> ToolResultBudgetResult: """Process a model request without mutating session Events.""" - if not self._runtime.config.enabled: - return ToolResultBudgetResult(0, 0, 0) - await self._runtime.initialize() - async with self._session_lock(session_id): + if not self._tool_result_budget_config.enabled: + return ToolResultBudgetResult() + if self._runtime.scope: + return await self._apply_scoped(request, ctx.session) + runtime = self._runtime.for_session(ctx.session) + processor = self._scoped_processors.get(runtime.scope) + if processor is None: + processor = copy.copy(self) + processor._runtime = runtime + processor._states = {} + processor._session_locks = {} + self._scoped_processors[runtime.scope] = processor + return await processor.apply(request, ctx=ctx) + + async def _apply_scoped(self, request: LlmRequest, session: Session) -> ToolResultBudgetResult: + """Apply budgeting while ``_runtime`` is bound to the current tenant.""" + async with self._session_lock(session.id): request.contents = [content.model_copy(deep=True) for content in request.contents] - state = await self._load_state(session_id) - groups = self._collect_candidates(request) + state = await self._load_state(session.id) + groups = self._collect_candidates(request, session) for group in groups: for candidate in group: known_hash = state.result_hashes.get(candidate.result_id) current_hash = tool_result_sha256(candidate.serialized_result) if known_hash is not None and known_hash != current_hash: raise ValueError(f"Tool result id {candidate.result_id!r} is reused with different content") - selected = self._select_replacements(session_id, groups, state) + selected = self._select_replacements(groups, state) for replacement in selected: - await self._persist_replacement(session_id, replacement) state.replacements[replacement.candidate.result_id] = replacement.replacement_response state.result_hashes[replacement.candidate.result_id] = tool_result_sha256( replacement.candidate.serialized_result) @@ -349,7 +301,6 @@ async def apply(self, request: "LlmRequest", *, session_id: str) -> ToolResultBu for group in groups: for candidate in group: if candidate.result_id not in state.seen_ids and candidate.result_id not in selected_ids: - await self._persist_seen_decision(session_id, candidate) state.seen_ids.add(candidate.result_id) state.result_hashes[candidate.result_id] = tool_result_sha256(candidate.serialized_result) @@ -374,39 +325,21 @@ async def apply(self, request: "LlmRequest", *, session_id: str) -> ToolResultBu ) -class ToolResultBudgetCallback: +class ToolResultBudgetHandler(BaseCompactSummarizerHandler): """Adapt the tool-result budget processor to before_model_callback.""" - advanced_memory_stage = 10 - - def __init__(self, budget: ToolResultBudget) -> None: + def __init__(self) -> None: """Store the budget processor run before model requests.""" - self._budget = budget - - @property - def budget(self) -> ToolResultBudget: - """Return the budget processor used by this callback.""" - return self._budget + self._budget: ToolResultBudget | None = None - async def __call__(self, ctx: "InvocationContext", request: "LlmRequest") -> None: + @override + async def handle(self, ctx: InvocationContext, request: LlmRequest) -> None: """Apply tool-result budgeting without truncating model calls.""" - await self._budget.apply(request, session_id=ctx.session_id) + summarizer = self.get_summarizer(ctx) + from ._auto_compact import AdvancedAutoCompactSummarizer + if not isinstance(summarizer, AdvancedAutoCompactSummarizer): + raise ValueError("Summarizer is not an AdvancedAutoCompactSummarizer") + if self._budget is None: + self._budget = ToolResultBudget(summarizer.runtime) + await self._budget.apply(request, ctx=ctx) return None - - -def setup_tool_result_budget( - agent: "LlmAgent", - memory_runtime: AdvancedMemoryRuntime, -) -> ToolResultBudget: - """Install the budget callback while preserving existing callbacks.""" - budget = ToolResultBudget(memory_runtime) - callback = ToolResultBudgetCallback(budget) - existing_budget = install_staged_callback( - agent, - callback, - callback_type=ToolResultBudgetCallback, - component_attribute="budget", - memory_runtime=memory_runtime, - conflict_message="Tool result budget is already configured with another runtime", - ) - return existing_budget or budget diff --git a/trpc_agent_sdk/sessions/compact/advanced/_utils.py b/trpc_agent_sdk/sessions/compact/advanced/_utils.py new file mode 100644 index 000000000..40dea0d9a --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/_utils.py @@ -0,0 +1,74 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Advanced utils for compact session manager.""" + +import hashlib +import json +from contextlib import contextmanager +from typing import Any +from typing import Iterator + +from trpc_agent_sdk.context import AgentContext +from trpc_agent_sdk.types import Content + +INTERNAL_COMPACTION_METADATA_KEY = "_trpc_agent_internal_compaction" +_MISSING = object() + + +@contextmanager +def internal_compaction_call(agent_context: AgentContext | None) -> Iterator[None]: + """Prevent the compact filter from recursively handling its own model call.""" + if agent_context is None: + yield + return + previous = agent_context.get_metadata(INTERNAL_COMPACTION_METADATA_KEY, _MISSING) + agent_context.with_metadata(INTERNAL_COMPACTION_METADATA_KEY, True) + try: + yield + finally: + if previous is _MISSING: + agent_context.metadata.pop(INTERNAL_COMPACTION_METADATA_KEY, None) + else: + agent_context.with_metadata(INTERNAL_COMPACTION_METADATA_KEY, previous) + + +def content_signature(content: Content) -> str: + """Generate a stable signature that preserves message identity.""" + parts: list[dict[str, Any]] = [] + for part in content.parts or []: + if part.text is not None: + parts.append({ + "type": "text", + "sha256": hashlib.sha256(part.text.encode("utf-8")).hexdigest(), + }) + elif part.function_call is not None: + parts.append({ + "type": "function_call", + "id": getattr(part.function_call, "id", None), + "name": part.function_call.name, + }) + elif part.function_response is not None: + parts.append({ + "type": "function_response", + "id": getattr(part.function_response, "id", None), + "name": part.function_response.name, + }) + elif part.executable_code is not None: + parts.append({"type": "executable_code"}) + elif part.code_execution_result is not None: + parts.append({"type": "code_execution_result"}) + else: + parts.append({"type": "other"}) + serialized = json.dumps( + { + "role": content.role, + "parts": parts + }, + ensure_ascii=False, + sort_keys=True, + separators=(",", ":"), + ) + return hashlib.sha256(serialized.encode("utf-8")).hexdigest() diff --git a/trpc_agent_sdk/sessions/compact/default/__init__.py b/trpc_agent_sdk/sessions/compact/default/__init__.py new file mode 100644 index 000000000..092b2dda5 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/default/__init__.py @@ -0,0 +1,34 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Default compact session manager.""" + +from ._checker import CheckSummarizerFunction +from ._checker import set_summarizer_token_threshold +from ._checker import set_summarizer_events_count_threshold +from ._checker import set_summarizer_time_interval_threshold +from ._checker import set_summarizer_important_content_threshold +from ._checker import set_summarizer_conversation_threshold +from ._checker import set_summarizer_check_functions_by_and +from ._checker import set_summarizer_check_functions_by_or +from ._summarizer import DEFAULT_SUMMARIZER_PROMPT +from ._summarizer import DefaultSessionSummary +from ._summarizer import DefaultSessionSummarizer +from ._summarizer_manager import DefaultSessionSummarizerManager + +__all__ = [ + "DEFAULT_SUMMARIZER_PROMPT", + "DefaultSessionSummarizer", + "DefaultSessionSummarizerManager", + "DefaultSessionSummary", + "CheckSummarizerFunction", + "set_summarizer_token_threshold", + "set_summarizer_events_count_threshold", + "set_summarizer_time_interval_threshold", + "set_summarizer_important_content_threshold", + "set_summarizer_conversation_threshold", + "set_summarizer_check_functions_by_and", + "set_summarizer_check_functions_by_or", +] diff --git a/trpc_agent_sdk/sessions/_summarizer_checker.py b/trpc_agent_sdk/sessions/compact/default/_checker.py similarity index 96% rename from trpc_agent_sdk/sessions/_summarizer_checker.py rename to trpc_agent_sdk/sessions/compact/default/_checker.py index 37863cc9b..50d97d38d 100644 --- a/trpc_agent_sdk/sessions/_summarizer_checker.py +++ b/trpc_agent_sdk/sessions/compact/default/_checker.py @@ -14,7 +14,7 @@ from trpc_agent_sdk.events import Event from trpc_agent_sdk.log import logger -from ._session import Session +from ..._session import Session CheckSummarizerFunction = Callable[[Session], bool] @@ -137,10 +137,7 @@ def set_summarizer_conversation_threshold(conversation_count: int = 100) -> Chec """ def _decorator(session: Session) -> bool: - if session.conversation_count > conversation_count: - session.conversation_count = 0 - return True - return False + return session.conversation_count > conversation_count return _decorator diff --git a/trpc_agent_sdk/sessions/_session_summarizer.py b/trpc_agent_sdk/sessions/compact/default/_summarizer.py similarity index 96% rename from trpc_agent_sdk/sessions/_session_summarizer.py rename to trpc_agent_sdk/sessions/compact/default/_summarizer.py index da7171107..0c6d826b0 100644 --- a/trpc_agent_sdk/sessions/_session_summarizer.py +++ b/trpc_agent_sdk/sessions/compact/default/_summarizer.py @@ -34,11 +34,13 @@ from typing import Dict from typing import List from typing import Optional +from typing_extensions import override from pydantic import BaseModel from pydantic import ConfigDict from pydantic import Field +from trpc_agent_sdk.abc import CompactSummarizerABC from trpc_agent_sdk.context import InvocationContext from trpc_agent_sdk.events import Event from trpc_agent_sdk.log import logger @@ -47,10 +49,10 @@ from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import Part -from ._session import Session -from ._summarizer_checker import CheckSummarizerFunction -from ._summarizer_checker import set_summarizer_conversation_threshold -from ._utils import find_events_for_summary +from ..._session import Session +from ..._utils import find_events_for_summary +from ._checker import CheckSummarizerFunction +from ._checker import set_summarizer_conversation_threshold DEFAULT_SUMMARIZER_PROMPT = dedent("""\ Please summarize the following conversation, focusing on: @@ -68,7 +70,7 @@ Summary:""") -class SessionSummary(BaseModel): +class DefaultSessionSummary(BaseModel): """Represents a summary of a session's conversation history. This class encapsulates the summary information including the summary text, @@ -88,6 +90,8 @@ class SessionSummary(BaseModel): """The timestamp when the summary was created.""" metadata: Dict[str, Any] = Field(default_factory=dict) """Additional metadata about the summarization.""" + model_name: str = "" + """The name of the model used for summarization.""" def get_compression_ratio(self) -> float: """Get the compression ratio achieved by summarization. @@ -111,13 +115,13 @@ def to_dict(self) -> Dict[str, Any]: "original_event_count": self.original_event_count, "compressed_event_count": self.compressed_event_count, "summary_timestamp": self.summary_timestamp, - "model_name": self.model.name, + "model_name": self.model_name, "compression_ratio": self.get_compression_ratio(), "metadata": self.metadata, } -class SessionSummarizer: +class DefaultSessionSummarizer(CompactSummarizerABC): """Summarizes conversation history to reduce memory usage. This class provides functionality to compress long conversation histories @@ -155,6 +159,7 @@ def model(self) -> LLMModel: """Get the LLM model for summarization.""" return self._model + @override async def should_summarize(self, session: Session) -> bool: """Check if the session should be summarized. @@ -291,6 +296,7 @@ def _extract_conversation_text(self, events: List[Event]) -> str: current_author = author current_branch = branch current_text = "" + continue if is_partial and current_author == author and current_text and current_branch == branch: # Merge with current accumulated text current_text += event_text @@ -355,6 +361,7 @@ def _create_summarization_prompt(self, conversation_text: str) -> str: """ return self._summarizer_prompt.format(conversation_text=conversation_text) + @override async def create_session_summary_by_events( self, events: List[Event], @@ -413,6 +420,7 @@ async def create_session_summary_by_events( logger.error("Failed to compress session %s: %s", session_id, ex, exc_info=True) return None, events + @override async def create_session_summary(self, session: Session, ctx: InvocationContext | None = None, @@ -436,6 +444,7 @@ async def create_session_summary(self, store_historical_events=store_historical_events) return summary_text + @override def get_summary_metadata(self) -> Dict[str, Any]: """Get metadata about the summarizer configuration. diff --git a/trpc_agent_sdk/sessions/_summarizer_manager.py b/trpc_agent_sdk/sessions/compact/default/_summarizer_manager.py similarity index 80% rename from trpc_agent_sdk/sessions/_summarizer_manager.py rename to trpc_agent_sdk/sessions/compact/default/_summarizer_manager.py index 2a07f58c0..65fd5af0b 100644 --- a/trpc_agent_sdk/sessions/_summarizer_manager.py +++ b/trpc_agent_sdk/sessions/compact/default/_summarizer_manager.py @@ -31,18 +31,20 @@ from typing import Any from typing import Dict from typing import Optional +from typing_extensions import override -from trpc_agent_sdk.abc import SessionServiceABC +from trpc_agent_sdk.abc import CompactSummarizerManagerABC +from trpc_agent_sdk.abc import CompactTrigger from trpc_agent_sdk.context import InvocationContext from trpc_agent_sdk.log import logger from trpc_agent_sdk.models import LLMModel -from ._session import Session -from ._session_summarizer import SessionSummarizer -from ._session_summarizer import SessionSummary +from ..._session import Session +from ._summarizer import DefaultSessionSummarizer +from ._summarizer import DefaultSessionSummary -class SummarizerSessionManager: +class DefaultSessionSummarizerManager(CompactSummarizerManagerABC): """Session service with automatic summarization capabilities. This service extends the basic session service with automatic @@ -53,8 +55,9 @@ class SummarizerSessionManager: def __init__( self, model: LLMModel, - summarizer: Optional[SessionSummarizer] = None, + summarizer: Optional[DefaultSessionSummarizer] = None, auto_summarize: bool = True, + compact_trigger: CompactTrigger = CompactTrigger.AFTER_TURN, ): """Initialize the summarizer session service. @@ -64,33 +67,13 @@ def __init__( summarizer: The session summarizer to use auto_summarize: Whether to automatically summarize sessions """ - self._base_service = None if not summarizer and model: - summarizer = SessionSummarizer(model=model) - self._summarizer: SessionSummarizer = summarizer + summarizer = DefaultSessionSummarizer(model=model) + super().__init__(summarizer, compact_trigger=compact_trigger) self._auto_summarize = auto_summarize - self._summarizer_cache: Dict[str, Dict[str, Dict[str, SessionSummary]]] = {} - - def set_session_service(self, session_service: SessionServiceABC, force: bool = False) -> None: - """Set the session service to use. - - Args: - session_service: The session service to use - force: Whether to force update even if already set - """ - if not self._base_service or force: - self._base_service = session_service - - def set_summarizer(self, summarizer: SessionSummarizer, force: bool = False) -> None: - """Set the summarizer to use. - - Args: - summarizer: The summarizer to use - force: Whether to force update even if already set - """ - if not self._summarizer or force: - self._summarizer = summarizer + self._summarizer_cache: Dict[str, Dict[str, Dict[str, DefaultSessionSummary]]] = {} + @override async def create_session_summary(self, session: Session, force: bool = False, @@ -100,6 +83,8 @@ async def create_session_summary(self, Args: session: The session to summarize """ + if self.compact_trigger != CompactTrigger.AFTER_TURN: + return is_should_summarize = await self.should_summarize_session(session) or force # Check if session should be summarized if is_should_summarize: @@ -122,25 +107,28 @@ async def create_session_summary(self, self._summarizer_cache[app_name] = {} if user_id not in self._summarizer_cache[app_name]: self._summarizer_cache[app_name][user_id] = {} - self._summarizer_cache[app_name][user_id][session.id] = SessionSummary( + self._summarizer_cache[app_name][user_id][session.id] = DefaultSessionSummary( session_id=session.id, summary_text=summary_text, original_event_count=original_event_count, compressed_event_count=len(session.events), summary_timestamp=time.time(), + model_name=self._summarizer.model.name, ) + session.conversation_count = 0 # Update the stored session if self._base_service: await self._base_service.update_session(session) - async def get_session_summary(self, session: Session) -> Optional[SessionSummary]: + @override + async def get_session_summary(self, session: Session) -> Optional[DefaultSessionSummary]: """Get a summary of a session. Args: session: The session to summarize Returns: - SessionSummary if successful, None otherwise + DefaultSessionSummary if successful, None otherwise """ if not self._summarizer or not self._summarizer_cache: return None diff --git a/trpc_agent_sdk/storage/_sql.py b/trpc_agent_sdk/storage/_sql.py index 539c62414..0d250490a 100644 --- a/trpc_agent_sdk/storage/_sql.py +++ b/trpc_agent_sdk/storage/_sql.py @@ -293,6 +293,18 @@ async def get(self, db: SqlSession, key: SqlKey) -> Any: return await db.get(key.storage_cls, key.key) return db.get(key.storage_cls, key.key) + async def get_for_update(self, db: SqlSession, key: SqlKey) -> Any: + """Get one row while holding a database row lock until commit.""" + stmt = select(key.storage_cls) + for column, value in zip(inspect(key.storage_cls).primary_key, key.key): + stmt = stmt.where(column == value) + stmt = stmt.with_for_update() + if isinstance(db, AsyncSession): + result = await db.execute(stmt) + else: + result = db.execute(stmt) + return result.scalars().first() + @override async def query(self, db: SqlSession, key: SqlKey, conditions: SqlCondition) -> Any: """Query the data""" diff --git a/trpc_agent_sdk/tools/__init__.py b/trpc_agent_sdk/tools/__init__.py index 20a937149..071b5365a 100644 --- a/trpc_agent_sdk/tools/__init__.py +++ b/trpc_agent_sdk/tools/__init__.py @@ -14,9 +14,9 @@ # Lazy re-export — see ``_LAZY_REEXPORTS`` below. from trpc_agent_sdk.agents.sub_agent import DynamicSubAgentTool as DynamicSubAgentTool # noqa: F401 from trpc_agent_sdk.agents.sub_agent import SpawnSubAgentTool as SpawnSubAgentTool # noqa: F401 - from trpc_agent_sdk.tools._advanced_memory_tool import AdvancedMemoryTools as AdvancedMemoryTools # noqa: F401 - from trpc_agent_sdk.tools._advanced_memory_tool import ( # noqa: F401 - create_advanced_memory_tools as create_advanced_memory_tools, ) + # from trpc_agent_sdk.tools._advanced_memory_tool import AdvancedMemoryTools as AdvancedMemoryTools # noqa: F401 + # from trpc_agent_sdk.tools._advanced_memory_tool import ( # noqa: F401 + # create_advanced_memory_tools as create_advanced_memory_tools, ) from ._agent_tool import AGENT_TOOL_APP_NAME_SUFFIX from ._agent_tool import AgentTool @@ -202,14 +202,14 @@ # the tools package free of optional file/web tool dependencies) but exposed # here for discoverability. Not in ``__all__`` so ``import *`` stays lazy. _LAZY_REEXPORTS = { - "AdvancedMemoryTools": ( - "trpc_agent_sdk.tools._advanced_memory_tool", - "AdvancedMemoryTools", - ), - "create_advanced_memory_tools": ( - "trpc_agent_sdk.tools._advanced_memory_tool", - "create_advanced_memory_tools", - ), + # "AdvancedMemoryTools": ( + # "trpc_agent_sdk.tools._advanced_memory_tool", + # "AdvancedMemoryTools", + # ), + # "create_advanced_memory_tools": ( + # "trpc_agent_sdk.tools._advanced_memory_tool", + # "create_advanced_memory_tools", + # ), "DynamicSubAgentTool": ("trpc_agent_sdk.agents.sub_agent", "DynamicSubAgentTool"), "SpawnSubAgentTool": ("trpc_agent_sdk.agents.sub_agent", "SpawnSubAgentTool"), } diff --git a/trpc_agent_sdk/tools/_advanced_memory_tool.py b/trpc_agent_sdk/tools/_advanced_memory_tool.py index f0825a774..933fd77d2 100644 --- a/trpc_agent_sdk/tools/_advanced_memory_tool.py +++ b/trpc_agent_sdk/tools/_advanced_memory_tool.py @@ -8,15 +8,15 @@ from __future__ import annotations import asyncio -import re from typing import Any -from trpc_agent_sdk.advanced_memory._formats import MemoryDocument -from trpc_agent_sdk.advanced_memory._formats import MemoryIndexEntry -from trpc_agent_sdk.advanced_memory._formats import MemoryType -from trpc_agent_sdk.advanced_memory._formats import memory_freshness -from trpc_agent_sdk.advanced_memory._formats import parse_memory_updated_at -from trpc_agent_sdk.advanced_memory._runtime import AdvancedMemoryRuntime +from trpc_agent_sdk.memory.advanced_memory._formats import MemoryDocument +from trpc_agent_sdk.memory.advanced_memory._formats import MemoryIndexEntry +from trpc_agent_sdk.memory.advanced_memory._formats import MemoryType +from trpc_agent_sdk.memory.advanced_memory._formats import memory_freshness +from trpc_agent_sdk.memory.advanced_memory._formats import parse_memory_updated_at +from trpc_agent_sdk.memory.advanced_memory._storage import parse_memory_index +from trpc_agent_sdk.memory.advanced_memory._runtime import AdvancedMemoryRuntime from ._function_tool import FunctionTool @@ -25,27 +25,25 @@ "read_memory", "list_memory_index", }) -_INDEX_PATTERN = re.compile(r"^- \[(?P.+?)\]((?P.+?)):(?P.+)$") + + +def _memory_index_reference(runtime: Any) -> str: + """Return a storage-accurate reference to the tenant memory index.""" + return runtime.paths.storage_reference("memory_index") def _parse_index(index: str) -> list[MemoryIndexEntry]: """Parse standard Advanced Memory index entries from MEMORY.md.""" - entries: list[MemoryIndexEntry] = [] - for line in index.splitlines(): - match = _INDEX_PATTERN.match(line.strip()) - if match is None: - continue - entries.append(MemoryIndexEntry(**match.groupdict())) - return entries + return parse_memory_index(index) class AdvancedMemoryTools: """Wrap long-term memory storage as three official Agent-callable tools.""" def __init__(self, runtime: AdvancedMemoryRuntime) -> None: - """Store the runtime and create the index update lock.""" + """Store the runtime and create tenant-scoped index update locks.""" self._runtime = runtime - self._index_lock = asyncio.Lock() + self._index_locks: dict[str, asyncio.Lock] = {} self._tools = ( FunctionTool(self.save_memory), FunctionTool(self.read_memory), @@ -61,10 +59,22 @@ def as_tools(self) -> list[FunctionTool]: """Return tools that can be appended directly to LlmAgent.tools.""" return list(self._tools) - def owns_tool(self, tool: Any) -> bool: - """Return whether this container created the given FunctionTool.""" - function = getattr(tool, "func", None) - return getattr(function, "__self__", None) is self + def _runtime_for_context(self, tool_context: Any | None) -> Any: + """Resolve storage from the authenticated session, never tool arguments.""" + if tool_context is None: + return self._runtime + session = getattr(tool_context, "session", None) + return self._runtime.for_session(session) + + def _index_lock(self, runtime: Any) -> asyncio.Lock: + """Return a lock for one long-term-memory tenant index.""" + scope = getattr(runtime, "scope", None) + key = scope.storage_key if scope is not None else str(runtime.paths.root_dir) + lock = self._index_locks.get(key) + if lock is None: + lock = asyncio.Lock() + self._index_locks[key] = lock + return lock async def save_memory( self, @@ -74,6 +84,7 @@ async def save_memory( memory_type: str, summary: str, content: str, + tool_context: Any | None = None, ) -> dict: """Save or overwrite a long-term memory file and update MEMORY.md.""" try: @@ -87,12 +98,13 @@ async def save_memory( memory_type=resolved_type, content=content, ) - async with self._index_lock: - path = await self._runtime.long_term_memory.write_topic( + runtime = self._runtime_for_context(tool_context) + async with self._index_lock(runtime): + path = await runtime.long_term_memory.write_topic( filename, document, ) - entries = _parse_index(await self._runtime.long_term_memory.read_index()) + entries = _parse_index(await runtime.long_term_memory.read_index()) new_entry = MemoryIndexEntry( name=name, filename=path.name, @@ -100,19 +112,19 @@ async def save_memory( ) entries = [entry for entry in entries if entry.filename != new_entry.filename] entries.insert(0, new_entry) - await self._runtime.long_term_memory.write_index(entries) - updated_at = parse_memory_updated_at(await self._runtime.long_term_memory.read_topic(filename) or "") + await runtime.long_term_memory.write_index(entries) + updated_at = parse_memory_updated_at(await runtime.long_term_memory.read_topic(filename) or "") return { "saved": True, "filename": path.name, - "path": str(path), + "path": runtime.paths.storage_reference("memory_topic", topic_name=path.name), "memory_type": resolved_type.value, "updated_at": updated_at.isoformat() if updated_at is not None else None, } - async def read_memory(self, filename: str) -> dict: + async def read_memory(self, filename: str, tool_context: Any | None = None) -> dict: """Read a complete long-term memory by its filename in MEMORY.md.""" - content = await self._runtime.long_term_memory.read_topic(filename) + content = await self._runtime_for_context(tool_context).long_term_memory.read_topic(filename) if content is None: return {"found": False, "filename": filename} updated_at = parse_memory_updated_at(content) @@ -133,14 +145,15 @@ async def read_memory(self, filename: str) -> dict: "update this memory if it is outdated or incorrect."), } - async def list_memory_index(self) -> dict: - """Return the current long-term memory index and its disk path.""" + async def list_memory_index(self, tool_context: Any | None = None) -> dict: + """Return the current long-term memory index and its storage reference.""" + runtime = self._runtime_for_context(tool_context) return { - "index_path": str(self._runtime.paths.memory_index_path), - "index": await self._runtime.long_term_memory.read_index(), + "index_path": _memory_index_reference(runtime), + "index": await runtime.long_term_memory.read_index(), } -def create_advanced_memory_tools(runtime: AdvancedMemoryRuntime, ) -> list[FunctionTool]: +def create_advanced_memory_tools(runtime: AdvancedMemoryRuntime) -> list[FunctionTool]: """Create the official Advanced Memory tools bound to the given runtime.""" return AdvancedMemoryTools(runtime).as_tools()