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/run_agent.py b/examples/memory_service_with_advanced_memory/run_agent.py index 097a3271d..9d34a0dd1 100644 --- a/examples/memory_service_with_advanced_memory/run_agent.py +++ b/examples/memory_service_with_advanced_memory/run_agent.py @@ -11,26 +11,47 @@ 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 AdvancedCompactConfig +from trpc_agent_sdk.sessions.compact import AdvancedSessionCompactManager 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 = AdvancedSessionCompactManager( + config=AdvancedCompactConfig(), ) + return InMemorySessionService( + session_config=SessionServiceConfig( + ttl=SessionServiceConfig.create_ttl_config( + enable=True, + ttl_seconds=60, + cleanup_interval_seconds=5, + ), + store_historical_events=True, + ), + session_compact_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 +78,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 +118,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..0a21b66ba --- /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_service_with_advanced_memory_redis/.env b/examples/session_service_with_advanced_memory_redis/.env new file mode 100644 index 000000000..4858f369a --- /dev/null +++ b/examples/session_service_with_advanced_memory_redis/.env @@ -0,0 +1,10 @@ +TRPC_AGENT_API_KEY= +TRPC_AGENT_BASE_URL= +TRPC_AGENT_MODEL_NAME= + +REDIS_USER= +REDIS_PASSWORD= +REDIS_HOST=127.0.0.1 +REDIS_PORT=6379 +REDIS_DB=0 +SESSION_ID=simple-demo diff --git a/examples/session_service_with_advanced_memory_redis/README.md b/examples/session_service_with_advanced_memory_redis/README.md new file mode 100644 index 000000000..9039db31a --- /dev/null +++ b/examples/session_service_with_advanced_memory_redis/README.md @@ -0,0 +1,103 @@ +# Redis SessionService + Session Compact + +本示例只演示如何在已有 `RedisSessionService` 上增加: + +- Tool Result Budget +- History Snip +- Microcompact +- AutoCompact +- AutoCompact 触发时生成的 Session Memory + +压缩能力来自 `trpc_agent_sdk.sessions.compact`,长期记忆仍属于独立的 +`trpc_agent_sdk.memory.advanced_memory`。 + +## 组装关系 + +```text +AdvancedCompactConfig + ↓ AdvancedSessionCompactManager +RedisSessionService +├── AdvancedSessionCompactManager +├── events: summary + recent Events +├── historical_events: 被压缩的原始 Events +└── state["_trpc_agent:summary"] + +``` + +核心调用: + +```python +session_config = SessionServiceConfig( + store_historical_events=True, +) +compact_config = AdvancedCompactConfig( + model_context_window_tokens=4096, + token_autocompact_ratio=0.30, +) +session_service = RedisSessionService( + db_url=redis_url, + is_async=True, + session_config=session_config, + session_compact_manager=AdvancedSessionCompactManager(config=compact_config), +) + +runner = Runner( + app_name=app_name, + agent=agent, + session_service=session_service, +) +``` + +`RedisSessionService` 会接收 `session_compact_manager`。Compact 只使用 SessionService 的 +`events`、`historical_events` 和 `state`,不创建额外的 Redis 存储。 + +## 兼容已有 Session + +旧数据不需要包含 `_trpc_agent:summary`: + +```python +summary = session.state.get("_trpc_agent:summary") +``` + +不存在时正常返回 `None`。只有上下文达到 AutoCompact 阈值后,子 Agent 才会 +根据当前可读 Events 生成第一份 Summary。 + +前三个阶段只修改发给模型的 `LlmRequest`。AutoCompact 成功后还会把同一份 +Session Memory 作为 summary Event 写到 `session.events[0]`,并把被替换的 +原始 Events 移入 `session.historical_events`。因此下一轮直接读取 +`summary + recent events`,无需重新加载已经压缩的活跃 Events。 + +## 配置与运行 + +复制并修改 `.env`: + +```dotenv +TRPC_AGENT_API_KEY=your-api-key +TRPC_AGENT_BASE_URL=your-model-base-url +TRPC_AGENT_MODEL_NAME=your-model-name +REDIS_USER= +REDIS_PASSWORD= +REDIS_HOST=127.0.0.1 +REDIS_PORT=6379 +REDIS_DB=0 +SESSION_ID=simple-demo +``` + +运行: + +```bash +cd examples/session_service_with_advanced_memory_redis +source ../../.venv/bin/activate +python run_agent.py +``` + +脚本默认使用 `simple-demo`,可通过 `SESSION_ID` 修改。重复运行可以验证 +活跃窗口、历史原始 Events 和 Session Memory 都能跨进程恢复。 + +运行结束会输出 `Active Events`、`Historical Events`、活跃窗口是否以 summary +开头,以及 Session Memory state 是否存在,方便直接确认压缩是否触发。 + +## 存储职责 + +- `RedisSessionService`:Session、活跃 Events、historical Events 和 state。 +- Compact 不创建独立的 Redis transcript、Tool Result 或 session-memory 存储。 diff --git a/examples/session_service_with_advanced_memory_redis/agent/__init__.py b/examples/session_service_with_advanced_memory_redis/agent/__init__.py new file mode 100644 index 000000000..bc6e483f9 --- /dev/null +++ b/examples/session_service_with_advanced_memory_redis/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_service_with_advanced_memory_redis/agent/agent.py b/examples/session_service_with_advanced_memory_redis/agent/agent.py new file mode 100644 index 000000000..57093a8c1 --- /dev/null +++ b/examples/session_service_with_advanced_memory_redis/agent/agent.py @@ -0,0 +1,39 @@ +# 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 for the Advanced Memory Redis session 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.tools import FunctionTool + +from .config import get_model_config +from .prompts import INSTRUCTION +from .tools import large_report + + +def _create_model() -> LLMModel: + """Create the configured model.""" + api_key, base_url, model_name = get_model_config() + return OpenAIModel( + model_name=model_name, + api_key=api_key, + base_url=base_url, + ) + + +def create_agent() -> LlmAgent: + """Create the report Agent used by the session example.""" + return LlmAgent( + name="redis_compression_demo", + description="Demonstrate Redis session context compression.", + model=_create_model(), + instruction=INSTRUCTION, + tools=[FunctionTool(large_report)], + ) + + +root_agent = create_agent() diff --git a/examples/session_service_with_advanced_memory_redis/agent/config.py b/examples/session_service_with_advanced_memory_redis/agent/config.py new file mode 100644 index 000000000..9ff843472 --- /dev/null +++ b/examples/session_service_with_advanced_memory_redis/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. +"""Model configuration for the Advanced Memory Redis session example.""" + +import os + + +def get_model_config() -> tuple[str, str, str]: + """Read required model configuration from environment variables.""" + 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/session_service_with_advanced_memory_redis/agent/prompts.py b/examples/session_service_with_advanced_memory_redis/agent/prompts.py new file mode 100644 index 000000000..8913fbef9 --- /dev/null +++ b/examples/session_service_with_advanced_memory_redis/agent/prompts.py @@ -0,0 +1,10 @@ +# 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 the Advanced Memory Redis session example.""" + +INSTRUCTION = """You are a helpful assistant. +Use large_report when the user requests a report. Keep continuity with earlier +messages and answer concisely from the available context.""" diff --git a/examples/session_service_with_advanced_memory_redis/agent/tools.py b/examples/session_service_with_advanced_memory_redis/agent/tools.py new file mode 100644 index 000000000..9a109117a --- /dev/null +++ b/examples/session_service_with_advanced_memory_redis/agent/tools.py @@ -0,0 +1,11 @@ +# 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 Memory Redis session example.""" + + +def large_report(topic: str) -> dict[str, str]: + """Return a deliberately large result for the compression demo.""" + return {"output": f"Report for {topic}\n" + ("detail " * 2_000)} diff --git a/examples/session_service_with_advanced_memory_redis/run_agent.py b/examples/session_service_with_advanced_memory_redis/run_agent.py new file mode 100644 index 000000000..77efbaca2 --- /dev/null +++ b/examples/session_service_with_advanced_memory_redis/run_agent.py @@ -0,0 +1,120 @@ +#!/usr/bin/env python3 + +# 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. +"""Run native Session compaction over the standard RedisSessionService.""" + +from __future__ import annotations + +import asyncio +import os + +from dotenv import load_dotenv + +from trpc_agent_sdk.sessions.compact import AdvancedCompactConfig +from trpc_agent_sdk.sessions.compact import AdvancedSessionCompactManager +from trpc_agent_sdk.runners import Runner +from trpc_agent_sdk.sessions import RedisSessionService +from trpc_agent_sdk.sessions import SessionServiceConfig +from trpc_agent_sdk.types import Content +from trpc_agent_sdk.types import Part + +load_dotenv() + + +def redis_url() -> str: + """Build the Redis connection URL from environment variables.""" + db_user = os.environ.get("REDIS_USER", "") + db_password = os.environ.get("REDIS_PASSWORD", "") + db_host = os.environ.get("REDIS_HOST", "127.0.0.1") + db_port = os.environ.get("REDIS_PORT", "6379") + db_name = os.environ.get("REDIS_DB", "0") + + if db_password: + if db_user: + return f"redis://{db_user}:{db_password}@{db_host}:{db_port}/{db_name}" + return f"redis://:{db_password}@{db_host}:{db_port}/{db_name}" + return f"redis://{db_host}:{db_port}/{db_name}" + + +def create_compact_config() -> AdvancedCompactConfig: + """Configure only the settings needed to demonstrate one compaction.""" + return AdvancedCompactConfig( + model_context_window_tokens=4096, + max_output_tokens=256, + token_warning_ratio=0.25, + token_autocompact_ratio=0.30, + token_blocking_ratio=0.95, + session_memory_initial_tokens=500, + session_memory_update_tokens=500, + autocompact_keep_recent_contents=2, + ) + + +async def main() -> None: + """Attach Session Compact to RedisSessionService and run the demo.""" + app_name = "session-service-advanced-memory-redis" + user_id = "demo-user" + session_id = os.getenv("SESSION_ID", "simple-demo") + from agent.agent import create_agent + + agent = create_agent() + compact_config = create_compact_config() + compact_manager = AdvancedSessionCompactManager(config=compact_config) + session_config = SessionServiceConfig(store_historical_events=True) + session_service = RedisSessionService( + db_url=redis_url(), + is_async=True, + session_config=session_config, + session_compact_manager=compact_manager, + ) + runner = Runner( + app_name=app_name, + agent=agent, + session_service=session_service, + ) + try: + for prompt in ( + "Generate a large report about Redis session persistence.", + "What are the key points and persistence options?", + "List the main operational risks and mitigations.", + "Summarize our work so far and preserve the important state.", + ): + print(f"\nUser: {prompt}") + 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 event.content and not event.partial: + for part in event.content.parts: + if part.text and not part.thought: + print(f"Assistant: {part.text}") + + stored = await session_service.get_session( + app_name=app_name, + user_id=user_id, + session_id=session_id, + ) + if stored is not None: + print(f"\nActive Events: {len(stored.events)}") + print(f"Historical Events: {len(stored.historical_events)}") + print( + "Active window starts with summary:", + bool(stored.events and stored.events[0].is_summary_event()), + ) + print( + "Session Memory state present:", + "_trpc_agent:summary" in stored.state, + ) + print("Event IDs:", [event.id for event in stored.events]) + print("Historical IDs:", [event.id for event in stored.historical_events]) + finally: + await runner.close() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/examples/session_service_with_advanced_memory_sql/.env b/examples/session_service_with_advanced_memory_sql/.env new file mode 100644 index 000000000..693f8ecb1 --- /dev/null +++ b/examples/session_service_with_advanced_memory_sql/.env @@ -0,0 +1,11 @@ +TRPC_AGENT_API_KEY= +TRPC_AGENT_BASE_URL= +TRPC_AGENT_MODEL_NAME= + +MYSQL_USER= +MYSQL_PASSWORD= +MYSQL_HOST= +MYSQL_PORT= +MYSQL_DB=trpc_agent_session +SESSION_ID=simple-demo + diff --git a/examples/session_service_with_advanced_memory_sql/README.md b/examples/session_service_with_advanced_memory_sql/README.md new file mode 100644 index 000000000..c703b700f --- /dev/null +++ b/examples/session_service_with_advanced_memory_sql/README.md @@ -0,0 +1,100 @@ +# SQL SessionService + Session Compact + +本示例只演示如何在已有 `SqlSessionService` 上增加: + +- Tool Result Budget +- History Snip +- Microcompact +- AutoCompact +- AutoCompact 触发时生成的 Session Memory + +压缩能力来自 `trpc_agent_sdk.sessions.compact`,长期记忆仍属于独立的 +`trpc_agent_sdk.memory.advanced_memory`。SQL 表结构不变,但活跃/历史 Event +会按原 Session 语义重新分区。 + +## 组装关系 + +```text +AdvancedCompactConfig + ↓ AdvancedSessionCompactManager +SqlSessionService +├── AdvancedSessionCompactManager +├── events: summary + recent Events +├── sessions.historical_events: 被压缩的原始 Events +└── sessions.state["_trpc_agent:summary"] + +``` + +核心调用: + +```python +session_config = SessionServiceConfig( + store_historical_events=True, +) +compact_config = AdvancedCompactConfig( + model_context_window_tokens=4096, + token_autocompact_ratio=0.30, +) +session_service = SqlSessionService( + db_url=sql_url, + is_async=False, + session_config=session_config, + session_compact_manager=AdvancedSessionCompactManager(config=compact_config), +) + +runner = Runner( + app_name=app_name, + agent=agent, + session_service=session_service, +) +``` + +`SqlSessionService` 会接收 `session_compact_manager`。Compact 只使用 SessionService 的 +`events`、`historical_events` 和 `state`,不创建额外的 SQL 表。 + +## 兼容已有 Session + +旧 `sessions.state` 不需要预先包含 `_trpc_agent:summary`。Key 不存在时继续使用 +原 Events;达到 AutoCompact 阈值后才生成并写入第一份结构化 Summary。 + +Session Memory 更新通过 `patch_session_state()` 完成。AutoCompact 成功后, +同一份内容会作为 summary Event 写入活跃 `events` 表;被替换的 Event 从活跃表 +移入 `sessions.historical_events`。下一轮直接读取 `summary + recent events`。 + +## 配置与运行 + +默认使用 MySQL: + +```dotenv +TRPC_AGENT_API_KEY=your-api-key +TRPC_AGENT_BASE_URL=your-model-base-url +TRPC_AGENT_MODEL_NAME=your-model-name +MYSQL_USER=root +MYSQL_PASSWORD= +MYSQL_HOST=127.0.0.1 +MYSQL_PORT=3306 +MYSQL_DB=trpc_agent_session +SESSION_ID=simple-demo +``` + +示例使用同步 `pymysql` 驱动。如果需要异步连接,可以将连接地址改为 +`mysql+aiomysql://...`,安装 `aiomysql`,并将 `is_async` 改为 `True`。 + +运行: + +```bash +cd examples/session_service_with_advanced_memory_sql +source ../../.venv/bin/activate +python run_agent.py +``` + +脚本默认使用 `simple-demo`,可通过 `SESSION_ID` 修改。重复运行可以验证 +活跃窗口、历史原始 Events、Session Memory 和完整 Tool Result 能够恢复。 + +运行结束会输出 `Active Events`、`Historical Events`、活跃窗口是否以 summary +开头,以及 Session Memory state 是否存在,方便直接确认压缩是否触发。 + +## 存储职责 + +- `SqlSessionService`:Session、活跃 Events、historical Events 和 state。 +- Compact 不创建独立的 SQL transcript、Tool Result 或 session-memory 表。 diff --git a/examples/session_service_with_advanced_memory_sql/agent/__init__.py b/examples/session_service_with_advanced_memory_sql/agent/__init__.py new file mode 100644 index 000000000..bc6e483f9 --- /dev/null +++ b/examples/session_service_with_advanced_memory_sql/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_service_with_advanced_memory_sql/agent/agent.py b/examples/session_service_with_advanced_memory_sql/agent/agent.py new file mode 100644 index 000000000..5501a1b0b --- /dev/null +++ b/examples/session_service_with_advanced_memory_sql/agent/agent.py @@ -0,0 +1,39 @@ +# 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 for the Advanced Memory SQL session 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.tools import FunctionTool + +from .config import get_model_config +from .prompts import INSTRUCTION +from .tools import large_report + + +def _create_model() -> LLMModel: + """Create the configured model.""" + api_key, base_url, model_name = get_model_config() + return OpenAIModel( + model_name=model_name, + api_key=api_key, + base_url=base_url, + ) + + +def create_agent() -> LlmAgent: + """Create the report Agent used by the session example.""" + return LlmAgent( + name="sql_compression_demo", + description="Demonstrate SQL session context compression.", + model=_create_model(), + instruction=INSTRUCTION, + tools=[FunctionTool(large_report)], + ) + + +root_agent = create_agent() diff --git a/examples/session_service_with_advanced_memory_sql/agent/config.py b/examples/session_service_with_advanced_memory_sql/agent/config.py new file mode 100644 index 000000000..91236eaf9 --- /dev/null +++ b/examples/session_service_with_advanced_memory_sql/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. +"""Model configuration for the Advanced Memory SQL session example.""" + +import os + + +def get_model_config() -> tuple[str, str, str]: + """Read required model configuration from environment variables.""" + 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/session_service_with_advanced_memory_sql/agent/prompts.py b/examples/session_service_with_advanced_memory_sql/agent/prompts.py new file mode 100644 index 000000000..d7213fa0e --- /dev/null +++ b/examples/session_service_with_advanced_memory_sql/agent/prompts.py @@ -0,0 +1,10 @@ +# 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 the Advanced Memory SQL session example.""" + +INSTRUCTION = """You are a helpful assistant. +Use large_report when the user requests a report. Keep continuity with earlier +messages and answer concisely from the available context.""" diff --git a/examples/session_service_with_advanced_memory_sql/agent/tools.py b/examples/session_service_with_advanced_memory_sql/agent/tools.py new file mode 100644 index 000000000..cf472e2b6 --- /dev/null +++ b/examples/session_service_with_advanced_memory_sql/agent/tools.py @@ -0,0 +1,11 @@ +# 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 Memory SQL session example.""" + + +def large_report(topic: str) -> dict[str, str]: + """Return a deliberately large result for the compression demo.""" + return {"output": f"Report for {topic}\n" + ("detail " * 2_000)} diff --git a/examples/session_service_with_advanced_memory_sql/run_agent.py b/examples/session_service_with_advanced_memory_sql/run_agent.py new file mode 100644 index 000000000..ffee4b3df --- /dev/null +++ b/examples/session_service_with_advanced_memory_sql/run_agent.py @@ -0,0 +1,116 @@ +#!/usr/bin/env python3 + +# 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. +"""Run native Session compaction over the standard SqlSessionService.""" + +from __future__ import annotations + +import asyncio +import os + +from dotenv import load_dotenv + +from trpc_agent_sdk.sessions.compact import AdvancedCompactConfig +from trpc_agent_sdk.sessions.compact import AdvancedSessionCompactManager +from trpc_agent_sdk.runners import Runner +from trpc_agent_sdk.sessions import SessionServiceConfig +from trpc_agent_sdk.sessions import SqlSessionService +from trpc_agent_sdk.types import Content +from trpc_agent_sdk.types import Part + +load_dotenv() + + +def sql_url() -> str: + """Build the MySQL connection URL from environment variables.""" + db_user = os.environ.get("MYSQL_USER", "root") + db_password = os.environ.get("MYSQL_PASSWORD", "") + db_host = os.environ.get("MYSQL_HOST", "127.0.0.1") + db_port = os.environ.get("MYSQL_PORT", "3306") + db_name = os.environ.get("MYSQL_DB", "trpc_agent_session") + return (f"mysql+pymysql://{db_user}:{db_password}@" + f"{db_host}:{db_port}/{db_name}?charset=utf8mb4") + + +def create_compact_config() -> AdvancedCompactConfig: + """Configure only the settings needed to demonstrate one compaction.""" + return AdvancedCompactConfig( + model_context_window_tokens=4096, + max_output_tokens=256, + token_warning_ratio=0.25, + token_autocompact_ratio=0.30, + token_blocking_ratio=0.95, + session_memory_initial_tokens=500, + session_memory_update_tokens=500, + autocompact_keep_recent_contents=2, + ) + + +async def main() -> None: + """Attach Session Compact to SqlSessionService and run the demo.""" + app_name = "session-service-advanced-memory-sql" + user_id = "demo-user" + session_id = os.getenv("SESSION_ID", "simple-demo") + from agent.agent import create_agent + + agent = create_agent() + compact_config = create_compact_config() + compact_manager = AdvancedSessionCompactManager(config=compact_config) + session_config = SessionServiceConfig(store_historical_events=True) + session_service = SqlSessionService( + db_url=sql_url(), + is_async=False, + session_config=session_config, + session_compact_manager=compact_manager, + ) + runner = Runner( + app_name=app_name, + agent=agent, + session_service=session_service, + ) + try: + for prompt in ( + "Generate a report about SQL session persistence.", + "What are the key points and persistence options?", + "List the main operational risks and mitigations.", + "Summarize our work so far and preserve the important state.", + ): + print(f"\nUser: {prompt}") + 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 event.content and not event.partial: + for part in event.content.parts: + if part.text and not part.thought: + print(f"Assistant: {part.text}") + + stored = await session_service.get_session( + app_name=app_name, + user_id=user_id, + session_id=session_id, + ) + if stored is not None: + print(f"\nActive Events: {len(stored.events)}") + print(f"Historical Events: {len(stored.historical_events)}") + print( + "Active window starts with summary:", + bool(stored.events and stored.events[0].is_summary_event()), + ) + print( + "Session Memory state present:", + "_trpc_agent:summary" in stored.state, + ) + print("Event IDs:", [event.id for event in stored.events]) + print("Historical IDs:", [event.id for event in stored.historical_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..dc2e181b2 100644 --- a/tests/advanced_memory/test_advanced_memory_tools.py +++ b/tests/advanced_memory/test_advanced_memory_tools.py @@ -3,21 +3,24 @@ 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 AdvancedMemoryRuntime +from trpc_agent_sdk.memory.advanced_memory import AdvancedMemoryServiceConfig +from trpc_agent_sdk.memory.advanced_memory import AdvancedMemoryPaths +from trpc_agent_sdk.memory.advanced_memory import AdvancedMemoryRuntime from trpc_agent_sdk.tools import AdvancedMemoryTools from trpc_agent_sdk.tools import create_advanced_memory_tools 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..4a4b319b7 100644 --- a/tests/advanced_memory/test_memory_context.py +++ b/tests/advanced_memory/test_memory_context.py @@ -7,40 +7,42 @@ 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 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.advanced_memory import AdvancedMemoryServiceConfig +from trpc_agent_sdk.memory.advanced_memory import AdvancedMemoryRuntime +from trpc_agent_sdk.memory.advanced_memory import LongTermMemoryContext +from trpc_agent_sdk.memory.advanced_memory import LongTermMemoryContextCallback +from trpc_agent_sdk.memory.advanced_memory import MemoryDocument +from trpc_agent_sdk.memory.advanced_memory import MemoryIndexEntry +from trpc_agent_sdk.memory.advanced_memory import MemoryType +from trpc_agent_sdk.abc import MemoryServiceABC +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, )) +@pytest.mark.asyncio +async def test_advanced_memory_service_implements_memory_service_contract(tmp_path: Path) -> None: + """Ensure the tool-driven service remains compatible with the base API.""" + memory_service = AdvancedMemoryService(runtime=_runtime(tmp_path)) + + assert isinstance(memory_service, MemoryServiceABC) + assert memory_service.enabled is True + await memory_service.store_session(SimpleNamespace()) + response = await memory_service.search_memory("user", "anything") + assert response.memories == [] + + await memory_service.close() + + def test_staged_callback_rejects_invalid_stage(tmp_path: Path) -> None: """Ensure a newly installed callback must declare an integer stage.""" agent = SimpleNamespace(before_model_callback=None) @@ -81,6 +83,15 @@ class StagedCallback: async def test_long_term_memory_index_is_injected_once(tmp_path: Path) -> None: """Ensure the index, paths, and on-demand read guidance are injected.""" runtime = _runtime(tmp_path) + await runtime.long_term_memory.write_topic( + "project.md", + MemoryDocument( + name="项目约定", + description="项目代码规范", + memory_type=MemoryType.PROJECT, + content="使用清晰的项目代码规范。", + ), + ) await runtime.long_term_memory.write_index( [MemoryIndexEntry( name="项目约定", @@ -104,61 +115,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..8af531bb3 100644 --- a/tests/advanced_memory/test_preload_memory.py +++ b/tests/advanced_memory/test_preload_memory.py @@ -5,11 +5,15 @@ 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 MemoryDocument -from trpc_agent_sdk.advanced_memory import MemoryPreloader -from trpc_agent_sdk.advanced_memory import MemoryType +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 MemoryDocument +from trpc_agent_sdk.memory.advanced_memory import MemoryPreloader +from trpc_agent_sdk.memory.advanced_memory import MemoryCandidate +from trpc_agent_sdk.memory.advanced_memory import ModelMemoryRelevanceSelector +from trpc_agent_sdk.memory.advanced_memory import MemoryType +from trpc_agent_sdk.types import Content +from trpc_agent_sdk.types import Part class _FakeSelector: @@ -29,17 +33,55 @@ async def select(self, query, candidates, ctx, *, limit): raise RuntimeError("selector failed") +class _FakeModel: + """Return one deterministic selector response.""" + + name = "test-model" + + async def generate_async(self, request, *, stream, ctx): + """Return the requested memory filename without using a Runner.""" + assert request.model == self.name + assert stream is False + assert ctx is None + yield SimpleNamespace( + content=Content(parts=[Part.from_text(text='{"selected_memories": ["project.md"]}')]), + error_code=None, + error_message=None, + ) + + +async def test_model_selector_uses_direct_llm_call() -> None: + """Ensure preload selection does not construct an Agent or Runner.""" + candidate = MemoryCandidate( + filename="project.md", + name="Project", + description="Project details", + memory_type="project", + updated_at=None, + ) + ctx = SimpleNamespace(agent=SimpleNamespace(model=_FakeModel())) + + selected = await ModelMemoryRelevanceSelector().select( + "project", + [candidate], + ctx, + limit=1, + ) + + assert selected == ["project.md"] + + 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 +90,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 +110,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 +126,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 +145,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 93% rename from tests/advanced_memory/test_coordination.py rename to tests/sessions/compact/test_coordination.py index 030d01bc5..2b2fb7ee1 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._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..f8de41c30 --- /dev/null +++ b/tests/sessions/compact/test_session_compact.py @@ -0,0 +1,128 @@ +"""Tests for SessionService-owned Session Compact.""" + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest + +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.sessions import SessionServiceConfig +from trpc_agent_sdk.sessions.compact import AdvancedCompactConfig +from trpc_agent_sdk.sessions.compact import AdvancedSessionCompactManager +from trpc_agent_sdk.sessions.compact import SessionCompactRuntime +from trpc_agent_sdk.sessions.compact import SessionMemoryDocument +from trpc_agent_sdk.sessions.compact import SessionMemoryExtractor +from trpc_agent_sdk.sessions.compact import ToolResultBudget +from trpc_agent_sdk.types import Content +from trpc_agent_sdk.types import FunctionResponse +from trpc_agent_sdk.types import Part + + +class _MemoryGenerator: + + async def generate(self, extraction_input, ctx) -> SessionMemoryDocument: + del ctx + return SessionMemoryDocument( + session_title="Test session", + current_state=extraction_input.last_event_id, + ) + + +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 = SessionCompactRuntime.create(AdvancedCompactConfig()) + + assert not hasattr(runtime, "transcripts") + assert not hasattr(runtime, "tool_results") + assert not hasattr(runtime, "paths") + + +@pytest.mark.asyncio +async def test_session_service_accepts_a_configured_compact_manager() -> None: + manager = AdvancedSessionCompactManager(config=AdvancedCompactConfig()) + service = InMemorySessionService( + session_config=SessionServiceConfig(store_historical_events=True), + session_compact_manager=manager, + ) + + assert service.session_compact_manager is manager + 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")])), + ) + extractor = SessionMemoryExtractor( + SessionCompactRuntime.create( + AdvancedCompactConfig( + session_memory_initial_chars=1, + session_memory_update_chars=1, + )), + _MemoryGenerator(), + session_service=service, + ) + + result = await extractor.extract_if_needed( + session, + SimpleNamespace(session=session, agent=SimpleNamespace(model="test")), + force=True, + ) + + assert result.extracted is True + assert "_trpc_agent:summary" in session.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)]) + budget = ToolResultBudget( + SessionCompactRuntime.create(AdvancedCompactConfig( + tool_result_max_chars=100, + tool_result_preview_chars=20, + ))) + + await budget.apply( + request, + session_id=session.id, + 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 85% rename from tests/advanced_memory/test_token_budget.py rename to tests/sessions/compact/test_token_budget.py index 0431f9515..9e2c918a6 100644 --- a/tests/advanced_memory/test_token_budget.py +++ b/tests/sessions/compact/test_token_budget.py @@ -4,8 +4,8 @@ from types import SimpleNamespace -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig -from trpc_agent_sdk.advanced_memory import TokenContextTracker +from trpc_agent_sdk.sessions.compact import AdvancedCompactConfig +from trpc_agent_sdk.sessions.compact import TokenContextTracker from trpc_agent_sdk.models import LlmRequest 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( + AdvancedCompactConfig( 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(AdvancedCompactConfig(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(AdvancedCompactConfig(enabled=True)).estimate(request, ctx) assert estimate.source == "estimated" assert estimate.tokens < 999_999 @@ -94,9 +93,8 @@ 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( + AdvancedCompactConfig( enabled=True, - root_dir=tmp_path, model_context_window_tokens=10_000, max_output_tokens=2_000, )) @@ -111,8 +109,7 @@ def test_budget_reserves_max_output_and_calculates_three_thresholds(tmp_path) -> 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")) + budget = TokenContextTracker(AdvancedCompactConfig(enabled=True)).budget(_request("compatibility request")) assert not budget.token_mode_enabled assert budget.estimate.source == "estimated" 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_in_memory_session_service.py b/tests/sessions/test_in_memory_session_service.py index 174daa311..51b14af33 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.patch_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..e5f154bf5 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.patch_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.patch_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_sql_session_service.py b/tests/sessions/test_sql_session_service.py index e1730ec6c..1eaa27853 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.patch_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/trpc_agent_sdk/abc/_session_service.py b/trpc_agent_sdk/abc/_session_service.py index 419fac067..d26b7f397 100644 --- a/trpc_agent_sdk/abc/_session_service.py +++ b/trpc_agent_sdk/abc/_session_service.py @@ -125,6 +125,19 @@ async def update_session(self, session: SessionABC) -> None: session: The session to update """ + async def patch_session_state( + self, + session: SessionABC, + state_delta: dict[str, Any], + ) -> None: + """Atomically merge session-scoped state without replacing Events. + + Session services that support Advanced Memory session summaries must + override this method. It is intentionally non-abstract so existing + third-party implementations remain source compatible. + """ + raise NotImplementedError(f"{type(self).__name__} does not support atomic session state patches") + @abstractmethod async def create_session_summary(self, session: SessionABC, ctx: "InvocationContext" = None) -> None: """Summarize a session.""" 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/_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..6e8e5e8fd 100644 --- a/trpc_agent_sdk/evaluation/_eval_session_service.py +++ b/trpc_agent_sdk/evaluation/_eval_session_service.py @@ -9,12 +9,16 @@ from typing import Any from typing import Optional +from typing import TYPE_CHECKING from typing_extensions import override from trpc_agent_sdk.events import Event from trpc_agent_sdk.sessions import BaseSessionService from trpc_agent_sdk.sessions import Session +if TYPE_CHECKING: + from trpc_agent_sdk.sessions.compact import BaseSessionCompactManager + class EvalSessionService(BaseSessionService): """Wraps a SessionService: on create_session, if context_messages were passed in, @@ -25,6 +29,24 @@ 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 + + @property + def session_compact_manager(self) -> Optional["BaseSessionCompactManager"]: + """Expose Session Compact installed on the storage service.""" + return self._inner.session_compact_manager + + def set_session_compact_manager( + self, + compact_manager: "BaseSessionCompactManager", + force: bool = False, + ) -> None: + """Install Session Compact on the service that owns persistence.""" + self._inner.set_session_compact_manager(compact_manager, force=force) + @override async def create_session( self, @@ -86,6 +108,17 @@ async def append_event(self, session: Session, event: Event) -> Event: async def update_session(self, session: Session) -> None: return await self._inner.update_session(session=session) + @override + async def patch_session_state( + self, + session: Session, + state_delta: dict[str, Any], + ) -> None: + return await self._inner.patch_session_state( + session=session, + state_delta=state_delta, + ) + @override async def create_session_summary(self, session: Session, ctx: Any = None) -> None: return await self._inner.create_session_summary(session=session, ctx=ctx) 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..fa16d565a --- /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._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..418d09519 100644 --- a/trpc_agent_sdk/runners.py +++ b/trpc_agent_sdk/runners.py @@ -227,12 +227,12 @@ def __init__( # 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) + compact_manager = getattr(session_service, "session_compact_manager", None) + if compact_manager is not None: + compact_manager.setup(agent) 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..9ce67dc43 100644 --- a/trpc_agent_sdk/sessions/__init__.py +++ b/trpc_agent_sdk/sessions/__init__.py @@ -53,7 +53,10 @@ "ListSessionsResponse", "State", "BaseSessionService", - "AdvancedMemorySessionService", + "BaseSessionCompactManager", + "AdvancedCompactConfig", + "AdvancedSessionCompactManager", + "AutoCompact", "HistoryRecord", "InMemorySessionService", "SessionWithTTL", @@ -92,8 +95,13 @@ def __getattr__(name: str): """Lazily expose Advanced Memory without creating an import cycle.""" - if name == "AdvancedMemorySessionService": - from ._advanced_memory_session_service import AdvancedMemorySessionService + if name in { + "AdvancedCompactConfig", + "AdvancedSessionCompactManager", + "AutoCompact", + "BaseSessionCompactManager", + }: + from . import compact - return AdvancedMemorySessionService + return getattr(compact, name) 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..6cbd8fbed 100644 --- a/trpc_agent_sdk/sessions/_base_session_service.py +++ b/trpc_agent_sdk/sessions/_base_session_service.py @@ -25,6 +25,7 @@ from __future__ import annotations from typing import Optional +from typing import TYPE_CHECKING from typing_extensions import override from trpc_agent_sdk.abc import SessionServiceABC @@ -36,6 +37,9 @@ from ._summarizer_manager import SummarizerSessionManager from ._types import SessionServiceConfig +if TYPE_CHECKING: + from .compact import BaseSessionCompactManager + class BaseSessionService(SessionServiceABC): """Abstract base class for session management services. @@ -45,14 +49,17 @@ class BaseSessionService(SessionServiceABC): def __init__(self, summarizer_manager: Optional[SummarizerSessionManager] = None, - session_config: Optional[SessionServiceConfig] = None): + session_config: Optional[SessionServiceConfig] = None, + session_compact_manager: Optional["BaseSessionCompactManager"] = None): """Initialize the base session service. Args: summarizer_manager: Optional summarizer manager for session summarization session_config: Optional session configuration + session_compact_manager: Optional pluggable Session Compact manager """ self._summarizer_manager = summarizer_manager + self._session_compact_manager: Optional[BaseSessionCompactManager] = None if session_config is None: session_config = SessionServiceConfig() # Clean up the TTL configuration if not set @@ -60,6 +67,8 @@ def __init__(self, self._session_config = session_config if self._summarizer_manager: self._summarizer_manager.set_session_service(self) + if session_compact_manager is not None: + self.set_session_compact_manager(session_compact_manager) @property def summarizer_manager(self) -> Optional[SummarizerSessionManager]: @@ -71,6 +80,11 @@ def session_config(self) -> SessionServiceConfig: """Get the session service configuration.""" return self._session_config + @property + def session_compact_manager(self) -> Optional["BaseSessionCompactManager"]: + """Get the Session Compact lifecycle manager.""" + return self._session_compact_manager + def set_summarizer_manager(self, summarizer_manager: SummarizerSessionManager, force: bool = False) -> None: """Set the summarizer manager to use. @@ -78,10 +92,27 @@ def set_summarizer_manager(self, summarizer_manager: SummarizerSessionManager, f summarizer_manager: The summarizer manager to use force: Whether to force update even if already set """ + if self._session_compact_manager is not None: + raise ValueError("SummarizerSessionManager and BaseSessionCompactManager are mutually exclusive") if not self._summarizer_manager or force: self._summarizer_manager = summarizer_manager self._summarizer_manager.set_session_service(self) + def set_session_compact_manager( + self, + compact_manager: "BaseSessionCompactManager", + force: bool = False, + ) -> None: + """Attach Session Compact through the native manager lifecycle.""" + if self._summarizer_manager is not None: + raise ValueError("SummarizerSessionManager and BaseSessionCompactManager are mutually exclusive") + if self._session_compact_manager is not None and not force: + if self._session_compact_manager is compact_manager: + return + raise ValueError("A Session Compact manager is already configured") + self._session_compact_manager = compact_manager + compact_manager.set_session_service(self, force=force) + @override async def append_event(self, session: Session, event: Event) -> Event: """Appends an event to a session object.""" @@ -174,6 +205,8 @@ async def create_session_summary(self, session: Session, ctx: Optional[Invocatio """ if self._summarizer_manager: await self._summarizer_manager.create_session_summary(session, ctx=ctx) + elif self._session_compact_manager: + await self._session_compact_manager.create_session_summary(session, ctx=ctx) @override async def get_session_summary(self, session: Session) -> Optional[str]: @@ -189,6 +222,8 @@ async def get_session_summary(self, session: Session) -> Optional[str]: summary = await self._summarizer_manager.get_session_summary(session) if summary: return summary.summary_text + if self._session_compact_manager: + return await self._session_compact_manager.get_session_summary(session) return None def filter_events(self, session: Session, need_copy: bool = False) -> Session: @@ -211,4 +246,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._session_compact_manager: + await self._session_compact_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..fdd1d1ce3 100644 --- a/trpc_agent_sdk/sessions/_in_memory_session_service.py +++ b/trpc_agent_sdk/sessions/_in_memory_session_service.py @@ -31,6 +31,7 @@ import uuid from typing import Any from typing import Optional +from typing import TYPE_CHECKING from typing_extensions import override from pydantic import BaseModel @@ -51,6 +52,9 @@ from ._utils import extract_state_delta from ._utils import merge_state +if TYPE_CHECKING: + from .compact._base_manager import BaseSessionCompactManager + class SessionWithTTL(BaseModel): """Wrapper for session with TTL support.""" @@ -108,8 +112,13 @@ class InMemorySessionService(BaseSessionService): def __init__(self, summarizer_manager: Optional[SummarizerSessionManager] = None, - session_config: Optional[SessionServiceConfig] = None): - super().__init__(summarizer_manager=summarizer_manager, session_config=session_config) + session_config: Optional[SessionServiceConfig] = None, + session_compact_manager: BaseSessionCompactManager | None = None): + super().__init__( + summarizer_manager=summarizer_manager, + session_config=session_config, + session_compact_manager=session_compact_manager, + ) # Storage with TTL support # Map: app_name -> user_id -> session_id -> SessionWithTTL self._sessions: dict[str, dict[str, dict[str, SessionWithTTL]]] = {} @@ -213,9 +222,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: @@ -294,6 +302,21 @@ async def update_session(self, session: Session) -> None: # Update the stored session and refresh TTL self._set_session(app_name, user_id, session_id, session) + @override + async def patch_session_state( + self, + session: Session, + state_delta: dict[str, Any], + ) -> None: + """Merge state into the stored session without replacing its Events.""" + stored = (self._sessions.get(session.app_name, {}).get(session.user_id, {}).get(session.id)) + if stored is None: + raise ValueError(f"Session {session.id} was not found") + stored.session.state.update(state_delta) + stored.ttl.update_expired_at() + session.state.update(state_delta) + session.last_update_time = time.time() + def _cleanup_expired(self) -> None: """Remove all expired sessions and states. diff --git a/trpc_agent_sdk/sessions/_redis_session_service.py b/trpc_agent_sdk/sessions/_redis_session_service.py index 8bec47af1..2fce605a8 100644 --- a/trpc_agent_sdk/sessions/_redis_session_service.py +++ b/trpc_agent_sdk/sessions/_redis_session_service.py @@ -8,10 +8,12 @@ from __future__ import annotations +import json import time import uuid from typing import Any from typing import Optional +from typing import TYPE_CHECKING from typing_extensions import override from trpc_agent_sdk.abc import ListSessionsResponse @@ -35,6 +37,9 @@ from ._utils import session_key from ._utils import user_state_key +if TYPE_CHECKING: + from .compact._base_manager import BaseSessionCompactManager + def _session_key_prefix(app_name: str, user_id: Optional[str] = None) -> str: """Generate a Redis key prefix for listing sessions. @@ -54,6 +59,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. @@ -79,15 +93,32 @@ def __init__(self, summarizer_manager: Optional[SummarizerSessionManager] = None, session_config: Optional[SessionServiceConfig] = None, is_async: bool = False, + session_compact_manager: BaseSessionCompactManager | None = 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, + session_compact_manager=session_compact_manager, + ) 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. @@ -251,6 +282,75 @@ async def update_session(self, session: Session) -> None: return await self._set_session(redis_session, session) + @override + async def patch_session_state( + self, + session: Session, + state_delta: dict[str, Any], + ) -> None: + """Atomically merge state while preserving concurrently written Events.""" + script = """ +local raw = redis.call('GET', KEYS[1]) +if not raw then + return false +end +local value = cjson.decode(raw) +local delta = cjson.decode(ARGV[1]) +if not value.state then + value.state = {} +end +for key, item in pairs(delta) do + value.state[key] = item +end +if type(value.events) == 'table' and next(value.events) == nil then + value.events = cjson.empty_array +end +if type(value.historical_events) == 'table' and next(value.historical_events) == nil then + value.historical_events = cjson.empty_array +end +if type(value.historicalEvents) == 'table' and next(value.historicalEvents) == nil then + value.historicalEvents = cjson.empty_array +end +local timestamp = tonumber(ARGV[2]) +if value.last_update_time ~= nil then + value.last_update_time = timestamp +end +if value.lastUpdateTime ~= nil then + value.lastUpdateTime = timestamp +end +local encoded = cjson.encode(value) +local ttl = tonumber(ARGV[3]) +if ttl > 0 then + redis.call('SET', KEYS[1], encoded, 'EX', ttl) +else + redis.call('SET', KEYS[1], encoded) +end +return encoded +""" + timestamp = time.time() + ttl = (int(self._session_config.ttl.ttl_seconds) if self._session_config.ttl.need_ttl_expire() else 0) + key = session_key(session.app_name, session.user_id, session.id) + async with self._redis_storage.create_db_session() as redis_session: + result = await self._redis_storage.execute_command( + redis_session, + RedisCommand( + method="eval", + args=( + script, + 1, + key, + json.dumps(state_delta, default=str), + timestamp, + ttl, + ), + ), + ) + if not result: + raise ValueError(f"Session {session.id} was not found") + stored_session = _session_from_storage_json(result) + session.state.update(state_delta) + session.last_update_time = stored_session.last_update_time + @override async def close(self) -> None: """Close the service and release resources.""" @@ -410,7 +510,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..1f998ba9b 100644 --- a/trpc_agent_sdk/sessions/_sql_session_service.py +++ b/trpc_agent_sdk/sessions/_sql_session_service.py @@ -34,6 +34,7 @@ from typing import Any from typing import List from typing import Optional +from typing import TYPE_CHECKING from typing_extensions import override from sqlalchemy import Boolean @@ -77,6 +78,9 @@ from ._utils import extract_state_delta from ._utils import merge_state +if TYPE_CHECKING: + from .compact._base_manager import BaseSessionCompactManager + def _event_field_or_default(field_name: str, value: Any) -> Any: """Use Event's default when legacy SQL rows contain NULL for non-null Event fields.""" @@ -391,18 +395,39 @@ def __init__(self, summarizer_manager: Optional[SummarizerSessionManager] = None, is_async: bool = False, session_config: Optional[SessionServiceConfig] = None, + session_compact_manager: BaseSessionCompactManager | None = 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, + session_compact_manager=session_compact_manager, + ) 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 +568,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 @@ -651,6 +679,29 @@ async def update_session(self, session: Session) -> None: session.last_update_time = storage_session.update_timestamp_tz + @override + async def patch_session_state( + self, + session: Session, + state_delta: dict[str, Any], + ) -> None: + """Merge state under a row lock without touching persisted Events.""" + key = SqlKey( + key=(session.app_name, session.user_id, session.id), + storage_cls=StorageSession, + ) + async with self._sql_storage.create_db_session() as sql_session: + storage_session: Optional[StorageSession] = (await self._sql_storage.get_for_update(sql_session, key)) + if storage_session is None: + raise ValueError(f"Session {session.id} was not found") + merged_state = dict(storage_session.state or {}) + merged_state.update(state_delta) + storage_session.state = merged_state # type: ignore + await self._sql_storage.commit(sql_session) + await self._sql_storage.refresh(sql_session, storage_session) + session.state.update(state_delta) + session.last_update_time = storage_session.update_timestamp_tz + @override async def close(self) -> None: self._stop_cleanup_task() @@ -704,6 +755,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 +769,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 +786,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/advanced_memory/__init__.py b/trpc_agent_sdk/sessions/compact/__init__.py similarity index 59% rename from trpc_agent_sdk/advanced_memory/__init__.py rename to trpc_agent_sdk/sessions/compact/__init__.py index c346cd2d3..1fe406010 100644 --- a/trpc_agent_sdk/advanced_memory/__init__.py +++ b/trpc_agent_sdk/sessions/compact/__init__.py @@ -3,7 +3,7 @@ # 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.""" +"""Canonical context-compression package for session management.""" from ._autocompact import AutoCompact from ._autocompact import AutoCompactCallback @@ -11,34 +11,24 @@ 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 ._base_manager import BaseSessionCompactManager +from ._config import AdvancedCompactConfig +from ._formats import build_session_memory_state +from ._formats import parse_session_memory_state 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 ._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 ._manager import AdvancedSessionCompactManager 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 @@ -46,87 +36,61 @@ 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 +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 ._runtime import ScopedSessionCompactRuntime +from ._runtime import SessionCompactRuntime __all__ = [ + "AdvancedCompactConfig", "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", + "HeuristicTokenEstimator", "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", + "SESSION_MEMORY_STATE_KEY", "SessionMemoryDocument", "SessionMemoryExtractionInput", "SessionMemoryExtractionResult", "SessionMemoryExtractor", - "SessionMemoryStore", - "TRANSCRIPT_SCHEMA_VERSION", + "BaseSessionCompactManager", + "AdvancedSessionCompactManager", + "TokenContextTracker", + "TokenEstimator", "ToolResultBudget", "ToolResultBudgetCallback", "ToolResultBudgetResult", - "ToolResultStore", - "TokenContextTracker", - "TokenEstimator", - "TranscriptSessionService", - "TranscriptStore", - "setup_autocompact", - "setup_advanced_memory", + "SessionCompactRuntime", + "ScopedSessionCompactRuntime", + "build_session_memory_prompt", + "build_session_memory_state", + "content_signature", + "estimate_request_chars", + "has_session_memory_content", "limit_session_memory_document", + "parse_session_memory_state", + "setup_autocompact", "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/sessions/compact/_autocompact.py similarity index 69% rename from trpc_agent_sdk/advanced_memory/_autocompact.py rename to trpc_agent_sdk/sessions/compact/_autocompact.py index e7efb436a..af4630b5a 100644 --- a/trpc_agent_sdk/advanced_memory/_autocompact.py +++ b/trpc_agent_sdk/sessions/compact/_autocompact.py @@ -8,6 +8,7 @@ from __future__ import annotations import asyncio +import copy import hashlib import json import re @@ -18,6 +19,7 @@ from typing import TYPE_CHECKING from trpc_agent_sdk.agents import LlmAgent +from trpc_agent_sdk.events import Event from trpc_agent_sdk.models import LlmResponse from trpc_agent_sdk.runners import Runner from trpc_agent_sdk.sessions import InMemorySessionService @@ -26,24 +28,26 @@ from ._callbacks import install_staged_callback 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 AdvancedMemoryRuntime +from ._runtime import SessionCompactRuntime 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 + from ._session_memory import SessionMemoryExtractor -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. +The complete original events remain available in the SessionService. """ _LEGACY_SESSION_MEMORY_SECTION_LIST = "\n".join(f"- # {section}" for section in SESSION_MEMORY_SECTIONS) @@ -68,6 +72,8 @@ class AutoCompactRecord: boundary_occurrence: int summary: str source: str + boundary_event_id: str | None = None + compaction_id: str | None = None @dataclass @@ -222,59 +228,53 @@ class AutoCompact: def __init__( self, - memory_runtime: AdvancedMemoryRuntime, + memory_runtime: SessionCompactRuntime, summary_generator: LegacySummaryGenerator | None = None, *, model: Any | None = None, + session_memory_extractor: "SessionMemoryExtractor | 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._session_memory_extractor = session_memory_extractor self._states: dict[str, AutoCompactState] = {} self._session_locks: dict[str, asyncio.Lock] = {} + self._scoped_processors: dict[object, "AutoCompact"] = {} @property - def runtime(self) -> AdvancedMemoryRuntime: + def runtime(self) -> SessionCompactRuntime: """Return the runtime bound to this compressor.""" return self._runtime + def attach_session_memory_extractor( + self, + extractor: "SessionMemoryExtractor", + ) -> None: + """Attach the extractor invoked only when AutoCompact is reached.""" + if (self._session_memory_extractor is not None and self._session_memory_extractor is not extractor): + raise ValueError("Autocompact session memory extractor is already configured") + self._session_memory_extractor = extractor + def _session_lock(self, session_id: str) -> asyncio.Lock: """Return the unique 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) -> AutoCompactState: - """Restore the latest compaction and failure count from the transcript.""" - state = self._states.get(session_id) + """Restore process-local compaction 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) - 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 + state = AutoCompactState(latest_compaction=None, consecutive_failures=0) + self._states[state_key] = state return state def _summary_content(self, summary: str) -> Content: @@ -285,12 +285,11 @@ def _summary_content(self, summary: str) -> Content: ) def _summary_with_recovery_path(self, summary: str, session_id: str) -> str: - """Append recovery paths for the full transcript and session memory.""" + """Tell the model where the authoritative compacted data lives.""" 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)}") + "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, @@ -307,6 +306,17 @@ def _find_signature_index( 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], @@ -372,44 +382,24 @@ def _apply_record(self, request: "LlmRequest", record: AutoCompactRecord) -> boo 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 + ctx: "InvocationContext", + ) -> tuple[str, str, int, str] | None: + """Read Session Memory and its checkpoint from Session.state.""" + state = getattr(ctx.session, "state", {}) + parsed = parse_session_memory_state(state.get(SESSION_MEMORY_STATE_KEY) if isinstance(state, dict) else None) + 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, @@ -419,6 +409,7 @@ def _compact_with_summary( boundary_index: int, source: str, strict_boundary: bool = False, + boundary_event_id: str | None = None, ) -> AutoCompactRecord: """Replace the old prefix with a summary and return a replay record.""" boundary_signature = content_signature(request.contents[boundary_index]) @@ -435,7 +426,95 @@ def _compact_with_summary( boundary_occurrence, summary, source, + boundary_event_id, + f"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 getattr(ctx.session, "events", []) or []: + 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 (getattr(ctx.session, "events", []) or []) if getattr(event, "content", None) is not None + ] + if len(content_events) <= 1: + return None + keep_count = min( + self._runtime.config.autocompact_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: AutoCompactRecord, + ) -> None: + """Persist the compacted active window through the original SessionService.""" + compact_events = getattr(ctx.session, "compact_events", None) + if not callable(compact_events): + # AutoCompact 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 AutoCompact boundary to an active Session Event") + + compaction_id = record.compaction_id or f"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.""" @@ -471,66 +550,28 @@ async def _legacy_summary( 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( + async def apply( self, + request: "LlmRequest", + *, 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( + ctx: "InvocationContext", + force: bool = False, + ) -> AutoCompactResult: + """Run compaction against the current session's tenant namespace.""" + if hasattr(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, session_id=session_id, ctx=ctx, force=force) + + async def _apply_scoped( self, request: "LlmRequest", *, @@ -544,7 +585,6 @@ async def apply( 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) @@ -556,6 +596,7 @@ async def apply( 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 >= config.autocompact_blocking_chars) if state.consecutive_failures >= config.autocompact_max_failures and blocking_reached: @@ -594,38 +635,48 @@ async def apply( 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 self._session_memory_extractor is not None: + await self._session_memory_extractor.extract_if_needed( + ctx.session, + ctx, + force=True, + ) + session_memory = await self._latest_session_memory_record( + session_id, + ctx, + ) 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, + memory, boundary_signature, boundary_occurrence, boundary_event_id = session_memory + boundary_index = self._find_signature_index( + request.contents, + boundary_signature, + boundary_occurrence, ) - if boundary is not None: - boundary_signature, boundary_occurrence = boundary - boundary_index = self._find_signature_index( + if boundary_index is None and reapplied: + boundary_index = self._find_last_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, - ) + 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, + boundary_event_id=boundary_event_id, + ) + if token_mode: 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 + <= token_budget_before.warning_threshold_tokens) + else: + target_reached = 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( @@ -648,23 +699,18 @@ async def apply( ), boundary_index=boundary_index, source="legacy", + boundary_event_id=self._legacy_boundary_event_id(ctx), ) request_chars_after = estimate_request_chars(request) - if request_chars_after >= request_chars_before: + 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("Autocompact did not reduce request token estimate") + elif 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, - ) + await self._persist_session_compaction(ctx, compact_record) state.latest_compaction = compact_record state.consecutive_failures = 0 return AutoCompactResult( @@ -675,19 +721,13 @@ async def apply( 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, + 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, ) 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, @@ -743,7 +783,7 @@ async def __call__( def setup_autocompact( agent: "ParentLlmAgent", - memory_runtime: AdvancedMemoryRuntime, + memory_runtime: SessionCompactRuntime, summary_generator: LegacySummaryGenerator | None = None, *, model: Any | None = None, diff --git a/trpc_agent_sdk/sessions/compact/_base_manager.py b/trpc_agent_sdk/sessions/compact/_base_manager.py new file mode 100644 index 000000000..2a861ec8e --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/_base_manager.py @@ -0,0 +1,52 @@ +# 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. +"""Define the Session Compact manager lifecycle contract.""" + +from __future__ import annotations + +from abc import ABC +from abc import abstractmethod +from typing import Any +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from trpc_agent_sdk.abc import SessionServiceABC + from trpc_agent_sdk.context import InvocationContext + from trpc_agent_sdk.sessions import Session + + +class BaseSessionCompactManager(ABC): + """Coordinate one Session Compact implementation with a SessionService.""" + + @abstractmethod + def setup(self, agent: Any) -> None: + """Initialize this manager and install its Agent callbacks.""" + + @abstractmethod + def set_session_service( + self, + session_service: "SessionServiceABC", + force: bool = False, + ) -> None: + """Bind this manager to the SessionService that owns its sessions.""" + + @abstractmethod + async def create_session_summary( + self, + session: "Session", + force: bool = False, + ctx: "InvocationContext | None" = None, + ) -> None: + """Update compact state through the SessionService post-turn hook.""" + + @abstractmethod + async def get_session_summary(self, session: "Session") -> str | None: + """Return the compact representation exposed as a session summary.""" + + @abstractmethod + async def close(self) -> None: + """Release resources owned by this manager.""" diff --git a/trpc_agent_sdk/advanced_memory/_callbacks.py b/trpc_agent_sdk/sessions/compact/_callbacks.py similarity index 94% rename from trpc_agent_sdk/advanced_memory/_callbacks.py rename to trpc_agent_sdk/sessions/compact/_callbacks.py index 99b422346..a678fffb1 100644 --- a/trpc_agent_sdk/advanced_memory/_callbacks.py +++ b/trpc_agent_sdk/sessions/compact/_callbacks.py @@ -9,7 +9,7 @@ from typing import Any -from ._runtime import AdvancedMemoryRuntime +from ._runtime import SessionCompactRuntime def install_staged_callback( @@ -18,7 +18,7 @@ def install_staged_callback( *, callback_type: type, component_attribute: str, - memory_runtime: AdvancedMemoryRuntime, + memory_runtime: SessionCompactRuntime, conflict_message: str, ) -> Any | None: """Install a staged callback idempotently and validate runtime ownership.""" diff --git a/trpc_agent_sdk/sessions/compact/_config.py b/trpc_agent_sdk/sessions/compact/_config.py new file mode 100644 index 000000000..1692bf2d6 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/_config.py @@ -0,0 +1,130 @@ +# 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 typing import Any + +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 AdvancedCompactConfig: + """Configure compression that is persisted by the SessionService.""" + + 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=None) + max_output_tokens: int = 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 + + def __post_init__(self) -> None: + """Validate compression limits and token thresholds.""" + _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, + 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, + autocompact_target_chars=self.autocompact_target_chars, + 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, + microcompact_gap_seconds=self.microcompact_gap_seconds, + microcompact_trigger_count=self.microcompact_trigger_count, + microcompact_keep_recent=self.microcompact_keep_recent, + ) + _require_non_negative( + max_output_tokens=self.max_output_tokens, + session_memory_request_overhead_tokens=self.session_memory_request_overhead_tokens, + ) + if self.model_context_window_tokens is not None: + _require_positive(model_context_window_tokens=self.model_context_window_tokens) + if 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") + if self.tool_result_preview_chars >= self.tool_result_max_chars: + raise ValueError("tool_result_preview_chars must be smaller than tool_result_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.autocompact_trigger_chars <= self.autocompact_target_chars: + raise ValueError("autocompact_trigger_chars must be greater than autocompact_target_chars") + if self.autocompact_blocking_chars <= self.autocompact_trigger_chars: + raise ValueError("autocompact_blocking_chars must be greater than autocompact_trigger_chars") + _require_non_empty_names("history_snip_tool_names", self.history_snip_tool_names) + _require_non_empty_names("microcompact_tool_names", self.microcompact_tool_names) diff --git a/trpc_agent_sdk/advanced_memory/_coordination.py b/trpc_agent_sdk/sessions/compact/_coordination.py similarity index 100% rename from trpc_agent_sdk/advanced_memory/_coordination.py rename to trpc_agent_sdk/sessions/compact/_coordination.py diff --git a/trpc_agent_sdk/advanced_memory/_formats.py b/trpc_agent_sdk/sessions/compact/_formats.py similarity index 75% rename from trpc_agent_sdk/advanced_memory/_formats.py rename to trpc_agent_sdk/sessions/compact/_formats.py index f0fa25ad1..33c3841b6 100644 --- a/trpc_agent_sdk/advanced_memory/_formats.py +++ b/trpc_agent_sdk/sessions/compact/_formats.py @@ -5,10 +5,14 @@ # 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 import re +from dataclasses import asdict from dataclasses import dataclass +from dataclasses import fields from datetime import datetime from datetime import timezone from enum import Enum @@ -137,6 +141,8 @@ def memory_freshness(updated_at: datetime | None, *, now: datetime | None = None "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", @@ -189,3 +195,53 @@ def to_markdown(self) -> str: ) ] 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/_history_snip.py similarity index 84% rename from trpc_agent_sdk/advanced_memory/_history_snip.py rename to trpc_agent_sdk/sessions/compact/_history_snip.py index 67d8ff1a6..5f8867551 100644 --- a/trpc_agent_sdk/advanced_memory/_history_snip.py +++ b/trpc_agent_sdk/sessions/compact/_history_snip.py @@ -8,13 +8,14 @@ from __future__ import annotations import asyncio +import copy import json from dataclasses import dataclass from typing import Any from typing import TYPE_CHECKING from ._callbacks import install_staged_callback -from ._runtime import AdvancedMemoryRuntime +from ._runtime import SessionCompactRuntime 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 @@ -26,7 +27,6 @@ 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]" @@ -83,46 +83,38 @@ 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, memory_runtime: SessionCompactRuntime) -> None: """Initialize history-snip state and per-session async locks.""" self._runtime = memory_runtime 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) -> SessionCompactRuntime: """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]: @@ -158,29 +150,6 @@ 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", @@ -191,12 +160,32 @@ async def apply( ) -> 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: request_chars = estimate_request_chars(request) return HistorySnipResult(None, 0, 0, 0, request_chars, request_chars) - - await self._runtime.initialize() + if ctx is None or hasattr(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, session_id=session_id, 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.""" + config = self._runtime.config + tracker = TokenContextTracker(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) @@ -246,7 +235,6 @@ async def apply( 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 @@ -292,7 +280,7 @@ async def __call__(self, ctx: "InvocationContext", request: "LlmRequest") -> Non def setup_history_snip( agent: "LlmAgent", - memory_runtime: AdvancedMemoryRuntime, + memory_runtime: SessionCompactRuntime, ) -> HistorySnip: """Install history snip while preserving context stage order.""" history_snip = HistorySnip(memory_runtime) diff --git a/trpc_agent_sdk/sessions/compact/_manager.py b/trpc_agent_sdk/sessions/compact/_manager.py new file mode 100644 index 000000000..f969d0cda --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/_manager.py @@ -0,0 +1,135 @@ +# 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 import Any +from typing import TYPE_CHECKING + +from ._base_manager import BaseSessionCompactManager +from ._formats import parse_session_memory_state +from ._formats import SESSION_MEMORY_STATE_KEY + +if TYPE_CHECKING: + from trpc_agent_sdk.abc import SessionServiceABC + from trpc_agent_sdk.context import InvocationContext + from trpc_agent_sdk.sessions import Session + +from ._autocompact import LegacySummaryGenerator +from ._config import AdvancedCompactConfig +from ._runtime import SessionCompactRuntime +from ._session_memory import SessionMemoryExtractor +from ._session_memory import SessionMemoryGenerator + + +class AdvancedSessionCompactManager(BaseSessionCompactManager): + """Coordinate Advanced Compact state without wrapping a SessionService.""" + + def __init__( + self, + config: AdvancedCompactConfig, + *, + summary_generator: "LegacySummaryGenerator | None" = None, + compact_model: Any | None = None, + session_memory_generator: "SessionMemoryGenerator | None" = None, + session_memory_model: Any | None = None, + ) -> None: + """Store configuration until Runner supplies the Agent.""" + self._config = config + self._summary_generator = summary_generator + self._compact_model = compact_model + self._session_memory_generator = session_memory_generator + self._session_memory_model = session_memory_model + self._runtime: SessionCompactRuntime | None = None + self._session_memory_extractor: SessionMemoryExtractor | None = None + self._session_service: SessionServiceABC | None = None + + def setup(self, agent: Any) -> None: + """Initialize the runtime and install all compression callbacks.""" + if self._session_service is None: + raise RuntimeError("Session Compact manager must be bound to a SessionService first") + if self._runtime is not None: + return + from ._autocompact import setup_autocompact + from ._history_snip import setup_history_snip + from ._microcompact import setup_microcompact + from ._tool_result_budget import setup_tool_result_budget + + runtime = SessionCompactRuntime.create(self._config) + extractor = SessionMemoryExtractor( + runtime, + self._session_memory_generator, + model=self._session_memory_model, + ) + setup_tool_result_budget(agent, runtime) + setup_history_snip(agent, runtime) + setup_microcompact(agent, runtime) + autocompact = setup_autocompact( + agent, + runtime, + self._summary_generator, + model=self._compact_model, + ) + autocompact.attach_session_memory_extractor(extractor) + extractor.attach_session_service(self._session_service) + self._runtime = runtime + self._session_memory_extractor = extractor + + @property + def runtime(self) -> "SessionCompactRuntime": + """Return the runtime shared by all compact stages.""" + if self._runtime is None: + raise RuntimeError("Session Compact manager has not been initialized by Runner") + return self._runtime + + @property + def session_memory_extractor(self) -> "SessionMemoryExtractor": + """Return the post-turn Session Memory extractor.""" + if self._session_memory_extractor is None: + raise RuntimeError("Session Compact manager has not been initialized by Runner") + return self._session_memory_extractor + + def set_session_service( + self, + session_service: "SessionServiceABC", + force: bool = False, + ) -> None: + """Bind the manager to the original persistence service.""" + if self._session_service is not None and self._session_service is not session_service and not force: + raise ValueError("AdvancedSessionCompactManager is already bound to another SessionService") + session_config = getattr(session_service, "session_config", None) + if session_config is None or not getattr(session_config, "store_historical_events", False): + raise ValueError("Advanced Session Compact requires " + "SessionServiceConfig(store_historical_events=True)") + self._session_service = session_service + if self._session_memory_extractor is not None: + self._session_memory_extractor.attach_session_service(session_service) + + async def create_session_summary( + self, + session: "Session", + force: bool = False, + ctx: "InvocationContext | None" = None, + ) -> None: + """Use the native post-turn hook to update persistent Session Memory.""" + if ctx is not None and self._session_memory_extractor is not None: + await self._session_memory_extractor.extract_if_needed( + session, + ctx, + 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 + + async def close(self) -> None: + """Release Compact resources owned by the manager.""" diff --git a/trpc_agent_sdk/advanced_memory/_microcompact.py b/trpc_agent_sdk/sessions/compact/_microcompact.py similarity index 81% rename from trpc_agent_sdk/advanced_memory/_microcompact.py rename to trpc_agent_sdk/sessions/compact/_microcompact.py index eeaabdd36..896c45212 100644 --- a/trpc_agent_sdk/advanced_memory/_microcompact.py +++ b/trpc_agent_sdk/sessions/compact/_microcompact.py @@ -8,13 +8,14 @@ from __future__ import annotations import asyncio +import copy import time from dataclasses import dataclass from typing import Any from typing import TYPE_CHECKING from ._callbacks import install_staged_callback -from ._runtime import AdvancedMemoryRuntime +from ._runtime import SessionCompactRuntime 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 @@ -25,7 +26,6 @@ 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]" @@ -74,46 +74,38 @@ def find_last_assistant_timestamp(ctx: "InvocationContext") -> float | None: class Microcompact: """Local compressor that cleans old tool results by time or count.""" - def __init__(self, memory_runtime: AdvancedMemoryRuntime) -> None: + def __init__(self, memory_runtime: SessionCompactRuntime) -> None: """Initialize mechanical-compaction state and per-session locks.""" self._runtime = memory_runtime self._states: dict[str, MicrocompactState] = {} self._session_locks: dict[str, asyncio.Lock] = {} + self._scoped_processors: dict[object, "Microcompact"] = {} @property - def runtime(self) -> AdvancedMemoryRuntime: + def runtime(self) -> SessionCompactRuntime: """Return the runtime bound to this mechanical compressor.""" return self._runtime 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) + """Return process-local microcompact 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, + 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]: @@ -148,42 +140,52 @@ def _cleared_response(self) -> dict[str, str]: """Return the minimal placeholder shared by cleanups.""" return {"output": MICROCOMPACT_CLEARED_MESSAGE} - async def _persist_clear( - 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", - ) - async def apply( self, request: "LlmRequest", *, session_id: str, last_assistant_timestamp: float | None, + ctx: "InvocationContext | None" = 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() + if ctx is None or hasattr(self._runtime, "scope"): + return await self._apply_scoped( + request, + session_id=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, + session_id=session_id, + last_assistant_timestamp=last_assistant_timestamp, + ctx=ctx, + now=now, + ) + + async def _apply_scoped( + self, + request: "LlmRequest", + *, + session_id: str, + last_assistant_timestamp: float | None, + now: float | None, + ) -> MicrocompactResult: + """Apply one tenant-bound mechanical compaction.""" + config = self._runtime.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) @@ -221,7 +223,6 @@ async def apply( 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 @@ -255,13 +256,14 @@ async def __call__(self, ctx: "InvocationContext", request: "LlmRequest") -> Non request, session_id=ctx.session_id, last_assistant_timestamp=find_last_assistant_timestamp(ctx), + ctx=ctx, ) return None def setup_microcompact( agent: "LlmAgent", - memory_runtime: AdvancedMemoryRuntime, + memory_runtime: SessionCompactRuntime, ) -> Microcompact: """Install the mechanical callback while preserving existing order.""" microcompact = Microcompact(memory_runtime) diff --git a/trpc_agent_sdk/sessions/compact/_runtime.py b/trpc_agent_sdk/sessions/compact/_runtime.py new file mode 100644 index 000000000..f5734e2c5 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/_runtime.py @@ -0,0 +1,53 @@ +# 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 + +from ._config import AdvancedCompactConfig +from ._coordination import SessionOperationCoordinator + + +@dataclass +class SessionCompactRuntime: + """Hold compression configuration and per-session coordination only.""" + + config: AdvancedCompactConfig + coordination: SessionOperationCoordinator + + @classmethod + def create(cls, config: AdvancedCompactConfig | None = None) -> "SessionCompactRuntime": + return cls(config or AdvancedCompactConfig(), SessionOperationCoordinator()) + + def for_session(self, session: object) -> "ScopedSessionCompactRuntime": + 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("Session Compact requires session app_name and user_id") + return ScopedSessionCompactRuntime(self, f"{app_name}\0{user_id}") + + +@dataclass +class ScopedSessionCompactRuntime: + """Session-scoped view used by compression callbacks.""" + + root: SessionCompactRuntime + scope: str + + @property + def config(self) -> AdvancedCompactConfig: + return self.root.config + + @property + def coordination(self) -> SessionOperationCoordinator: + return self.root.coordination + + def session_key(self, session_id: str) -> str: + return f"{self.scope}\0{session_id}" + + def for_session(self, session: object) -> "ScopedSessionCompactRuntime": + return self.root.for_session(session) diff --git a/trpc_agent_sdk/advanced_memory/_session_memory.py b/trpc_agent_sdk/sessions/compact/_session_memory.py similarity index 82% rename from trpc_agent_sdk/advanced_memory/_session_memory.py rename to trpc_agent_sdk/sessions/compact/_session_memory.py index 39f9ab454..2e9961289 100644 --- a/trpc_agent_sdk/advanced_memory/_session_memory.py +++ b/trpc_agent_sdk/sessions/compact/_session_memory.py @@ -11,6 +11,8 @@ from collections import Counter from dataclasses import dataclass from dataclasses import fields +from datetime import datetime +from datetime import timezone import re from typing import Any from typing import Protocol @@ -26,15 +28,18 @@ 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 SessionCompactRuntime from ._token_budget import TokenContextTracker if TYPE_CHECKING: from trpc_agent_sdk.abc import SessionABC + from trpc_agent_sdk.abc import SessionServiceABC 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 +149,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. """ @@ -336,10 +341,11 @@ class SessionMemoryExtractor: def __init__( self, - memory_runtime: AdvancedMemoryRuntime, + memory_runtime: SessionCompactRuntime, generator: SessionMemoryGenerator | None = None, *, model: Any | None = None, + session_service: "SessionServiceABC | None" = None, ) -> None: """Initialize extraction and per-session serialization locks.""" if generator is not None and model is not None: @@ -349,19 +355,58 @@ def __init__( model, section_max_chars=memory_runtime.config.session_memory_section_max_chars, ) + self._session_service = session_service @property - def runtime(self) -> AdvancedMemoryRuntime: + def runtime(self) -> SessionCompactRuntime: """Return the runtime bound to this extractor.""" return self._runtime + def attach_session_service(self, session_service: "SessionServiceABC") -> None: + """Attach the service used for atomic state-only writes.""" + if self._session_service is not None and self._session_service is not session_service: + raise ValueError("Session memory extractor is already bound to another service") + self._session_service = session_service + + def _session_event_records(self, session: "SessionABC") -> 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(getattr(session, "events", None) or []) + for event in events: + is_summary_event = getattr(event, "is_summary_event", None) + if callable(is_summary_event) and is_summary_event(): + continue + event_id = getattr(event, "id", None) + if not isinstance(event_id, str) or event_id in seen: + continue + seen.add(event_id) + timestamp = float(getattr(event, "timestamp", 0.0) or 0.0) + records.append({ + "kind": "event", + "event_id": event_id, + "recorded_at": datetime.fromtimestamp( + timestamp, + tz=timezone.utc, + ).isoformat(), + "event": event.model_dump( + mode="json", + by_alias=True, + exclude_none=True, + ), + }) + return records + def _event_records_after_checkpoint( self, records: list[dict[str, Any]], 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 +426,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, @@ -487,7 +522,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,7 +537,7 @@ 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) @@ -607,14 +642,56 @@ 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: "SessionABC") -> 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: "SessionABC", + ) -> 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: "SessionABC", + event_id: str, + ) -> tuple[str, int] | None: + """Return a model-content signature and occurrence for one Event.""" + from ._autocompact import content_signature + + signatures: list[str] = [] + # AutoCompact matches against the active model request, so occurrence + # counts must not include archived Events. + events = list(getattr(session, "events", None) or []) + seen_ids: set[str] = set() + for event in events: + current_id = getattr(event, "id", None) + if not isinstance(current_id, str) or current_id in seen_ids: + continue + seen_ids.add(current_id) + content = getattr(event, "content", None) + if content is None: + continue + signature = content_signature(content) + signatures.append(signature) + if current_id == event_id: + return signature, signatures.count(signature) + return None async def _persist_checkpoint( self, - session_id: str, + session: "SessionABC", included_records: list[dict[str, Any]], document: SessionMemoryDocument, context_tokens: int | None, @@ -634,20 +711,31 @@ 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", + if self._session_service is None: + raise RuntimeError("Session Memory requires a SessionService") + 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, + ) + await self._session_service.patch_session_state( + session, + {SESSION_MEMORY_STATE_KEY: payload}, ) async def extract_if_needed( @@ -661,12 +749,13 @@ async def extract_if_needed( 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: + runtime = self._runtime.for_session(session) + session_key = runtime.session_key(session.id) + async with self._runtime.coordination.guard(session_key) as acquired: if not acquired: return SessionMemoryExtractionResult(False, "coordination-timeout") - records = await self._runtime.transcripts.read_all(session.id) - checkpoint = self._last_checkpoint(records) + 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( @@ -681,8 +770,6 @@ async def extract_if_needed( tracker = TokenContextTracker(config) 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 @@ -699,7 +786,7 @@ async def extract_if_needed( return SessionMemoryExtractionResult(False, "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, @@ -715,9 +802,8 @@ async def extract_if_needed( max_chars=config.session_memory_section_max_chars, total_max_chars=config.session_memory_total_max_chars, ) - await self._runtime.session_memory.write(session.id, document) await self._persist_checkpoint( - session.id, + session, included, document, context_tokens, diff --git a/trpc_agent_sdk/advanced_memory/_token_budget.py b/trpc_agent_sdk/sessions/compact/_token_budget.py similarity index 98% rename from trpc_agent_sdk/advanced_memory/_token_budget.py rename to trpc_agent_sdk/sessions/compact/_token_budget.py index 544fefd2a..aad9af666 100644 --- a/trpc_agent_sdk/advanced_memory/_token_budget.py +++ b/trpc_agent_sdk/sessions/compact/_token_budget.py @@ -218,6 +218,10 @@ 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 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 | None" = None) -> bool: """Return whether the configuration resolves a model context window.""" return self._resolve_window_tokens(ctx) is not None diff --git a/trpc_agent_sdk/advanced_memory/_tool_result_budget.py b/trpc_agent_sdk/sessions/compact/_tool_result_budget.py similarity index 70% rename from trpc_agent_sdk/advanced_memory/_tool_result_budget.py rename to trpc_agent_sdk/sessions/compact/_tool_result_budget.py index 73a710aa2..3ad6cda17 100644 --- a/trpc_agent_sdk/advanced_memory/_tool_result_budget.py +++ b/trpc_agent_sdk/sessions/compact/_tool_result_budget.py @@ -8,15 +8,15 @@ 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 ._callbacks import install_staged_callback -from ._runtime import AdvancedMemoryRuntime +from ._runtime import SessionCompactRuntime if TYPE_CHECKING: from trpc_agent_sdk.agents import LlmAgent @@ -40,6 +40,7 @@ class ToolResultCandidate: """Describe a function response candidate in a model request.""" result_id: str + event_id: str | None tool_name: str serialized_result: str original_size: int @@ -48,10 +49,9 @@ class ToolResultCandidate: @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 @@ -66,7 +66,7 @@ class ToolResultBudgetResult: 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 +91,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 +112,58 @@ 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, memory_runtime: SessionCompactRuntime) -> None: """Initialize the budget processor and per-session state locks.""" self._runtime = memory_runtime self._states: dict[str, ToolResultBudgetState] = {} self._session_locks: dict[str, asyncio.Lock] = {} + self._scoped_processors: dict[object, "ToolResultBudget"] = {} @property - def runtime(self) -> AdvancedMemoryRuntime: + def runtime(self) -> SessionCompactRuntime: """Return the runtime bound to this budget processor.""" return self._runtime 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) 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) -> 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) 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) - 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: Any | None = None, + ) -> list[list[ToolResultCandidate]]: """Group function responses from consecutive user contents.""" + event_ids: dict[str, str] = {} + for event in getattr(session, "events", []) or []: + event_id = getattr(event, "id", None) + if not isinstance(event_id, str): + continue + event_content = getattr(event, "content", None) + for event_part in getattr(event_content, "parts", []) or []: + response = getattr(event_part, "function_response", None) + response_id = getattr(response, "id", None) + if isinstance(response_id, str): + event_ids[response_id] = event_id candidate_groups: list[list[ToolResultCandidate]] = [] serialized_by_result_id: dict[str, str] = {} current_group: list[ToolResultCandidate] = [] @@ -191,6 +187,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,11 +199,9 @@ 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, @@ -216,24 +211,21 @@ def _build_replacement( "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]: @@ -248,7 +240,7 @@ 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) + selected[candidate.result_id] = self._build_replacement(candidate) visible_size = 0 remaining_fresh: list[ToolResultCandidate] = [] @@ -267,80 +259,55 @@ def _select_replacements( for candidate in sorted(remaining_fresh, key=lambda item: item.original_size, reverse=True): if visible_size <= config.tool_results_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( + async def apply( self, + request: "LlmRequest", + *, 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: + ctx: "InvocationContext | None" = None, + ) -> ToolResultBudgetResult: """Process a model request without mutating session Events.""" if not self._runtime.config.enabled: return ToolResultBudgetResult(0, 0, 0) - await self._runtime.initialize() + if ctx is None or hasattr(self._runtime, "scope"): + return await self._apply_scoped(request, session_id, getattr(ctx, "session", None)) + 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, session_id=session_id, ctx=ctx) + + async def _apply_scoped( + self, + request: "LlmRequest", + session_id: str, + session: Any | None, + ) -> 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) + 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 +316,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) @@ -390,13 +356,13 @@ def budget(self) -> ToolResultBudget: async def __call__(self, ctx: "InvocationContext", request: "LlmRequest") -> None: """Apply tool-result budgeting without truncating model calls.""" - await self._budget.apply(request, session_id=ctx.session_id) + await self._budget.apply(request, session_id=ctx.session_id, ctx=ctx) return None def setup_tool_result_budget( agent: "LlmAgent", - memory_runtime: AdvancedMemoryRuntime, + memory_runtime: SessionCompactRuntime, ) -> ToolResultBudget: """Install the budget callback while preserving existing callbacks.""" budget = ToolResultBudget(memory_runtime) 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/_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()