From 36ee472f49de6fdad96b933ff2323b03fa6c4b3e Mon Sep 17 00:00:00 2001 From: congkechen Date: Wed, 9 Sep 2026 10:32:21 +0800 Subject: [PATCH 1/6] =?UTF-8?q?feature:=20advanced=20memory=20=E6=9C=8D?= =?UTF-8?q?=E5=8A=A1=E5=8C=96=EF=BC=8C=E6=94=AF=E6=8C=81=20redis/sql=20=20?= =?UTF-8?q?Please=20enter=20the=20commit=20message=20for=20your=20changes.?= =?UTF-8?q?=20Lines=20starting?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../memory_service_with_advanced_memory/.env | 4 + .../README.md | 101 +++- .../run_agent.py | 31 +- .../.env | 11 + .../README.md | 428 +++++++++++++ .../agent/__init__.py | 1 + .../agent/agent.py | 27 + .../agent/tools.py | 21 + .../run_agent.py | 150 +++++ .../.env | 17 + .../README.md | 194 ++++++ .../agent/__init__.py | 1 + .../agent/agent.py | 25 + .../agent/config.py | 14 + .../agent/prompts.py | 8 + .../agent/tools.py | 21 + .../run_agent.py | 139 +++++ .../test_advanced_memory_session_service.py | 39 +- .../test_advanced_memory_tools.py | 2 +- tests/advanced_memory/test_autocompact.py | 11 +- tests/advanced_memory/test_memory_context.py | 18 + tests/advanced_memory/test_preload_memory.py | 20 +- tests/advanced_memory/test_redis_stores.py | 104 ++++ .../test_session_memory_extractor.py | 20 +- tests/advanced_memory/test_sql_stores.py | 84 +++ tests/advanced_memory/test_storage.py | 88 +++ .../test_tool_result_budget.py | 22 + .../test_transcript_session_service.py | 10 +- trpc_agent_sdk/advanced_memory/__init__.py | 8 + .../advanced_memory/_autocompact.py | 33 +- trpc_agent_sdk/advanced_memory/_config.py | 29 + .../advanced_memory/_history_snip.py | 39 +- .../advanced_memory/_memory_context.py | 26 +- .../advanced_memory/_microcompact.py | 49 +- trpc_agent_sdk/advanced_memory/_paths.py | 44 +- .../advanced_memory/_preload_memory.py | 13 +- .../advanced_memory/_redis_stores.py | 299 ++++++++++ trpc_agent_sdk/advanced_memory/_runtime.py | 195 ++++++ .../advanced_memory/_session_memory.py | 25 +- .../advanced_memory/_session_service.py | 51 +- trpc_agent_sdk/advanced_memory/_sql_stores.py | 560 ++++++++++++++++++ trpc_agent_sdk/advanced_memory/_storage.py | 215 ++++++- .../advanced_memory/_storage_backend.py | 30 + .../advanced_memory/_tool_result_budget.py | 56 +- .../memory/_advanced_memory_service.py | 2 +- .../_advanced_memory_session_service.py | 128 ++-- .../sessions/_sql_session_service.py | 10 + trpc_agent_sdk/storage/_sql.py | 12 + trpc_agent_sdk/tools/_advanced_memory_tool.py | 44 +- 49 files changed, 3239 insertions(+), 240 deletions(-) create mode 100644 examples/memory_service_with_advanced_memory_redis/.env create mode 100644 examples/memory_service_with_advanced_memory_redis/README.md create mode 100644 examples/memory_service_with_advanced_memory_redis/agent/__init__.py create mode 100644 examples/memory_service_with_advanced_memory_redis/agent/agent.py create mode 100644 examples/memory_service_with_advanced_memory_redis/agent/tools.py create mode 100644 examples/memory_service_with_advanced_memory_redis/run_agent.py create mode 100644 examples/memory_service_with_advanced_memory_sql/.env create mode 100644 examples/memory_service_with_advanced_memory_sql/README.md create mode 100644 examples/memory_service_with_advanced_memory_sql/agent/__init__.py create mode 100644 examples/memory_service_with_advanced_memory_sql/agent/agent.py create mode 100644 examples/memory_service_with_advanced_memory_sql/agent/config.py create mode 100644 examples/memory_service_with_advanced_memory_sql/agent/prompts.py create mode 100644 examples/memory_service_with_advanced_memory_sql/agent/tools.py create mode 100644 examples/memory_service_with_advanced_memory_sql/run_agent.py create mode 100644 tests/advanced_memory/test_redis_stores.py create mode 100644 tests/advanced_memory/test_sql_stores.py create mode 100644 trpc_agent_sdk/advanced_memory/_redis_stores.py create mode 100644 trpc_agent_sdk/advanced_memory/_sql_stores.py create mode 100644 trpc_agent_sdk/advanced_memory/_storage_backend.py diff --git a/examples/memory_service_with_advanced_memory/.env b/examples/memory_service_with_advanced_memory/.env index 2da17e1ce..e4183ff5b 100644 --- a/examples/memory_service_with_advanced_memory/.env +++ b/examples/memory_service_with_advanced_memory/.env @@ -6,3 +6,7 @@ TRPC_AGENT_MODEL_NAME= # Set both model limits to enable token-based context budgeting. TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS= TRPC_AGENT_MAX_OUTPUT_TOKENS= + +# Optional TTL settings. Leave empty to disable automatic expiration. +M_TTL=120 +SESSION_TTL=60 diff --git a/examples/memory_service_with_advanced_memory/README.md b/examples/memory_service_with_advanced_memory/README.md index 0b430210f..17b534f21 100644 --- a/examples/memory_service_with_advanced_memory/README.md +++ b/examples/memory_service_with_advanced_memory/README.md @@ -2,25 +2,22 @@ ## Advanced Memory 简介 -`Advanced Memory` 是一套面向 Agent 的本地化记忆与上下文管理机制,重点增强 -Agent 在长期信息沉淀和超长对话处理方面的能力: - -- **本地化持久存储**:记忆和上下文数据以本地文件形式持久化,存储位置、数据边界 - 和组织方式清晰可控,适合本地开发、调试、迁移和审计。 -- **更强的长期记忆能力**:支持将对话中的稳定事实、用户偏好和重要经验主动沉淀为 - 可组织、可更新、可跨 Session 使用的长期记忆,而不是简单堆积历史消息。 -- **分层记忆管理**:分别管理原始对话、Session 级记忆和跨 Session 长期记忆,让不同 - 类型的信息以合适的粒度参与后续推理。 -- **上下文管理**:根据上下文规模、信息类型和使用情况,对历史消息、工具结果及记忆 - 内容进行统一治理,在保留关键信息的同时控制模型输入规模。 -- **上下文压缩**:支持对历史上下文和工具结果进行渐进式裁剪、压缩和摘要,降低长 - 对话导致的上下文膨胀以及超出模型窗口限制的风险。 -- **结构化记忆提取**:从持续增长的对话中提取结构化信息,形成更稳定、更易维护的 - Session Memory,提升后续对话对历史信息的利用效率。 - -本示例演示如何使用 `AdvancedMemorySessionService`。它把 Session 持久化和 -Advanced Memory 上下文管理整合到一个 SessionService 中,用户不需要显式调用 -`setup_advanced_memory()`,也不需要再创建 `InMemorySessionService`。 +`Advanced Memory` 是一套面向 Agent 的本地化记忆与上下文管理机制,重点增强 Agent 在长期信息沉淀和超长对话处理方面的能力: + +- **更强的长期记忆能力**:支持将对话中的稳定事实、用户偏好和重要经验主动沉淀为可组织、可更新、可跨 Session 使用的长期记忆,而不是简单堆积历史消息。 +- **分层记忆管理**:分别管理原始对话、Session 级记忆和跨 Session 长期记忆,让不同类型的信息以合适的粒度参与后续推理。 +- **上下文管理**:根据上下文规模、信息类型和使用情况,对历史消息、工具结果及记忆内容进行统一治理,在保留关键信息的同时控制模型输入规模。 +- **上下文压缩**:支持对历史上下文和工具结果进行渐进式裁剪、压缩和摘要,降低长对话导致的上下文膨胀以及超出模型窗口限制的风险。 +- **结构化记忆提取**:从持续增长的对话中提取结构化信息,形成更稳定、更易维护的Session Memory,提升后续对话对历史信息的利用效率。 +- **本地化持久存储**:记忆和上下文数据以本地文件形式持久化,存储位置、数据边界和组织方式清晰可控,适合本地开发、调试、迁移和审计。 + +本示例演示如何使用 `AdvancedMemorySessionService`。它把 Session 持久化和Advanced Memory 上下文管理整合到一个 SessionService 中,用户不需要显式调用`setup_advanced_memory()`,也不需要再创建 `InMemorySessionService`。 + +**Advanced Memory 在 Redis 存储:** +[Redis `run_agent.py`](../memory_service_with_advanced_memory_redis/run_agent.py) + +**Advanced Memory 在 SQL 存储:** +[SQL `run_agent.py`](../memory_service_with_advanced_memory_sql/run_agent.py) ## 示例流程 @@ -44,6 +41,12 @@ from trpc_agent_sdk.runners import Runner session_service = AdvancedMemorySessionService( config=AdvancedMemoryConfig( root_dir=Path(__file__).resolve().parent, + memory_ttl_seconds=120, + session_ttl_seconds=60, + memory_focus_instruction=( + "特别关注并主动记住用户长期稳定的兴趣爱好、" + "编程语言偏好、开发习惯和测试习惯。" + ), ) ) @@ -51,6 +54,7 @@ runner = Runner( app_name="advanced_memory_demo", agent=agent, session_service=session_service, + defer_post_turn_processing=True, # True 时开启,后台线程异步执行子 Agent 摘要 ) ``` @@ -67,6 +71,21 @@ runner = Runner( `AdvancedMemoryConfig` 默认已经启用这些能力,本示例直接使用默认配置。 +## 不同存储后端的 SessionService 选择 + +`AdvancedMemorySessionService` 是本地文件版 SessionService。使用 Redis 或 SQL 时,不要继续使用它,否则可能形成 Session 数据与 Advanced Memory 数据分开存储的混合模式。 + +推荐组合: + +- local:`AdvancedMemorySessionService` +- Redis:`RedisSessionService` + `AdvancedMemoryService` +- SQL:`SqlSessionService` + `AdvancedMemoryService` + +Redis 和 SQL 的完整示例分别见: + +- [Advanced Memory Redis 示例](../memory_service_with_advanced_memory_redis/README.md) +- [Advanced Memory SQL 示例](../memory_service_with_advanced_memory_sql/README.md) + ## 数据目录 运行后,数据默认写入当前示例目录: @@ -112,18 +131,22 @@ python3 run_agent.py - `TRPC_AGENT_MODEL_NAME` - `TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS`(可选,模型总上下文窗口大小,单位为 token) - `TRPC_AGENT_MAX_OUTPUT_TOKENS`(可选,模型最大输出窗口大小,单位为 token) +- `M_TTL`(可选,长期 memory 过期时间,单位为秒) +- `SESSION_TTL`(可选,session 相关数据过期时间,单位为秒) + +`M_TTL` 和 `SESSION_TTL` 未配置时不会自动删除数据。Session 的后台清理检查间隔由示例内部设置,不需要单独配置。 + +本示例提供的 `.env` 默认使用 `M_TTL=120` 和 `SESSION_TTL=60`,方便直接观察 +过期清理;如果不希望自动删除,将这两个值留空即可。 -`.env` 中留空的变量不会覆盖默认值;如果同时在 Python 中传入 -`model_context_window_tokens` 或 `max_output_tokens`,Python 显式配置优先。 +`.env` 中留空的变量不会覆盖默认值;如果同时在 Python 中传入`model_context_window_tokens` 或 `max_output_tokens`,Python 显式配置优先。 -如果配置了模型上下文窗口,Advanced Memory 会用 -`TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS - TRPC_AGENT_MAX_OUTPUT_TOKENS` +如果配置了模型上下文窗口,Advanced Memory 会用`TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS - TRPC_AGENT_MAX_OUTPUT_TOKENS` 作为可用于输入内容的窗口;两个变量都留空时使用字符数阈值。 ## `AdvancedMemoryConfig` 配置项 -下面列出当前所有可直接传入 `AdvancedMemoryConfig` 的配置项。**没有特殊需求时, -只设置 `root_dir` 即可**;示例中的值均为默认值。 +下面列出当前所有可直接传入 `AdvancedMemoryConfig` 的配置项。**没有特殊需求时,只设置 `root_dir` 即可**;其中 TTL 和记忆重点使用本示例的演示值。 ```python session_service = AdvancedMemorySessionService( @@ -139,10 +162,18 @@ session_service = AdvancedMemorySessionService( encoding="utf-8", # 文件编码 transcript_fsync=False, # transcript 写入后是否 fsync + # TTL(单位:秒;None 表示不过期) + memory_ttl_seconds=120, # 长期记忆 TTL(秒) + session_ttl_seconds=60, # 会话记忆 TTL(秒) + # 长期记忆 memory_index_max_lines=200, # 注入 prompt 的索引最大行数 memory_index_max_bytes=25_000, # 注入 prompt 的索引最大字节数 long_term_memory_injection_enabled=True, # 是否注入 MEMORY.md + memory_focus_instruction=( # 可选:重点记忆要求 + "特别关注并主动记住用户长期稳定的兴趣爱好、" + "编程语言偏好、开发习惯和测试习惯。" + ), # 工具结果 tool_result_max_chars=50_000, # 单个工具结果最大字符数 @@ -212,6 +243,26 @@ session_service = AdvancedMemorySessionService( ) ``` +`memory_focus_instruction` 可以传入应用级的自定义记忆偏好,例如: + +```python +memory_focus_instruction="特别关注用户长期稳定的兴趣爱好和开发习惯。" +``` + +它会追加到长期记忆的 system instruction 中,提示模型优先关注这些内容。 + +本示例还会把同一个 `SESSION_TTL` 传给 `SessionServiceConfig`,用于清理`session.json` 和 Session 目录;`cleanup_interval_seconds=5` 只是内部检查频率,不是另一个需要用户配置的 TTL: + +```python +session_config = SessionServiceConfig( + ttl=SessionServiceConfig.create_ttl_config( + enable=True, + ttl_seconds=60, # SESSION_TTL + cleanup_interval_seconds=5, # 内部检查频率 + ) +) +``` + `preload_memory_model` 不是 `AdvancedMemoryConfig` 字段,而是 `AdvancedMemorySessionService` 的可选参数,用于指定轻量筛选模型: diff --git a/examples/memory_service_with_advanced_memory/run_agent.py b/examples/memory_service_with_advanced_memory/run_agent.py index 097a3271d..0f2d564dd 100644 --- a/examples/memory_service_with_advanced_memory/run_agent.py +++ b/examples/memory_service_with_advanced_memory/run_agent.py @@ -8,6 +8,7 @@ """Run the two-session Advanced Memory demonstration.""" import asyncio +import os from pathlib import Path from dotenv import load_dotenv @@ -19,17 +20,27 @@ from agent.agent import create_agent -load_dotenv() +load_dotenv(Path(__file__).with_name(".env")) def create_session_service() -> AdvancedMemorySessionService: """Create the persistent Advanced Memory session service.""" + memory_ttl = os.getenv("M_TTL") + session_ttl = os.getenv("SESSION_TTL") + session_ttl_seconds = int(session_ttl) if session_ttl else 0 return AdvancedMemorySessionService( - config=AdvancedMemoryConfig(root_dir=Path(__file__).resolve().parent), + config=AdvancedMemoryConfig( + root_dir=Path(__file__).resolve().parent, + memory_ttl_seconds=int(memory_ttl) if memory_ttl else None, + session_ttl_seconds=session_ttl_seconds or None, + memory_focus_instruction=("特别关注并主动记住用户长期稳定的兴趣爱好、" + "编程语言偏好、开发习惯和测试习惯。"), + ), session_config=SessionServiceConfig(ttl=SessionServiceConfig.create_ttl_config( - ttl_seconds=60, + enable=bool(session_ttl), + ttl_seconds=session_ttl_seconds, cleanup_interval_seconds=5, - )), + ), ), ) @@ -64,6 +75,10 @@ async def main() -> None: agent=agent, session_service=session_service, ) + memory_ttl = os.getenv("M_TTL") + memory_ttl_seconds = int(memory_ttl) if memory_ttl else 0 + session_ttl = os.getenv("SESSION_TTL") + session_ttl_seconds = int(session_ttl) if session_ttl else 0 try: session_one_prompts = [ ("Please remember that my favorite programming language is Python. " @@ -95,9 +110,11 @@ 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.") + wait_seconds = max(memory_ttl_seconds, session_ttl_seconds) + if wait_seconds: + print(f"\n⏳ Waiting for TTL cleanup ({wait_seconds + 5}s)...") + await asyncio.sleep(wait_seconds + 5) + print("🧹 Expired Advanced Memory data 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..ed021e751 --- /dev/null +++ b/examples/memory_service_with_advanced_memory_redis/.env @@ -0,0 +1,11 @@ +REDIS_URL=redis://localhost:6379/0 + +# 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= \ 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..c2181ce46 --- /dev/null +++ b/examples/memory_service_with_advanced_memory_redis/README.md @@ -0,0 +1,428 @@ +# Advanced Memory Redis 示例 + +本示例演示如何将 Advanced Memory 的本地文件存储切换为 Redis,并验证: + +- Redis:`RedisSessionService` + `AdvancedMemoryService` +- 长期 memory 可以跨 Python 进程持久化; +- 同一用户在不同 `session_id` 中可以读取自己的长期 memory; +- session 相关数据和长期 memory 可以分别设置 TTL; +- Redis 中的 Markdown、Stream 和索引数据如何组织。 + +示例使用两个服务: + +```text +RedisSessionService +└── 保存 Session、app state、user state + +AdvancedMemoryService(storage_backend="redis") +└── 保存长期 memory、session memory、transcript、tool result +``` + +## 环境要求 + +- Python 3.10+,推荐 Python 3.12; +- 可访问的 Redis 服务; +- 可正常调用的模型服务。 + +如果还没有 Redis,可以使用 Docker: + +```bash +docker run --name advanced-memory-redis \ + -p 6379:6379 \ + -d redis:7-alpine +``` + +容器已创建过时不要重复执行 `docker run`,直接启动: + +```bash +docker start advanced-memory-redis +``` + +检查 Redis: + +```bash +docker exec advanced-memory-redis redis-cli PING +# PONG +``` + +## Redis 配置方式 + +### 方式一:使用完整连接串 + +在当前目录的 `.env` 中配置: + +```dotenv +REDIS_URL=redis://localhost:6379/0 +``` + +带密码: + +```dotenv +REDIS_URL=redis://:password@redis.example.com:6379/0 +``` + +Redis ACL 用户名和密码: + +```dotenv +REDIS_URL=redis://username:password@redis.example.com:6379/0 +``` + +启用 TLS: + +```dotenv +REDIS_URL=rediss://:password@redis.example.com:6380/0 +``` + +密码包含 `@`、`:`、`/`、`#` 等特殊字符时,需要进行 URL 编码。 + +### 方式二:分别配置连接参数 + +也可以不设置 `REDIS_URL`,改为: + +```dotenv +REDIS_HOST=127.0.0.1 +REDIS_PORT=6379 +REDIS_DB=0 +REDIS_USER= +REDIS_PASSWORD= +REDIS_TLS=false +``` + +云 Redis 使用示例: + +```dotenv +REDIS_HOST=your-redis.example.com +REDIS_PORT=6379 +REDIS_DB=0 +REDIS_USER=your-user +REDIS_PASSWORD=your-password +REDIS_TLS=true +``` + +代码会优先使用 `REDIS_URL`;未设置时才根据上述字段构造连接串。 + +## 模型和 TTL 配置 + +`.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_URL=redis://localhost:6379/0 + +# 长期 memory 的 TTL,单位为秒 +M_TTL=120 + +# 所有 session 相关内容的 TTL,单位为秒 +SESSION_TTL=60 +``` + +TTL 规则: + +- `M_TTL` 管理用户级长期 memory 的全部 Redis key; +- `SESSION_TTL` 管理 session memory、transcript、tool result、去重 key; +- `SESSION_TTL` 也传给 `RedisSessionService`,用于 Session 和 state; +- TTL 会在访问或写入时刷新,是“最后一次活动后过期”; +- 两个 TTL 必须设置为大于 0 的整数。 + +更多 Advanced Memory 配置请参考[Advanced Memory README](../memory_service_with_advanced_memory/README.md)。 + +## 运行示例 + +```bash +cd examples/memory_service_with_advanced_memory_redis +source ../../.venv/bin/activate +python run_agent.py +``` + +脚本会自动启动两个独立的 Python 子进程: + +```text +RUNNER A PROCESS +├── 使用 7 条对话模拟记忆建立过程 +└── Alice 的姓名和 favorite color 会被保存到长期 memory + +RUNNER B PROCESS +├── 使用新的 session +├── 询问 Alice 的 name +└── 询问 Alice 的 favorite color +``` + +两个进程使用相同的: + +```text +app_name = advanced-memory-redis-demo +user_id = redis-demo-user +``` + +但使用不同的 `session_id`。第二个进程应该能够回答: + +```text +name: Alice +favorite color: blue +``` + +这证明了 Redis 数据可以跨进程、跨 session 持久化。 + +也可以单独运行某个阶段: + +```bash +python run_agent.py --phase write # Runner A +python run_agent.py --phase read # Runner B +``` + +## 最基本的构建方式 + +Redis 版本最核心的构建过程可以简化为三步: + +```python +redis_url = "redis://:password@localhost:6379/0" + +memory_service = AdvancedMemoryService( + AdvancedMemoryConfig( + storage_backend="redis", + redis_url=redis_url, + memory_ttl_seconds=120, # from M_TTL; omit to disable expiration + session_ttl_seconds=60, # from SESSION_TTL; omit to disable expiration + ) +) + +session_config = SessionServiceConfig( + ttl=SessionServiceConfig.create_ttl_config( + enable=True, + ttl_seconds=60, # same value as SESSION_TTL + cleanup_interval_seconds=60, + ) +) +session_service = RedisSessionService( + db_url=redis_url, + is_async=True, + session_config=session_config, +) + +runner = Runner( + app_name="advanced-memory-redis-demo", + agent=create_agent(), + session_service=session_service, + memory_service=memory_service, +) +``` + +其中: + +- 用户只需要配置 `M_TTL` 和 `SESSION_TTL` 两个 TTL; +- `AdvancedMemoryService` 负责长期 memory、session memory、transcript 和 tool result; +- `RedisSessionService` 负责框架 Session、app state 和 user state; +- `Runner` 将 Agent、Session Service 和 Memory Service 组合起来; +- 运行请求时通过 `user_id` 和 `session_id` 指定用户及会话。 + +## 运行结果(实测) + +```text + user: Do you remember my name? +🔧 tool call: list_memory_index({}) +📊 Tool Result: {'index_path': '/data/workspace/trpc-agent-python-am-service/examples/memory_service_with_advanced_memory_redis/tenants/advanced-memory-redis-demo/redis-demo-user/MEMORY/MEMORY.md', 'index': ''} +🤖 Assistant: I checked my long-term memory, but I'm afraid I don't have anything saved yet — the memory index is currently empty, so I don't know your name. + +If you'd like, just tell me your name (and anything else you'd like me to remember about you), and I'll save it so I can recall it in future conversations! + +----- Runner A, query 2 ----- + +📝 user: Do you remember my favorite color? +🔧 tool call: list_memory_index({}) +📊 Tool Result: {'index_path': '/data/workspace/trpc-agent-python-am-service/examples/memory_service_with_advanced_memory_redis/tenants/advanced-memory-redis-demo/redis-demo-user/MEMORY/MEMORY.md', 'index': ''} +🤖 Assistant: I checked my long-term memory, but I don't have anything saved about your favorite color yet — my memory index is currently empty. + +If you'd like, tell me your favorite color and I'll remember it for future conversations. 💬 + +----- Runner A, query 3 ----- + +📝 user: what is the weather like in paris? +🔧 tool call: get_weather_report({'city': 'Paris'}) +📊 Tool Result: {'status': 'success', 'report': 'The weather in Paris is sunny with a temperature of 25 degrees Celsius.'} +🤖 Assistant: The weather in Paris is **sunny** with a temperature of **25°C**. ☀️ + +...... + +==================== READ PROCESS ==================== + +----- Runner B, query 1 ----- + +📝 user: Do you remember my name? +🔧 tool call: read_memory({'filename': 'alice-identity.md'}) +📊 Tool Result: {'found': True, 'filename': 'alice-identity.md', 'content': "---\nname: Alice's identity\ndescription: User's name and personal preferences.\ntype: user\nupdated_at: 2026-09-08T07:17:40.559641+00:00\n---\n- Name: Alice\n- Favorite color: blue\n", 'updated_at': '2026-09-08T07:17:40.559641+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 do — your name is Alice! 😊 And I also remember that your favorite color is blue. + +----- Runner B, query 2 ----- + +📝 user: Do you remember my favorite color? +🔧 tool call: read_memory({'filename': 'alice-identity.md'}) +📊 Tool Result: {'found': True, 'filename': 'alice-identity.md', 'content': "---\nname: Alice's identity\ndescription: User's name and personal preferences.\ntype: user\nupdated_at: 2026-09-08T07:17:40.559641+00:00\n---\n- Name: Alice\n- Favorite color: blue\n", 'updated_at': '2026-09-08T07:17:40.559641+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! According to your memory profile, your favorite color is **blue**. 💙 +``` + +## 查看 Redis 中的数据 + +进入 Redis CLI: + +```bash +docker exec -it advanced-memory-redis redis-cli +``` + +查看本示例写入的全部 Redis key: + +```redis +SCAN 0 MATCH advanced-memory-redis-demo:v1:* COUNT 100 +``` + +也可以在命令行中直接查看全部 key: + +```bash +docker exec advanced-memory-redis redis-cli --scan \ + --pattern 'advanced-memory-redis-demo:v1:*' +``` + +`SCAN` 不会像 `KEYS *` 一样阻塞 Redis,适合共享或云 Redis 环境。 + +## 查看 TTL + +长期 memory: + +```redis +TTL "advanced-memory-redis-demo:v1:{advanced-memory-redis-demo:redis-demo-user}:memory:index" +TTL "advanced-memory-redis-demo:v1:{advanced-memory-redis-demo:redis-demo-user}:memory:topic:user_favorite_project_code.md" +``` + +预期接近 `120`。 + +session transcript: + +```redis +TTL "advanced-memory-redis-demo:v1:{advanced-memory-redis-demo:redis-demo-user:redis-write-session}:transcript" +``` + +预期接近 `60`。 + +TTL 含义: + +```text +-1 永不过期 +-2 key 不存在或已经过期 +大于 0 剩余秒数 +``` + +观察 session key: + +```bash +docker exec advanced-memory-redis redis-cli --scan \ + --pattern 'advanced-memory-redis-demo:v1:*:summary' + +docker exec advanced-memory-redis redis-cli --scan \ + --pattern 'advanced-memory-redis-demo:v1:*:transcript*' +``` + +## 清理测试数据 + +只删除本示例的 Advanced Memory key: + +```bash +docker exec advanced-memory-redis redis-cli --scan \ + --pattern 'advanced-memory-redis-demo:v1:*' \ + | xargs -r docker exec -i advanced-memory-redis redis-cli DEL +``` + +测试 Redis 独占一个数据库时,也可以清空当前数据库: + +```bash +docker exec -it advanced-memory-redis redis-cli FLUSHDB +``` + +`FLUSHDB` 会删除当前 Redis DB 中的所有数据,不要在共享或生产数据库执行。 + +## Redis 中的存储形式 + +### 长期 memory + +本地文件概念: + +```text +MEMORY/MEMORY.md +MEMORY/user_favorite_project_code.md +``` + +Redis 映射: + +```text +{prefix}:{app:user}:memory:index +{prefix}:{app:user}:memory:topic:user_favorite_project_code.md +``` + +类型都是 Redis String,内容是 Markdown。 + +topic 列表的辅助索引: + +```text +{prefix}:{app:user}:memory:topics +``` + +类型是 ZSet,member 是 topic 文件名,score 是更新时间。 + +memory TTL registry: + +```text +{prefix}:{app:user}:memory:keys +``` + +它记录该用户的所有长期 memory key,用于统一刷新 `M_TTL`。 + +### session memory + +本地文件概念: + +```text +SESSION/{session_id}/session_memory.md +``` + +Redis 映射: + +```text +{prefix}:{app:user:session}:summary +``` + +类型是 Redis String,内容是 Markdown。 + +### transcript + +本地文件概念: + +```text +SESSION/{session_id}/transcript.jsonl +``` + +Redis 映射: + +```text +{prefix}:{app:user:session}:transcript +``` + +类型是 Redis Stream,每条记录保存一份 JSON 数据。 + +### transcript 去重和 tool result + +```text +{prefix}:{app:user:session}:transcript:seen:{unique_key} +{prefix}:{app:user:session}:tool:{result_id} +``` + +去重 key 使用 Set,tool result 使用 String。 + +session TTL registry: + +```text +{prefix}:{app:user:session}:keys +``` + +它记录该 session 下的 summary、transcript、tool result 等 key,用于统一刷新 +`SESSION_TTL`,避免同一个 session 的不同内容出现 TTL 不一致。 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..5aafab3db --- /dev/null +++ b/examples/memory_service_with_advanced_memory_redis/run_agent.py @@ -0,0 +1,150 @@ +#!/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.advanced_memory import AdvancedMemoryConfig +from trpc_agent_sdk.memory import AdvancedMemoryService +from trpc_agent_sdk.runners import Runner +from trpc_agent_sdk.sessions import RedisSessionService, SessionServiceConfig +from trpc_agent_sdk.types import Content, Part + +load_dotenv(Path(__file__).with_name(".env")) + +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 Advanced Memory backed by the configured Redis instance.""" + memory_ttl = os.getenv("M_TTL") + session_ttl = os.getenv("SESSION_TTL") + config = AdvancedMemoryConfig( + storage_backend="redis", + redis_url=redis_url, + redis_key_prefix="advanced-memory-redis-demo:v1", + memory_ttl_seconds=int(memory_ttl) if memory_ttl else None, + session_ttl_seconds=int(session_ttl) if session_ttl else None, + ) + return AdvancedMemoryService(config) + + +def create_redis_session_service(redis_url: str) -> RedisSessionService: + """Create session storage with the Advanced Memory session TTL.""" + session_ttl = os.getenv("SESSION_TTL") + ttl_seconds = int(session_ttl) if session_ttl else 0 + return RedisSessionService( + db_url=redis_url, + is_async=True, + session_config=SessionServiceConfig(ttl=SessionServiceConfig.create_ttl_config( + enable=bool(session_ttl), + ttl_seconds=ttl_seconds, + cleanup_interval_seconds=ttl_seconds, + ), ), + ) + + +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() + memory_service = create_advanced_memory_service(redis_url) + session_service = create_redis_session_service(redis_url) + runner = Runner( + app_name=app_name, + agent=create_agent(), + session_service=session_service, + memory_service=memory_service, + ) + 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..a617a519c --- /dev/null +++ b/examples/memory_service_with_advanced_memory_sql/.env @@ -0,0 +1,17 @@ +# Model configuration +TRPC_AGENT_API_KEY= +TRPC_AGENT_BASE_URL= +TRPC_AGENT_MODEL_NAME= + +TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS= +TRPC_AGENT_MAX_OUTPUT_TOKENS= + +# 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 +M_TTL=120 +SESSION_TTL=60 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..18540b19f --- /dev/null +++ b/examples/memory_service_with_advanced_memory_sql/README.md @@ -0,0 +1,194 @@ +# Advanced Memory SQL 示例 + +本示例使用 SQL 保存 Advanced Memory,并验证同一用户的长期 memory 可以跨 Python 进程和不同 session 读取。 + +- SQL:`SqlSessionService` + `AdvancedMemoryService` + +```text +SqlSessionService +└── Session、app state、user state + +AdvancedMemoryService(storage_backend="sql") +└── 长期 memory、session memory、transcript、tool result +``` + +## 配置 + +默认使用 SQLite,运行示例不需要额外启动数据库: + +```dotenv +SQL_URL=sqlite:///advanced-memory-sql-demo.db +SQL_IS_ASYNC=false +``` + +使用 MySQL 时: + +```dotenv +SQL_URL=mysql+aiomysql://user:password@host:3306/trpc_agent_advanced_memory?charset=utf8mb4 +SQL_IS_ASYNC=true +``` + +也可以通过 `MYSQL_USER`、`MYSQL_PASSWORD`、`MYSQL_HOST`、`MYSQL_PORT` 和 +`MYSQL_DB` 构造 MySQL URL。模型配置需要设置: + +```dotenv +TRPC_AGENT_API_KEY=your-api-key +TRPC_AGENT_BASE_URL=your-base-url +TRPC_AGENT_MODEL_NAME=your-model-name +``` + +`M_TTL` 默认控制长期 memory 的过期时间,`SESSION_TTL` 控制 session 相关内容的过期时间, +单位都是秒。 + +更多 Advanced Memory 配置请参考[Advanced Memory README](../memory_service_with_advanced_memory/README.md)。 + +## 运行 + +```bash +source .venv/bin/activate +cd examples/memory_service_with_advanced_memory_sql +python run_agent.py +``` + +脚本会依次启动两个独立进程: + +```text +RUNNER A PROCESS +├── 使用 7 条对话模拟记忆建立过程 +└── Alice 的姓名和 favorite color 会被保存到长期 memory + +RUNNER B PROCESS +├── 使用新的 session +├── 询问 Alice 的 name +└── 询问 Alice 的 favorite color +``` + +Runner B 应该能够回答: + +```text +name: Alice +favorite color: blue +``` + +也可以单独运行: + +```bash +python run_agent.py --phase write # Runner A +python run_agent.py --phase read # Runner B +``` + +第一次运行后,SQLite 文件 `advanced-memory-sql-demo.db` 会自动创建, +Advanced Memory 的表也会自动创建。 + +## 最基本的构建方式 + +SQL 版本最核心的构建过程可以简化为三步: + +```python +sql_url = "mysql+aiomysql://user:password@localhost:3306/trpc_agent_advanced_memory" + +memory_service = AdvancedMemoryService( + AdvancedMemoryConfig( + storage_backend="sql", + sql_url=sql_url, + sql_is_async=True, + memory_ttl_seconds=120, # from M_TTL; omit to disable expiration + session_ttl_seconds=60, # from SESSION_TTL; omit to disable expiration + ) +) + +session_config = SessionServiceConfig( + ttl=SessionServiceConfig.create_ttl_config( + enable=True, + ttl_seconds=60, # same value as SESSION_TTL + cleanup_interval_seconds=60, + ) +) +session_service = SqlSessionService( + db_url=sql_url, + is_async=True, + session_config=session_config, +) + +runner = Runner( + app_name="advanced-memory-sql-demo", + agent=create_agent(), + session_service=session_service, + memory_service=memory_service, +) +``` + +其中: + +- 用户只需要配置 `M_TTL` 和 `SESSION_TTL` 两个 TTL; +- `AdvancedMemoryService` 负责长期 memory、session memory、transcript 和 tool result; +- `SqlSessionService` 负责框架 Session、app state 和 user state; +- `Runner` 将 Agent、Session Service 和 Memory Service 组合起来; +- 运行请求时通过 `user_id` 和 `session_id` 指定用户及会话; +- 多个节点只要使用相同的 SQL 数据库、`app_name` 和 `user_id`,就能访问同一份长期 memory。 + +## 运行结果(实测) + +```text + +==================== WRITE PROCESS ==================== + +----- Runner A, query 1 ----- +📝 user: Do you remember my name? +🔧 tool call: list_memory_index({}) +📊 Tool Result: {'index_path': '/data/workspace/trpc-agent-python-am-service/examples/memory_service_with_advanced_memory_sql/tenants/advanced-memory-sql-demo/sql-demo-user/MEMORY/MEMORY.md', 'index': ''} +🤖 Assistant: I checked my long-term memory index, and it's currently empty — I don't have any saved memories about you yet, so I don't remember your name. + +If you'd like, tell me your name (or anything else you'd like me to remember about you), and I'll save it to my memory so I can remember it across future conversations. + +----- Runner A, query 2 ----- +📝 user: Do you remember my favorite color? +🔧 tool call: list_memory_index({}) +📊 Tool Result: {'index_path': '/data/workspace/trpc-agent-python-am-service/examples/memory_service_with_advanced_memory_sql/tenants/advanced-memory-sql-demo/sql-demo-user/MEMORY/MEMORY.md', 'index': ''} +🤖 Assistant: I checked my long-term memory index, and it's currently empty — I don't have any saved memories yet, so I don't remember your favorite color. + +If you tell me your favorite color (or any other preferences you'd like me to keep in mind), I can save it for future conversations! + +----- Runner A, query 3 ----- +📝 user: what is the weather like in paris? +🔧 tool call: get_weather_report({'city': 'Paris'}) +📊 Tool Result: {'status': 'success', 'report': 'The weather in Paris is sunny with a temperature of 25 degrees Celsius.'} +🤖 Assistant: The weather in Paris is **sunny** with a temperature of **25°C**. ☀️ + +...... + +==================== READ PROCESS ==================== + +----- Runner B, query 1 ----- +📝 user: Do you remember my name? +🔧 tool call: list_memory_index({}) +🔧 tool call: read_memory({'filename': 'user_identity.md'}) +📊 Tool Result: {'index_path': '/data/workspace/trpc-agent-python-am-service/examples/memory_service_with_advanced_memory_sql/tenants/advanced-memory-sql-demo/sql-demo-user/MEMORY/MEMORY.md', 'index': '- [User identity](user_identity.md):Alice, favorite color blue.\n'} +📊 Tool Result: {'found': True, 'filename': 'user_identity.md', 'content': "---\nname: User identity\ndescription: The user's name and basic identity.\ntype: user\nupdated_at: 2026-09-08T07:26:18.283583+00:00\n---\n---\nname: User identity\ndescription: The user's name and basic identity.\ntype: user\nupdated_at: 2026-09-08T07:26:06.708029+00:00\n---\nThe user's name is Alice. She introduced herself in conversation. Her favorite color is blue.\n", 'updated_at': '2026-09-08T07:26:18.283583+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** — from my long-term memory, that's how you introduced yourself. I also have noted that your favorite color is blue. 😊 + +If any of that has changed, just let me know and I'll update my memory records. + +----- Runner B, query 2 ----- +📝 user: Do you remember my favorite color? +🔧 tool call: read_memory({'filename': 'user_identity.md'}) +📊 Tool Result: {'found': True, 'filename': 'user_identity.md', 'content': "---\nname: User identity\ndescription: The user's name and basic identity.\ntype: user\nupdated_at: 2026-09-08T07:26:18.283583+00:00\n---\n---\nname: User identity\ndescription: The user's name and basic identity.\ntype: user\nupdated_at: 2026-09-08T07:26:06.708029+00:00\n---\nThe user's name is Alice. She introduced herself in conversation. Her favorite color is blue.\n", 'updated_at': '2026-09-08T07:26:18.283583+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**, Alice. 💙 +``` + +## SQL 表 + +Advanced Memory 使用独立的表,不复用原始 `SqlMemoryService` 的 `mem_events`: + +```text +advanced_memory_indexes +advanced_memory_topics +advanced_memory_session_memory +advanced_memory_transcripts +advanced_memory_transcript_seen +advanced_memory_tool_results +``` + +Markdown 内容保存在 `TEXT` 字段;transcript 保存 JSON 字符串; +`expires_at` 用于 SQL TTL。SQL 后端在读取时过滤过期数据,并在访问或写入时刷新 +同一用户或同一 session 下相关记录的过期时间。 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..7c570be0c --- /dev/null +++ b/examples/memory_service_with_advanced_memory_sql/run_agent.py @@ -0,0 +1,139 @@ +#!/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.advanced_memory import AdvancedMemoryConfig +from trpc_agent_sdk.memory import AdvancedMemoryService +from trpc_agent_sdk.runners import Runner +from trpc_agent_sdk.sessions import SessionServiceConfig, SqlSessionService +from trpc_agent_sdk.types import Content, Part + +load_dotenv(Path(__file__).with_name(".env")) + +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 Advanced Memory backed by SQL.""" + memory_ttl = os.getenv("M_TTL") + session_ttl = os.getenv("SESSION_TTL") + config = AdvancedMemoryConfig( + storage_backend="sql", + sql_url=sql_url, + sql_is_async=sql_is_async(), + memory_ttl_seconds=int(memory_ttl) if memory_ttl else None, + session_ttl_seconds=int(session_ttl) if session_ttl else None, + ) + return AdvancedMemoryService(config) + + +def create_sql_session_service(sql_url: str) -> SqlSessionService: + """Create the SQL-backed framework session service.""" + session_ttl = os.getenv("SESSION_TTL") + ttl_seconds = int(session_ttl) if session_ttl else 0 + return SqlSessionService( + db_url=sql_url, + is_async=sql_is_async(), + session_config=SessionServiceConfig(ttl=SessionServiceConfig.create_ttl_config( + enable=bool(session_ttl), + ttl_seconds=ttl_seconds, + cleanup_interval_seconds=ttl_seconds, + ), ), + ) + + +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=create_sql_session_service(sql_url), + 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/tests/advanced_memory/test_advanced_memory_session_service.py b/tests/advanced_memory/test_advanced_memory_session_service.py index 3854dd2e1..25fbb51aa 100644 --- a/tests/advanced_memory/test_advanced_memory_session_service.py +++ b/tests/advanced_memory/test_advanced_memory_session_service.py @@ -51,7 +51,8 @@ async def test_session_service_persists_and_restores_events(tmp_path: Path) -> N }, ) 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")) + metadata = json.loads( + (first.runtime.for_session(session).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)) @@ -68,25 +69,28 @@ async def test_session_service_persists_and_restores_events(tmp_path: Path) -> N 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.""" +async def test_same_session_id_is_isolated_between_users(tmp_path: Path) -> None: + """Allow matching IDs because each user owns a separate session directory.""" service = AdvancedMemorySessionService(config=_config(tmp_path)) - await service.create_session( + first = await service.create_session( app_name="demo-app", user_id="user-a", session_id="shared-session", ) + second = await service.create_session( + app_name="demo-app", + user_id="user-b", + session_id="shared-session", + ) + await service.append_event(first, _event("event-a", "for user a")) + await service.append_event(second, _event("event-b", "for user b")) - 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") + assert (await service.get_session(app_name="demo-app", user_id="user-a", + session_id="shared-session")).events[0].id == "event-a" + assert (await service.get_session(app_name="demo-app", user_id="user-b", + session_id="shared-session")).events[0].id == "event-b" + assert service.runtime.for_session(first).paths.session_dir( + first.id) != service.runtime.for_session(second).paths.session_dir(second.id) async def test_delete_session_removes_persistent_session_data(tmp_path: Path) -> None: @@ -102,7 +106,8 @@ async def test_delete_session_removes_persistent_session_data(tmp_path: Path) -> }, ) 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")) + metadata = json.loads((service.runtime.for_session(session).paths.session_dir(session.id) / + "session.json").read_text(encoding="utf-8")) assert metadata["state"] == {} await service.delete_session( @@ -116,7 +121,7 @@ async def test_delete_session_removes_persistent_session_data(tmp_path: Path) -> user_id="demo-user", session_id=session.id, ) is None - assert not service.runtime.paths.session_dir(session.id).exists() + assert not service.runtime.for_session(session).paths.session_dir(session.id).exists() async def test_ttl_cleanup_removes_expired_persistent_sessions(tmp_path: Path) -> None: @@ -216,6 +221,6 @@ async def append() -> None: asyncio.run(create()) asyncio.run(append()) - records = asyncio.run(service.runtime.transcripts.read_all("wrapped-session")) + records = asyncio.run(service.runtime.for_scope("demo-app", "demo-user").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..112511699 100644 --- a/tests/advanced_memory/test_advanced_memory_tools.py +++ b/tests/advanced_memory/test_advanced_memory_tools.py @@ -17,7 +17,7 @@ def _runtime(tmp_path: Path) -> AdvancedMemoryRuntime: return AdvancedMemoryRuntime.create(AdvancedMemoryConfig( 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: diff --git a/tests/advanced_memory/test_autocompact.py b/tests/advanced_memory/test_autocompact.py index af4e77a33..782faacab 100644 --- a/tests/advanced_memory/test_autocompact.py +++ b/tests/advanced_memory/test_autocompact.py @@ -62,7 +62,7 @@ def _runtime( autocompact_max_failures=max_failures, autocompact_summary_input_max_chars=10_000, autocompact_summary_retries=2, - )) + )).for_scope("demo-app", "demo-user") def _request(count: int, *, text_size: int = 800) -> LlmRequest: @@ -83,6 +83,11 @@ def _ctx(session_id: str = "session-a"): return SimpleNamespace( session_id=session_id, app_name="demo-app", + session=SimpleNamespace( + app_name="demo-app", + user_id="demo-user", + id=session_id, + ), agent=SimpleNamespace(model="fake-model"), ) @@ -122,7 +127,7 @@ async def test_token_budget_triggers_autocompact_and_records_diagnostics(tmp_pat autocompact_summary_input_max_chars=10_000, model_context_window_tokens=1_100, max_output_tokens=100, - )) + )).for_scope("demo-app", "demo-user") result = await AutoCompact(runtime, FakeSummaryGenerator()).apply( _request(5), session_id="session-a", @@ -435,7 +440,7 @@ async def test_disabled_autocompact_does_not_copy_request(tmp_path: Path) -> Non assert result.compacted is False assert request.contents[0] is original_content - assert not (tmp_path / "SESSION").exists() + assert not (tmp_path / "tenants" / "demo-app" / "demo-user" / "SESSION").exists() def test_setup_orders_full_context_pipeline(tmp_path: Path) -> None: diff --git a/tests/advanced_memory/test_memory_context.py b/tests/advanced_memory/test_memory_context.py index 965be4f65..78c48a1a3 100644 --- a/tests/advanced_memory/test_memory_context.py +++ b/tests/advanced_memory/test_memory_context.py @@ -104,6 +104,24 @@ 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_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( + AdvancedMemoryConfig( + enabled=True, + root_dir=tmp_path, + memory_focus_instruction="重点记住用户长期稳定的兴趣爱好。", + )) + request = LlmRequest(model="test-model") + + applied = await LongTermMemoryContext(runtime).apply(request) + + instruction = str(request.config.system_instruction) + assert applied is True + assert "## Custom memory focus" in instruction + assert "重点记住用户长期稳定的兴趣爱好。" 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) diff --git a/tests/advanced_memory/test_preload_memory.py b/tests/advanced_memory/test_preload_memory.py index 8a854da8f..602421263 100644 --- a/tests/advanced_memory/test_preload_memory.py +++ b/tests/advanced_memory/test_preload_memory.py @@ -39,7 +39,7 @@ async def test_preloader_injects_selected_topic_with_budget(tmp_path: Path) -> N preload_memory_max_chars=200, )) await runtime.initialize() - await runtime.long_term_memory.write_topic( + await runtime.for_scope("demo-app", "demo-user").long_term_memory.write_topic( "project.md", MemoryDocument( name="Project", @@ -48,10 +48,15 @@ async def test_preloader_injects_selected_topic_with_budget(tmp_path: Path) -> N content="important project details", ), ) + ctx = SimpleNamespace(session=SimpleNamespace( + app_name="demo-app", + user_id="demo-user", + id="session-a", + )) result = await MemoryPreloader(runtime, _FakeSelector()).preload( "What is relevant?", - SimpleNamespace(), + ctx, ) assert result is not None @@ -70,7 +75,7 @@ async def test_preloader_marks_truncated_content(tmp_path: Path) -> None: preload_memory_max_chars=12, )) await runtime.initialize() - await runtime.long_term_memory.write_topic( + await runtime.for_scope("demo-app", "demo-user").long_term_memory.write_topic( "project.md", MemoryDocument( name="Project", @@ -79,10 +84,15 @@ async def test_preloader_marks_truncated_content(tmp_path: Path) -> None: content="important project details", ), ) + ctx = SimpleNamespace(session=SimpleNamespace( + app_name="demo-app", + user_id="demo-user", + id="session-a", + )) result = await MemoryPreloader(runtime, _FakeSelector()).preload( "What is relevant?", - SimpleNamespace(), + ctx, ) assert result is not None @@ -99,7 +109,7 @@ async def test_preloader_failure_is_best_effort(tmp_path: Path) -> None: 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_redis_stores.py b/tests/advanced_memory/test_redis_stores.py new file mode 100644 index 000000000..8d0e9f8a3 --- /dev/null +++ b/tests/advanced_memory/test_redis_stores.py @@ -0,0 +1,104 @@ +"""Tests for Redis Advanced Memory storage and TTL grouping.""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock +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 MemoryIndexEntry +from trpc_agent_sdk.advanced_memory import SessionMemoryDocument +from trpc_agent_sdk.advanced_memory._redis_stores import RedisLongTermMemoryStore +from trpc_agent_sdk.advanced_memory._redis_stores import RedisSessionMemoryStore + + +def _store(store_type: type, **overrides: object): + config = AdvancedMemoryConfig( + storage_backend="redis", + redis_url="redis://localhost:6379/0", + root_dir=Path("/tmp/advanced-memory-redis-tests"), + memory_ttl_seconds=120, + session_ttl_seconds=60, + **overrides, + ) + paths = AdvancedMemoryPaths(config).for_scope("app", "user") + store = store_type(config, paths, MagicMock()) + + async def command(method: str, *args: object, **kwargs: object): + if method == "set" and args and str(args[0]).endswith(":memory:lock"): + return True + return [] + + store._command = AsyncMock(side_effect=command) + return store + + +@pytest.mark.asyncio +async def test_memory_writes_refresh_all_memory_keys() -> None: + store = _store(RedisLongTermMemoryStore) + + await store.write_index([ + MemoryIndexEntry(name="Profile", filename="profile.md", summary="User profile"), + ]) + + commands = [call.args for call in store._command.await_args_list] + assert ("set", f"{store._user_base}:memory:index", "- [Profile](profile.md):User profile\n") in commands + assert ("sadd", f"{store._user_base}:memory:keys", f"{store._user_base}:memory:index") in commands + assert ("expire", f"{store._user_base}:memory:index", 120) in commands + assert ("expire", f"{store._user_base}:memory:keys", 120) in commands + + +@pytest.mark.asyncio +async def test_session_writes_refresh_all_session_keys() -> None: + store = _store(RedisSessionMemoryStore) + + await store.write("session-1", SessionMemoryDocument(session_title="Test session")) + + session_base = store._session_base("session-1") + commands = [call.args for call in store._command.await_args_list] + assert any(command[0] == "set" and command[1] == f"{session_base}:summary" for command in commands) + assert ("sadd", f"{session_base}:keys", f"{session_base}:summary") in commands + assert ("expire", f"{session_base}:summary", 60) in commands + assert ("expire", f"{session_base}:keys", 60) in commands + + +@pytest.mark.asyncio +async def test_ttl_refresh_includes_previously_tracked_keys() -> None: + store = _store(RedisSessionMemoryStore) + session_base = store._session_base("session-1") + old_key = f"{session_base}:transcript" + store._command = AsyncMock(side_effect=[ + None, # SADD + [old_key.encode()], # SMEMBERS + None, # EXPIRE old key + None, # EXPIRE current key + None, # EXPIRE registry + ]) + + await store._refresh_session_ttl("session-1", f"{session_base}:summary") + + commands = [call.args for call in store._command.await_args_list] + assert ("expire", old_key, 60) in commands + assert ("expire", f"{session_base}:summary", 60) in commands + + +@pytest.mark.asyncio +async def test_memory_write_lock_releases_with_token_check() -> None: + store = _store(RedisLongTermMemoryStore) + + async with store._memory_write_lock(): + pass + + lock_key = f"{store._user_base}:memory:lock" + lock_sets = [call for call in store._command.await_args_list if call.args[:2] == ("set", lock_key)] + releases = [call for call in store._command.await_args_list if call.args and call.args[0] == "eval"] + assert lock_sets + assert lock_sets[0].kwargs["nx"] is True + assert lock_sets[0].kwargs["ex"] == 30 + assert releases + assert releases[0].args[2] == 1 + assert releases[0].args[3] == lock_key + assert releases[0].args[4] == lock_sets[0].args[2] diff --git a/tests/advanced_memory/test_session_memory_extractor.py b/tests/advanced_memory/test_session_memory_extractor.py index 240e1f2fc..7b8e9273a 100644 --- a/tests/advanced_memory/test_session_memory_extractor.py +++ b/tests/advanced_memory/test_session_memory_extractor.py @@ -130,6 +130,11 @@ def _ctx(session): return SimpleNamespace(session=session, agent=SimpleNamespace(model="fake-model")) +def _scoped(runtime: AdvancedMemoryRuntime): + """Return the tenant runtime used by the test sessions.""" + return runtime.for_scope("demo-app", "demo-user") + + 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) @@ -143,8 +148,9 @@ async def test_first_extraction_writes_document_and_checkpoint(tmp_path: Path) - _ctx(session), ) - memory = await runtime.session_memory.read(session.id) - records = await runtime.transcripts.read_all(session.id) + scoped = _scoped(runtime) + memory = await scoped.session_memory.read(session.id) + records = await scoped.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 @@ -312,7 +318,7 @@ async def test_missing_checkpoint_recovers_only_newer_timestamped_events(tmp_pat runtime = _runtime(tmp_path) service, session = await _service_and_session(runtime) await service.append_event(session, _event("event-old", "旧内容")) - await runtime.transcripts.append( + await _scoped(runtime).transcripts.append( session.id, { "kind": "session-memory-checkpoint", @@ -384,7 +390,7 @@ async def test_empty_document_does_not_overwrite_or_advance_checkpoint(tmp_path: session_title="已有记忆", current_state="等待新事件。", ) - await runtime.session_memory.write(session.id, old_document) + await _scoped(runtime).session_memory.write(session.id, old_document) await service.append_event(session, _event("event-1", "first")) result = await SessionMemoryExtractor( @@ -396,9 +402,9 @@ async def test_empty_document_does_not_overwrite_or_advance_checkpoint(tmp_path: force=True, ) - records = await runtime.transcripts.read_all(session.id) + records = await _scoped(runtime).transcripts.read_all(session.id) assert result.reason == "extraction-failed" - assert await runtime.session_memory.read(session.id) == old_document.to_markdown() + assert await _scoped(runtime).session_memory.read(session.id) == old_document.to_markdown() assert not any(record.get("kind") == "session-memory-checkpoint" for record in records) @@ -471,7 +477,7 @@ async def test_session_service_runs_extractor_after_old_summary(tmp_path: Path) 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 + assert await _scoped(runtime).session_memory.read(session.id) is not None async def test_forked_generator_uses_isolated_runner_and_returns_memory() -> None: diff --git a/tests/advanced_memory/test_sql_stores.py b/tests/advanced_memory/test_sql_stores.py new file mode 100644 index 000000000..8b96e5a7b --- /dev/null +++ b/tests/advanced_memory/test_sql_stores.py @@ -0,0 +1,84 @@ +"""SQLite tests for the Advanced Memory SQL backend.""" + +from __future__ import annotations + +from pathlib import Path + +from trpc_agent_sdk.advanced_memory import ( + AdvancedMemoryConfig, + AdvancedMemoryRuntime, + MemoryDocument, + MemoryIndexEntry, + MemoryType, + SessionMemoryDocument, +) + + +def _runtime(tmp_path: Path) -> AdvancedMemoryRuntime: + return AdvancedMemoryRuntime.create( + AdvancedMemoryConfig( + storage_backend="sql", + sql_url=f"sqlite:///{tmp_path / 'advanced-memory.db'}", + sql_is_async=False, + memory_ttl_seconds=120, + session_ttl_seconds=60, + )) + + +async def test_sql_stores_round_trip_and_deduplicate(tmp_path: Path) -> None: + root = _runtime(tmp_path) + scoped = root.for_scope("app", "user") + await scoped.initialize() + + await scoped.long_term_memory.write_index([ + MemoryIndexEntry(name="Profile", filename="profile.md", summary="Profile"), + ]) + await scoped.long_term_memory.write_topic( + "profile", + MemoryDocument( + name="Profile", + description="Profile", + memory_type=MemoryType.USER, + content="A user profile", + ), + ) + await scoped.session_memory.write("session", SessionMemoryDocument(session_title="Test")) + await scoped.tool_results.write("session", "result", '{"ok": true}') + await scoped.transcripts.append("session", {"event_id": "one"}) + _, first = await scoped.transcripts.append_unique( + "session", + {"event_id": "two"}, + unique_key="event_id", + ) + _, second = await scoped.transcripts.append_unique( + "session", + {"event_id": "two"}, + unique_key="event_id", + ) + + assert first is True + assert second is False + assert "profile.md" in await scoped.long_term_memory.read_index() + assert await scoped.long_term_memory.read_topic("profile") + assert await scoped.session_memory.read("session") + assert await scoped.tool_results.read("session", "result") == '{"ok": true}' + assert len(await scoped.transcripts.read_all("session")) == 2 + + await root.close() + + +async def test_sql_stores_isolate_users(tmp_path: Path) -> None: + root = _runtime(tmp_path) + first = root.for_scope("app", "first") + second = root.for_scope("app", "second") + await first.initialize() + await second.initialize() + + await first.long_term_memory.write_index([ + MemoryIndexEntry(name="First", filename="first.md", summary="First"), + ]) + + assert "first.md" in await first.long_term_memory.read_index() + assert "first.md" not in await second.long_term_memory.read_index() + + await root.close() diff --git a/tests/advanced_memory/test_storage.py b/tests/advanced_memory/test_storage.py index d529e3e9c..99ac2b024 100644 --- a/tests/advanced_memory/test_storage.py +++ b/tests/advanced_memory/test_storage.py @@ -4,6 +4,7 @@ import asyncio import json +import os import threading from datetime import datetime from datetime import timezone @@ -150,6 +151,30 @@ async def test_session_memory_is_isolated_by_session_id(tmp_path: Path) -> None: assert all(f"# {section}" in first_content for section in SESSION_MEMORY_SECTIONS) +async def test_scoped_storage_isolates_users_and_allows_same_session_id(tmp_path: Path) -> None: + """Keep all Advanced Memory records inside the app and user namespace.""" + runtime = AdvancedMemoryRuntime.create(_enabled_config(tmp_path)) + first = runtime.for_scope("demo-app", "user-a") + second = runtime.for_scope("demo-app", "user-b") + await first.initialize() + await second.initialize() + + await first.long_term_memory.write_index([MemoryIndexEntry(name="A", filename="a.md", summary="A")]) + await second.long_term_memory.write_index([MemoryIndexEntry(name="B", filename="b.md", summary="B")]) + await first.session_memory.write("shared", SessionMemoryDocument(session_title="A")) + await second.session_memory.write("shared", SessionMemoryDocument(session_title="B")) + await first.transcripts.append("shared", {"kind": "event", "event_id": "a"}) + await second.transcripts.append("shared", {"kind": "event", "event_id": "b"}) + + assert "a.md" in await first.long_term_memory.read_index() + assert "b.md" not in await first.long_term_memory.read_index() + assert "b.md" in await second.long_term_memory.read_index() + assert (await first.session_memory.read("shared")) != await second.session_memory.read("shared") + assert [record["event_id"] for record in await first.transcripts.read_all("shared")] == ["a"] + assert [record["event_id"] for record in await second.transcripts.read_all("shared")] == ["b"] + assert first.paths.session_dir("shared") != second.paths.session_dir("shared") + + 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)) @@ -210,6 +235,38 @@ async def test_transcript_append_unique_uses_persisted_ids(tmp_path: Path) -> No assert len(await second_runtime.transcripts.read_all("session-a")) == 1 +async def test_transcript_unique_cache_is_reset_after_session_ttl(tmp_path: Path, ) -> None: + """Allow a reused session ID to append after local TTL expiration.""" + runtime = AdvancedMemoryRuntime.create(_enabled_config( + tmp_path, + session_ttl_seconds=1, + )) + transcript = runtime.transcripts + await transcript.append_unique( + "session-a", + { + "kind": "event", + "event_id": "event-1" + }, + unique_key="event_id", + ) + activity_path = runtime.paths.session_dir("session-a") / ".advanced-memory-activity" + os.utime(activity_path, (1.0, 1.0)) + + _, appended = await transcript.append_unique( + "session-a", + { + "kind": "event", + "event_id": "event-1" + }, + unique_key="event_id", + ) + + assert appended is True + assert len(await transcript.read_all("session-a")) == 1 + await runtime.close() + + async def test_transcript_read_waits_for_in_progress_append( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, @@ -268,6 +325,37 @@ async def test_memory_index_is_truncated_when_read_over_byte_budget(tmp_path: Pa assert await runtime.long_term_memory.read_index() == "" +async def test_local_ttl_expires_memory_and_session_groups(tmp_path: Path) -> None: + """Expire local memory groups after their last activity.""" + runtime = AdvancedMemoryRuntime.create(_enabled_config( + tmp_path, + memory_ttl_seconds=1, + session_ttl_seconds=1, + )) + scoped = runtime.for_scope("app", "user") + await scoped.initialize() + await scoped.long_term_memory.write_index([ + MemoryIndexEntry(name="Profile", filename="profile.md", summary="Profile"), + ]) + await scoped.long_term_memory.write_topic( + "profile", + MemoryDocument(name="Profile", description="Profile", memory_type=MemoryType.USER, content="data"), + ) + await scoped.session_memory.write("session", SessionMemoryDocument(session_title="Session")) + await scoped.tool_results.write("session", "result", "data") + await scoped.transcripts.append("session", {"event_id": "event"}) + + old = 1.0 + os.utime(scoped.paths.memory_index_path, (old, old)) + os.utime(scoped.paths.session_dir("session") / ".advanced-memory-activity", (old, old)) + + assert await scoped.long_term_memory.read_index() == "" + assert await scoped.long_term_memory.read_topic("profile") is None + assert await scoped.session_memory.read("session") is None + assert not scoped.paths.session_dir("session").exists() + await runtime.close() + + 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)) diff --git a/tests/advanced_memory/test_tool_result_budget.py b/tests/advanced_memory/test_tool_result_budget.py index 4df8d9036..3fbbd17c3 100644 --- a/tests/advanced_memory/test_tool_result_budget.py +++ b/tests/advanced_memory/test_tool_result_budget.py @@ -70,6 +70,28 @@ async def test_single_large_result_is_persisted_and_replaced(tmp_path: Path) -> assert "x" * 100 in persisted +async def test_sql_replacement_reports_sql_storage_path(tmp_path: Path) -> None: + """Expose the path returned by the SQL tool-result store.""" + root = AdvancedMemoryRuntime.create(AdvancedMemoryConfig( + enabled=True, + storage_backend="sql", + sql_url=f"sqlite:///{tmp_path / 'memory.db'}", + sql_is_async=False, + tool_result_max_chars=200, + tool_results_per_message_max_chars=5_000, + tool_result_preview_chars=40, + )) + runtime = root.for_scope("demo-app", "demo-user") + budget = ToolResultBudget(runtime) + request, _ = _request(("result-1", "x" * 500)) + + await budget.apply(request, session_id="session-a") + + replacement = request.contents[0].parts[0].function_response.response + assert replacement["persisted_output"]["path"].startswith("advanced-memory://sql/") + assert await runtime.tool_results.read("session-a", "result-1") is not None + + 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) diff --git a/tests/advanced_memory/test_transcript_session_service.py b/tests/advanced_memory/test_transcript_session_service.py index a23217011..47cf79322 100644 --- a/tests/advanced_memory/test_transcript_session_service.py +++ b/tests/advanced_memory/test_transcript_session_service.py @@ -42,7 +42,7 @@ async def test_append_event_writes_versioned_parent_chain(tmp_path: Path) -> Non 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) + records = await runtime.for_session(session).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" @@ -65,7 +65,7 @@ async def test_duplicate_event_id_is_not_written_twice(tmp_path: Path) -> None: await service.append_event(session, duplicate) await service.append_event(session, duplicate.model_copy(deep=True)) - records = await runtime.transcripts.read_all(session.id) + records = await runtime.for_session(session).transcripts.read_all(session.id) assert [record["event_id"] for record in records] == ["event-1"] @@ -79,7 +79,7 @@ async def test_old_duplicate_does_not_rewind_parent_chain(tmp_path: Path) -> Non 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) + records = await runtime.for_session(session).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" @@ -97,7 +97,7 @@ async def test_new_wrapper_restores_parent_from_existing_transcript(tmp_path: Pa 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) + records = await second_runtime.for_session(session).transcripts.read_all(session.id) assert records[-1]["parent_event_id"] == "event-1" @@ -124,4 +124,4 @@ async def test_partial_event_is_not_written_to_transcript(tmp_path: Path) -> Non await service.append_event(session, _event("partial-1", "chunk", partial=True)) assert session.events == [] - assert await runtime.transcripts.read_all(session.id) == [] + assert await runtime.for_session(session).transcripts.read_all(session.id) == [] diff --git a/trpc_agent_sdk/advanced_memory/__init__.py b/trpc_agent_sdk/advanced_memory/__init__.py index c346cd2d3..fe1afb8a5 100644 --- a/trpc_agent_sdk/advanced_memory/__init__.py +++ b/trpc_agent_sdk/advanced_memory/__init__.py @@ -33,12 +33,14 @@ from ._microcompact import MicrocompactResult from ._microcompact import setup_microcompact from ._paths import AdvancedMemoryPaths +from ._paths import MemoryScope 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 ._runtime import ScopedAdvancedMemoryRuntime from ._session_memory import build_session_memory_prompt from ._session_memory import ForkedSessionMemoryGenerator from ._session_memory import has_session_memory_content @@ -55,6 +57,8 @@ from ._storage import SessionMemoryStore from ._storage import ToolResultStore from ._storage import TranscriptStore +from ._storage_backend import AdvancedMemoryStorageBackend +from ._storage_backend import LocalAdvancedMemoryStorageBackend from ._tool_result_budget import setup_tool_result_budget from ._tool_result_budget import ToolResultBudget from ._tool_result_budget import ToolResultBudgetCallback @@ -71,11 +75,13 @@ "AutoCompact", "AutoCompactCallback", "AutoCompactResult", + "AdvancedMemoryStorageBackend", "AdvancedMemoryConfig", "AdvancedContextManagement", "AdvancedMemoryIntegration", "AdvancedMemoryPaths", "AdvancedMemoryRuntime", + "ScopedAdvancedMemoryRuntime", "ContextBudget", "ContextTokenEstimate", "build_session_memory_prompt", @@ -89,9 +95,11 @@ "HistorySnipResult", "HeuristicTokenEstimator", "LongTermMemoryStore", + "LocalAdvancedMemoryStorageBackend", "LongTermMemoryContext", "LongTermMemoryContextCallback", "MemoryDocument", + "MemoryScope", "MemoryIndexEntry", "MemoryType", "MemoryCandidate", diff --git a/trpc_agent_sdk/advanced_memory/_autocompact.py b/trpc_agent_sdk/advanced_memory/_autocompact.py index e7efb436a..6c80f1ac2 100644 --- a/trpc_agent_sdk/advanced_memory/_autocompact.py +++ b/trpc_agent_sdk/advanced_memory/_autocompact.py @@ -8,6 +8,7 @@ from __future__ import annotations import asyncio +import copy import hashlib import json import re @@ -234,6 +235,7 @@ def __init__( self._summary_generator = summary_generator or ForkedLegacySummaryGenerator(model) self._states: dict[str, AutoCompactState] = {} self._session_locks: dict[str, asyncio.Lock] = {} + self._scoped_processors: dict[object, "AutoCompact"] = {} @property def runtime(self) -> AdvancedMemoryRuntime: @@ -242,15 +244,17 @@ def runtime(self) -> AdvancedMemoryRuntime: 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) + 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) @@ -274,7 +278,7 @@ async def _load_state(self, session_id: str) -> AutoCompactState: elif record.get("kind") == "autocompact-failure": failures += 1 state = AutoCompactState(latest_compaction=latest, consecutive_failures=failures) - self._states[session_id] = state + self._states[state_key] = state return state def _summary_content(self, summary: str) -> Content: @@ -537,6 +541,27 @@ async def apply( session_id: str, 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", + *, + session_id: str, + ctx: "InvocationContext", + force: bool = False, ) -> AutoCompactResult: """Replay old compaction and compact again when pressure is high.""" config = self._runtime.config diff --git a/trpc_agent_sdk/advanced_memory/_config.py b/trpc_agent_sdk/advanced_memory/_config.py index 975875d3e..5be754e9e 100644 --- a/trpc_agent_sdk/advanced_memory/_config.py +++ b/trpc_agent_sdk/advanced_memory/_config.py @@ -12,6 +12,7 @@ from dataclasses import field from pathlib import Path from typing import Any +from typing import Literal DEFAULT_COMPACTABLE_TOOL_NAMES = ( "Read", @@ -101,6 +102,17 @@ class AdvancedMemoryConfig: 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 + session_ttl_seconds: int | None = None + 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" session_dir_name: str = "SESSION" memory_index_name: str = "MEMORY.md" @@ -109,6 +121,7 @@ class AdvancedMemoryConfig: 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 tool_result_max_chars: int = 50_000 tool_results_per_message_max_chars: int = 200_000 tool_result_preview_chars: int = 2_000 @@ -165,6 +178,22 @@ class AdvancedMemoryConfig: def __post_init__(self) -> None: """Validate the configuration and normalize the root directory.""" + 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.session_ttl_seconds is not None and self.session_ttl_seconds <= 0: + raise ValueError("session_ttl_seconds must be greater than zero when provided") + 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, diff --git a/trpc_agent_sdk/advanced_memory/_history_snip.py b/trpc_agent_sdk/advanced_memory/_history_snip.py index 67d8ff1a6..72f959320 100644 --- a/trpc_agent_sdk/advanced_memory/_history_snip.py +++ b/trpc_agent_sdk/advanced_memory/_history_snip.py @@ -8,6 +8,7 @@ from __future__ import annotations import asyncio +import copy import json from dataclasses import dataclass from typing import Any @@ -88,6 +89,7 @@ def __init__(self, memory_runtime: AdvancedMemoryRuntime) -> None: 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: @@ -96,15 +98,17 @@ def runtime(self) -> AdvancedMemoryRuntime: 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) + 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) @@ -122,7 +126,7 @@ async def _load_state(self, session_id: str) -> HistorySnipState: snipped_ids=snipped_ids, result_hashes=result_hashes, ) - self._states[session_id] = state + self._states[state_key] = state return state def _collect_candidates(self, request: "LlmRequest") -> list[HistorySnipCandidate]: @@ -191,12 +195,33 @@ 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"): + await self._runtime.initialize() + 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) diff --git a/trpc_agent_sdk/advanced_memory/_memory_context.py b/trpc_agent_sdk/advanced_memory/_memory_context.py index 947db2330..a3ca94b89 100644 --- a/trpc_agent_sdk/advanced_memory/_memory_context.py +++ b/trpc_agent_sdk/advanced_memory/_memory_context.py @@ -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 'Redis'}\n" + f"Index file: " + f"{runtime.paths.memory_index_path if config.storage_backend == 'local' else 'Redis 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/advanced_memory/_microcompact.py b/trpc_agent_sdk/advanced_memory/_microcompact.py index eeaabdd36..ec1ae94d5 100644 --- a/trpc_agent_sdk/advanced_memory/_microcompact.py +++ b/trpc_agent_sdk/advanced_memory/_microcompact.py @@ -8,6 +8,7 @@ from __future__ import annotations import asyncio +import copy import time from dataclasses import dataclass from typing import Any @@ -79,6 +80,7 @@ def __init__(self, memory_runtime: AdvancedMemoryRuntime) -> None: 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: @@ -87,15 +89,17 @@ def runtime(self) -> AdvancedMemoryRuntime: 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) + 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) @@ -113,7 +117,7 @@ async def _load_state(self, session_id: str) -> MicrocompactState: cleared_ids=cleared_ids, result_hashes=result_hashes, ) - self._states[session_id] = state + self._states[state_key] = state return state def _collect_candidates(self, request: "LlmRequest") -> list[MicrocompactCandidate]: @@ -177,13 +181,47 @@ async def apply( *, 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"): + await self._runtime.initialize() + 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) @@ -255,6 +293,7 @@ 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 diff --git a/trpc_agent_sdk/advanced_memory/_paths.py b/trpc_agent_sdk/advanced_memory/_paths.py index da1a41e07..4cf5eaca3 100644 --- a/trpc_agent_sdk/advanced_memory/_paths.py +++ b/trpc_agent_sdk/advanced_memory/_paths.py @@ -19,6 +19,10 @@ def _safe_component(value: str, *, field_name: str) -> str: """Convert an external identifier into a safe path component.""" + if value != value.strip() or any(character.isspace() and character not in {" "} for character in value): + raise ValueError(f"{field_name} must not contain leading/trailing or control whitespace") + if any(ord(character) < 32 or ord(character) == 127 for character in value): + raise ValueError(f"{field_name} must not contain control characters") normalized = _SAFE_COMPONENT_PATTERN.sub("_", value.strip()).strip("._") if not normalized: raise ValueError(f"{field_name} must contain at least one safe character") @@ -35,21 +39,57 @@ def _collision_safe_component(value: str, *, field_name: str) -> str: return f"{normalized}-{digest}" +@dataclass(frozen=True) +class MemoryScope: + """Identify the application and user that own Advanced Memory data.""" + + 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 a stable process-local key for locks and caches.""" + return repr((self.app_name, self.user_id)) + + @dataclass(frozen=True) class AdvancedMemoryPaths: """Build all disk paths for long-term and session memory.""" config: AdvancedMemoryConfig + scope: MemoryScope | None = None + + def for_scope(self, app_name: str, user_id: str) -> "AdvancedMemoryPaths": + """Return paths rooted in the given application's user namespace.""" + return AdvancedMemoryPaths(self.config, MemoryScope(app_name, user_id)) + + @property + def tenant_root_dir(self) -> Path: + """Return this scope's root, or the legacy root when unscoped.""" + 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 a key suitable for lock and cache partitioning.""" + return self.scope.storage_key if self.scope is not None else "legacy\0global" @property def memory_dir(self) -> Path: """Return the long-term memory directory.""" - return self.config.root_dir / self.config.memory_dir_name + return self.tenant_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 + return self.tenant_root_dir / self.config.session_dir_name @property def memory_index_path(self) -> Path: diff --git a/trpc_agent_sdk/advanced_memory/_preload_memory.py b/trpc_agent_sdk/advanced_memory/_preload_memory.py index 4f7f92eee..88c487753 100644 --- a/trpc_agent_sdk/advanced_memory/_preload_memory.py +++ b/trpc_agent_sdk/advanced_memory/_preload_memory.py @@ -228,18 +228,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 +248,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 +273,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/advanced_memory/_redis_stores.py b/trpc_agent_sdk/advanced_memory/_redis_stores.py new file mode 100644 index 000000000..3bf2bd763 --- /dev/null +++ b/trpc_agent_sdk/advanced_memory/_redis_stores.py @@ -0,0 +1,299 @@ +"""Redis implementations of the Advanced Memory storage contracts.""" + +from __future__ import annotations + +import asyncio +import json +from collections.abc import Mapping +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 AdvancedMemoryConfig +from ._formats import MemoryDocument, MemoryIndexEntry, SessionMemoryDocument +from ._paths import AdvancedMemoryPaths + +_APPEND_UNIQUE_SCRIPT = """ +if redis.call('SADD', KEYS[2], ARGV[1]) == 0 then return 0 end +redis.call('XADD', KEYS[1], '*', 'data', ARGV[2]) +return 1 +""" + +_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: AdvancedMemoryConfig, 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}}}" + self._app_base = f"{config.redis_key_prefix}:{{{app_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 _session_base(self, session_id: str) -> str: + safe_session_id = self._paths.session_dir(session_id).name + tenant = f"{self._paths.tenant_root_dir.parent.name}:{self._paths.tenant_root_dir.name}" + return f"{self._config.redis_key_prefix}:{{{tenant}:{safe_session_id}}}" + + def _session_registry(self, session_id: str) -> str: + return f"{self._session_base(session_id)}:keys" + + 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, + ) -> 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: + await self._command("expire", key, ttl) + await self._command("expire", registry, ttl) + + async def _refresh_session_ttl(self, session_id: str, *keys: str) -> None: + await self._refresh_ttl_group( + self._session_registry(session_id), + list(keys), + self._config.session_ttl_seconds, + ) + + async def _refresh_memory_ttl(self, *keys: str) -> None: + await self._refresh_ttl_group( + self._memory_registry(), + list(keys), + self._config.memory_ttl_seconds, + ) + + async def delete_session(self, session_id: str) -> None: + """Delete all Advanced Memory keys for one session.""" + session_base = self._session_base(session_id) + registry = self._session_registry(session_id) + keys: set[str] = {registry} + tracked = await self._command("smembers", registry) or [] + keys.update(value for value in (self._text(item) for item in tracked) if value) + + cursor: Any = 0 + pattern = f"{session_base}:*" + while True: + cursor, scanned = await self._command( + "scan", + cursor, + match=pattern, + count=100, + ) + keys.update(value for value in (self._text(item) for item in scanned) if value) + if int(cursor) == 0: + break + if keys: + await self._command("delete", *keys) + + @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() + lines, used_bytes = [], 0 + for line in value.splitlines(keepends=True)[:self._config.memory_index_max_lines]: + size = len(line.encode(self._config.encoding)) + if used_bytes + size > self._config.memory_index_max_bytes: + break + lines.append(line) + used_bytes += size + return "".join(lines) + + 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] + + +class RedisSessionMemoryStore(_RedisStore): + + async def read(self, session_id: str) -> str | None: + key = f"{self._session_base(session_id)}:summary" + value = await self._command("get", key) + await self._refresh_session_ttl(session_id, key) + return self._text(value) + + async def write(self, session_id: str, document: SessionMemoryDocument) -> Path: + key = f"{self._session_base(session_id)}:summary" + await self._command("set", key, document.to_markdown()) + await self._refresh_session_ttl(session_id, key) + return Path(f"advanced-memory://{key}") + + +class RedisToolResultStore(_RedisStore): + + async def write(self, session_id: str, result_id: str, serialized_result: str) -> Path: + key = f"{self._session_base(session_id)}:tool:{result_id}" + await self._command("set", key, serialized_result) + await self._refresh_session_ttl(session_id, key) + return Path(f"advanced-memory://{key}") + + async def read(self, session_id: str, result_id: str) -> str | None: + key = f"{self._session_base(session_id)}:tool:{result_id}" + value = await self._command("get", key) + await self._refresh_session_ttl(session_id, key) + return self._text(value) + + +class RedisTranscriptStore(_RedisStore): + + async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: + payload = dict(record) + payload.setdefault("recorded_at", datetime.now(timezone.utc).isoformat()) + stream = f"{self._session_base(session_id)}:transcript" + await self._command("xadd", stream, {"data": json.dumps(payload)}) + await self._refresh_session_ttl(session_id, stream) + return Path(f"advanced-memory://{stream}") + + async def append_unique(self, session_id: str, record: Mapping[str, Any], *, unique_key: str) -> tuple[Path, bool]: + payload = dict(record) + value = payload.get(unique_key) + if not isinstance(value, str) or not 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()) + stream = f"{self._session_base(session_id)}:transcript" + seen = f"{stream}:seen:{unique_key}" + async with self._storage.create_db_session() as connection: + added = await self._storage.execute_command( + connection, + RedisCommand( + method="eval", + args=(_APPEND_UNIQUE_SCRIPT, 2, stream, seen, value, json.dumps(payload, ensure_ascii=False)), + )) + await self._refresh_session_ttl(session_id, stream, seen) + return Path(f"advanced-memory://{stream}"), bool(added) + + async def read_all(self, session_id: str) -> list[dict[str, Any]]: + stream = f"{self._session_base(session_id)}:transcript" + entries = await self._command("xrange", stream, "-", "+") + await self._refresh_session_ttl(session_id, stream) + records: list[dict[str, Any]] = [] + for _, fields in entries: + value = fields.get(b"data") if isinstance(fields, dict) else None + value = value or fields.get("data") + text = self._text(value) + if text: + records.append(json.loads(text)) + return records diff --git a/trpc_agent_sdk/advanced_memory/_runtime.py b/trpc_agent_sdk/advanced_memory/_runtime.py index c26def35f..046357fcd 100644 --- a/trpc_agent_sdk/advanced_memory/_runtime.py +++ b/trpc_agent_sdk/advanced_memory/_runtime.py @@ -8,10 +8,17 @@ from __future__ import annotations from dataclasses import dataclass +from dataclasses import field +import asyncio +import shutil +import threading +from typing import Any from ._config import AdvancedMemoryConfig from ._coordination import SessionOperationCoordinator from ._paths import AdvancedMemoryPaths +from ._paths import MemoryScope +from ._storage import LocalAdvancedMemoryCleanup from ._storage import LongTermMemoryStore from ._storage import SessionMemoryStore from ._storage import ToolResultStore @@ -29,12 +36,46 @@ class AdvancedMemoryRuntime: session_memory: SessionMemoryStore tool_results: ToolResultStore transcripts: TranscriptStore + _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) @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) + 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 + sql_storage = SqlStorage( + is_async=resolved_config.sql_is_async, + db_url=resolved_config.sql_url, + metadata=AdvancedMemorySqlBase.metadata, + expire_on_commit=False, + ) + from ._sql_stores import SqlAdvancedMemoryCleanup + sql_cleanup = SqlAdvancedMemoryCleanup(resolved_config, sql_storage) + else: + local_cleanup = LocalAdvancedMemoryCleanup(resolved_config) return cls( config=resolved_config, paths=paths, @@ -43,11 +84,165 @@ def create(cls, config: AdvancedMemoryConfig | None = None) -> "AdvancedMemoryRu session_memory=SessionMemoryStore(resolved_config, paths), tool_results=ToolResultStore(resolved_config, paths), transcripts=TranscriptStore(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 + from ._redis_stores import RedisSessionMemoryStore + from ._redis_stores import RedisToolResultStore + from ._redis_stores import RedisTranscriptStore + + 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) + session_memory = RedisSessionMemoryStore(self.config, paths, storage) + tool_results = RedisToolResultStore(self.config, paths, storage) + transcripts = RedisTranscriptStore(self.config, paths, storage) + elif self.config.storage_backend == "sql": + from ._sql_stores import SqlLongTermMemoryStore + from ._sql_stores import SqlSessionMemoryStore + from ._sql_stores import SqlToolResultStore + from ._sql_stores import SqlTranscriptStore + 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) + session_memory = SqlSessionMemoryStore(self.config, paths, storage) + tool_results = SqlToolResultStore(self.config, paths, storage) + transcripts = SqlTranscriptStore(self.config, paths, storage) + else: + long_term_memory = LongTermMemoryStore(self.config, paths) + session_memory = SessionMemoryStore(self.config, paths) + tool_results = ToolResultStore(self.config, paths) + transcripts = TranscriptStore(self.config, paths) + runtime = ScopedAdvancedMemoryRuntime( + root=self, + scope=scope, + paths=paths, + long_term_memory=long_term_memory, + session_memory=session_memory, + tool_results=tool_results, + transcripts=transcripts, + ) + 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() and not legacy_paths.session_root_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)) + if legacy_paths.session_root_dir.exists(): + shutil.move(str(legacy_paths.session_root_dir), str(scoped.paths.session_root_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") + async with self._sql_storage.create_db_session(): + pass + if self._sql_cleanup is not None: + await self._sql_cleanup.start() + 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.""" + 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_storage is not None: + if self._sql_cleanup is not None: + await self._sql_cleanup.close() + await self._sql_storage.close() + + +@dataclass(frozen=True) +class ScopedAdvancedMemoryRuntime: + """A tenant-bound view of an :class:`AdvancedMemoryRuntime`.""" + + root: AdvancedMemoryRuntime + scope: MemoryScope + paths: AdvancedMemoryPaths + long_term_memory: LongTermMemoryStore + session_memory: SessionMemoryStore + tool_results: ToolResultStore + transcripts: TranscriptStore + + @property + def config(self) -> AdvancedMemoryConfig: + """Return the root runtime configuration.""" + return self.root.config + + @property + def coordination(self) -> SessionOperationCoordinator: + """Return the shared coordinator.""" + return self.root.coordination + + def session_key(self, session_id: str) -> str: + """Return a lock/cache key unique across all tenants.""" + return f"{self.scope.storage_key}\0{session_id}" + + async def initialize(self) -> bool: + """Initialize only this tenant's local directories.""" + if not self.config.enabled: + return False + if self.config.storage_backend == "sql" and self.root._sql_cleanup is not None: + await self.root._sql_cleanup.start() + if self.config.storage_backend == "local" and self.root._local_cleanup is not None: + await self.root._local_cleanup.start() + await self.long_term_memory.initialize() + return True + + async def delete_session(self, session_id: str) -> None: + """Delete all Advanced Memory data belonging to one session.""" + if self.config.storage_backend == "local": + session_dir = self.paths.session_dir(session_id) + await asyncio.to_thread(shutil.rmtree, session_dir, True) + return + delete_session = getattr(self.session_memory, "delete_session", None) + if delete_session is None: + raise RuntimeError("Configured Advanced Memory backend cannot delete sessions") + await delete_session(session_id) diff --git a/trpc_agent_sdk/advanced_memory/_session_memory.py b/trpc_agent_sdk/advanced_memory/_session_memory.py index 39f9ab454..05623d48d 100644 --- a/trpc_agent_sdk/advanced_memory/_session_memory.py +++ b/trpc_agent_sdk/advanced_memory/_session_memory.py @@ -607,14 +607,14 @@ def missing_context(end: int) -> list[str]: return [], None - async def _read_current_memory(self, session_id: str) -> str: + async def _read_current_memory(self, session: "SessionABC") -> str: """Read old session memory or return the complete empty template.""" - current = await self._runtime.session_memory.read(session_id) + current = await self._runtime.for_session(session).session_memory.read(session.id) return current if current is not None else SessionMemoryDocument().to_markdown() async def _persist_checkpoint( self, - session_id: str, + session: "SessionABC", included_records: list[dict[str, Any]], document: SessionMemoryDocument, context_tokens: int | None, @@ -634,8 +634,9 @@ async def _persist_checkpoint( document.key_results, document.worklog, ) - await self._runtime.transcripts.append_unique( - session_id, + runtime = self._runtime.for_session(session) + await runtime.transcripts.append_unique( + session.id, { "schema_version": SESSION_MEMORY_CHECKPOINT_SCHEMA_VERSION, "kind": "session-memory-checkpoint", @@ -661,11 +662,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) + await runtime.initialize() + 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) + records = await runtime.transcripts.read_all(session.id) checkpoint = self._last_checkpoint(records) 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 @@ -699,7 +702,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 +718,9 @@ 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 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/_session_service.py b/trpc_agent_sdk/advanced_memory/_session_service.py index a8cccd52d..3c6574cee 100644 --- a/trpc_agent_sdk/advanced_memory/_session_service.py +++ b/trpc_agent_sdk/advanced_memory/_session_service.py @@ -73,30 +73,33 @@ def attach_session_memory_extractor( raise ValueError("Session memory extractor uses another runtime") self._session_memory_extractor = extractor - async def _ensure_initialized(self) -> None: + async def _ensure_initialized(self, session: SessionABC) -> 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() + self._initialized = await self._memory_runtime.for_session(session).initialize() - def _session_lock(self, session_id: str) -> CrossLoopLock: + def _session_lock(self, session: SessionABC) -> CrossLoopLock: """Return an independent asynchronous write lock per session.""" - lock = self._session_locks.get(session_id) + key = self._memory_runtime.for_session(session).session_key(session.id) + lock = self._session_locks.get(key) if lock is None: lock = CrossLoopLock() - self._session_locks[session_id] = lock + self._session_locks[key] = lock return lock - async def _load_parent_if_needed(self, session_id: str) -> None: + async def _load_parent_if_needed(self, session: SessionABC) -> None: """Restore the parent-chain tail before the first session write.""" - if session_id in self._loaded_parent_sessions: + runtime = self._memory_runtime.for_session(session) + key = runtime.session_key(session.id) + if key 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) + records = await runtime.transcripts.read_all(session.id) + self._last_event_ids[key] = find_last_event_id(records) + self._loaded_parent_sessions.add(key) async def create_session( self, @@ -142,16 +145,20 @@ async def list_sessions( 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): + """Delete the framework session and all Advanced Memory session data.""" + runtime = self._memory_runtime.for_scope(app_name, user_id) + scope_key = runtime.session_key(session_id) + lock = self._session_locks.setdefault(scope_key, CrossLoopLock()) + async with lock: 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) + await runtime.delete_session(session_id) + self._session_locks.pop(scope_key, None) + self._loaded_parent_sessions.discard(scope_key) + self._last_event_ids.pop(scope_key, None) async def append_event(self, session: SessionABC, event: ResponseABC) -> ResponseABC: """Append each persisted non-streaming Event in order.""" @@ -167,21 +174,23 @@ async def append_event(self, session: SessionABC, event: ResponseABC) -> Respons 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) + await self._ensure_initialized(session) + runtime = self._memory_runtime.for_session(session) + key = runtime.session_key(session.id) + async with self._session_lock(session): + await self._load_parent_if_needed(session) record = build_event_transcript_record( session, persisted_event, - parent_event_id=self._last_event_ids.get(session.id), + parent_event_id=self._last_event_ids.get(key), ) - _, appended = await self._memory_runtime.transcripts.append_unique( + _, appended = await runtime.transcripts.append_unique( session.id, record, unique_key="event_id", ) if appended: - self._last_event_ids[session.id] = record["event_id"] + self._last_event_ids[key] = record["event_id"] return persisted_event async def update_session(self, session: SessionABC) -> None: diff --git a/trpc_agent_sdk/advanced_memory/_sql_stores.py b/trpc_agent_sdk/advanced_memory/_sql_stores.py new file mode 100644 index 000000000..a49a7d6cc --- /dev/null +++ b/trpc_agent_sdk/advanced_memory/_sql_stores.py @@ -0,0 +1,560 @@ +"""SQL implementations of the Advanced Memory storage contracts.""" + +from __future__ import annotations + +import json +import asyncio +import hashlib +import uuid +from datetime import datetime, timedelta, timezone +from dataclasses import replace +from pathlib import Path +from collections.abc import Mapping +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 AdvancedMemoryConfig +from ._formats import MemoryDocument, MemoryIndexEntry, SessionMemoryDocument +from ._paths import AdvancedMemoryPaths + + +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 SqlSessionMemory(AdvancedMemorySqlBase): + __tablename__ = "advanced_memory_session_memory" + + 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) + session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_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 SqlTranscript(AdvancedMemorySqlBase): + __tablename__ = "advanced_memory_transcripts" + + 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) + session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + record_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + payload: Mapped[str] = mapped_column(Text) + recorded_at: Mapped[datetime] = mapped_column(PreciseTimestamp, default=func.now()) + expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) + + +class SqlTranscriptSeen(AdvancedMemorySqlBase): + __tablename__ = "advanced_memory_transcript_seen" + + dedupe_id: Mapped[str] = mapped_column(String(64), primary_key=True) + app_name: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) + user_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) + session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) + unique_key: Mapped[str] = mapped_column(String(DEFAULT_MAX_VARCHAR_LENGTH)) + unique_value: Mapped[str] = mapped_column(String(DEFAULT_MAX_VARCHAR_LENGTH)) + expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) + + +class SqlToolResult(AdvancedMemorySqlBase): + __tablename__ = "advanced_memory_tool_results" + + 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) + session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + result_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_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: AdvancedMemoryConfig, 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 + + async def _refresh_session_scope(self, db: Any, session_id: str) -> None: + expiry = self._expiry(self._config.session_ttl_seconds) + if expiry is None: + return + tables = ( + (SqlSessionMemory, (self._app_name, self._user_id, session_id)), + (SqlTranscript, (self._app_name, self._user_id, session_id)), + (SqlTranscriptSeen, (self._app_name, self._user_id, session_id)), + (SqlToolResult, (self._app_name, self._user_id, session_id)), + ) + for model, key in tables: + rows = await self._storage.query( + db, + SqlKey(key=key, storage_cls=model), + SqlCondition(filters=[ + getattr(model, "app_name") == self._app_name, + getattr(model, "user_id") == self._user_id, + getattr(model, "session_id") == session_id, + getattr(model, "expires_at").is_(None) | (getattr(model, "expires_at") > self._now()), + ]), + ) + for row in rows: + row.expires_at = expiry + + async def delete_session(self, session_id: str) -> None: + """Delete all Advanced Memory rows for one session.""" + models = ( + SqlSessionMemory, + SqlTranscript, + SqlTranscriptSeen, + SqlToolResult, + ) + filters = { + SqlSessionMemory: [ + SqlSessionMemory.app_name == self._app_name, + SqlSessionMemory.user_id == self._user_id, + SqlSessionMemory.session_id == session_id, + ], + SqlTranscript: [ + SqlTranscript.app_name == self._app_name, + SqlTranscript.user_id == self._user_id, + SqlTranscript.session_id == session_id, + ], + SqlTranscriptSeen: [ + SqlTranscriptSeen.app_name == self._app_name, + SqlTranscriptSeen.user_id == self._user_id, + SqlTranscriptSeen.session_id == session_id, + ], + SqlToolResult: [ + SqlToolResult.app_name == self._app_name, + SqlToolResult.user_id == self._user_id, + SqlToolResult.session_id == session_id, + ], + } + async with self._storage.create_db_session() as db: + for model in models: + await self._storage.delete( + db, + SqlKey(key=tuple(), storage_cls=model), + SqlCondition(filters=filters[model]), + ) + await self._storage.commit(db) + + +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) + await self._storage.commit(db) + content = row.content + lines, used_bytes = [], 0 + for line in content.splitlines(keepends=True)[:self._config.memory_index_max_lines]: + size = len(line.encode(self._config.encoding)) + if used_bytes + size > self._config.memory_index_max_bytes: + break + lines.append(line) + used_bytes += size + return "".join(lines) + + 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 SqlSessionMemoryStore(_SqlStore): + + async def read(self, session_id: str) -> str | None: + async with self._storage.create_db_session() as db: + row = await self._storage.get( + db, SqlKey(key=(self._app_name, self._user_id, session_id), storage_cls=SqlSessionMemory)) + if row is None or self._expired(row.expires_at): + return None + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return row.content + + async def write(self, session_id: str, document: SessionMemoryDocument) -> Path: + async with self._storage.create_db_session() as db: + key = (self._app_name, self._user_id, session_id) + row = await self._storage.get(db, SqlKey(key=key, storage_cls=SqlSessionMemory)) + if row is None: + row = SqlSessionMemory(app_name=key[0], user_id=key[1], session_id=key[2]) + await self._storage.add(db, row) + row.content = document.to_markdown() + row.updated_at = self._now() + row.expires_at = self._expiry(self._config.session_ttl_seconds) + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/summary") + + +class SqlToolResultStore(_SqlStore): + + async def write(self, session_id: str, result_id: str, serialized_result: str) -> Path: + async with self._storage.create_db_session() as db: + key = (self._app_name, self._user_id, session_id, result_id) + row = await self._storage.get(db, SqlKey(key=key, storage_cls=SqlToolResult)) + if row is None: + row = SqlToolResult( + app_name=key[0], + user_id=key[1], + session_id=key[2], + result_id=key[3], + ) + await self._storage.add(db, row) + row.content = serialized_result + row.updated_at = self._now() + row.expires_at = self._expiry(self._config.session_ttl_seconds) + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/tool/{result_id}") + + async def read(self, session_id: str, result_id: str) -> str | None: + async with self._storage.create_db_session() as db: + row = await self._storage.get( + db, + SqlKey(key=(self._app_name, self._user_id, session_id, result_id), storage_cls=SqlToolResult), + ) + if row is None or self._expired(row.expires_at): + return None + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return row.content + + +class SqlTranscriptStore(_SqlStore): + + def _dedupe_id(self, session_id: str, unique_key: str, value: str) -> str: + raw = "\0".join((self._app_name, self._user_id, session_id, unique_key, value)) + return hashlib.sha256(raw.encode("utf-8")).hexdigest() + + async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: + payload = dict(record) + payload.setdefault("recorded_at", self._now().replace(tzinfo=timezone.utc).isoformat()) + async with self._storage.create_db_session() as db: + await self._storage.add( + db, + SqlTranscript( + app_name=self._app_name, + user_id=self._user_id, + session_id=session_id, + record_id=uuid.uuid4().hex, + payload=json.dumps(payload, ensure_ascii=False), + expires_at=self._expiry(self._config.session_ttl_seconds), + )) + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript") + + async def append_unique( + self, + session_id: str, + record: Mapping[str, Any], + *, + unique_key: str, + ) -> tuple[Path, bool]: + payload = dict(record) + value = payload.get(unique_key) + if not isinstance(value, str) or not value: + raise ValueError(f"Transcript unique key {unique_key!r} must be a non-empty string") + async with self._storage.create_db_session() as db: + dedupe_id = self._dedupe_id(session_id, unique_key, value) + seen_key = (self._app_name, self._user_id, session_id, unique_key, value) + seen = await self._storage.get( + db, + SqlKey(key=(dedupe_id, ), storage_cls=SqlTranscriptSeen), + ) + if seen is not None and not self._expired(seen.expires_at): + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript"), False + if seen is not None: + await self._storage.delete( + db, + SqlKey(key=(dedupe_id, ), storage_cls=SqlTranscriptSeen), + SqlCondition(filters=[ + SqlTranscriptSeen.dedupe_id == dedupe_id, + ]), + ) + payload.setdefault("recorded_at", self._now().replace(tzinfo=timezone.utc).isoformat()) + await self._storage.add( + db, + SqlTranscriptSeen( + dedupe_id=dedupe_id, + app_name=seen_key[0], + user_id=seen_key[1], + session_id=seen_key[2], + unique_key=seen_key[3], + unique_value=seen_key[4], + expires_at=self._expiry(self._config.session_ttl_seconds), + )) + await self._storage.add( + db, + SqlTranscript( + app_name=self._app_name, + user_id=self._user_id, + session_id=session_id, + record_id=uuid.uuid4().hex, + payload=json.dumps(payload, ensure_ascii=False), + expires_at=self._expiry(self._config.session_ttl_seconds), + )) + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript"), True + + async def read_all(self, session_id: str) -> list[dict[str, Any]]: + async with self._storage.create_db_session() as db: + rows = await self._storage.query( + db, + SqlKey(key=(self._app_name, self._user_id, session_id), storage_cls=SqlTranscript), + SqlCondition( + filters=[ + SqlTranscript.app_name == self._app_name, + SqlTranscript.user_id == self._user_id, + SqlTranscript.session_id == session_id, + SqlTranscript.expires_at.is_(None) | (SqlTranscript.expires_at > self._now()), + ], + order_func=SqlTranscript.recorded_at.asc, + ), + ) + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return [json.loads(row.payload) for row in rows] + + +class SqlAdvancedMemoryCleanup: + """Periodically remove expired Advanced Memory SQL rows.""" + + _models = ( + SqlMemoryIndex, + SqlMemoryTopic, + SqlSessionMemory, + SqlTranscript, + SqlTranscriptSeen, + SqlToolResult, + ) + + def __init__(self, config: AdvancedMemoryConfig, 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 + and self._config.session_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]), + ) + 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", + "SqlSessionMemoryStore", + "SqlToolResultStore", + "SqlTranscriptStore", +] diff --git a/trpc_agent_sdk/advanced_memory/_storage.py b/trpc_agent_sdk/advanced_memory/_storage.py index 1fb43591c..0a8a011de 100644 --- a/trpc_agent_sdk/advanced_memory/_storage.py +++ b/trpc_agent_sdk/advanced_memory/_storage.py @@ -10,8 +10,10 @@ import asyncio import json import os +import shutil import tempfile import threading +import time from collections.abc import Mapping from dataclasses import replace from datetime import datetime @@ -44,6 +46,57 @@ def _atomic_write_text(path: Path, content: str, *, encoding: str) -> None: raise +def _is_expired(path: Path, ttl: int | None) -> bool: + if ttl is None or not path.exists(): + return False + return time.time() - path.stat().st_mtime >= ttl + + +def _touch(path: Path) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.touch() + + +def _expire_memory_dir(memory_dir: Path, config: AdvancedMemoryConfig) -> bool: + """Expire the whole long-term memory group using index activity time.""" + index_path = memory_dir / config.memory_index_name + if not _is_expired(index_path, config.memory_ttl_seconds): + return False + for path in memory_dir.glob("*.md"): + path.unlink(missing_ok=True) + return True + + +def _refresh_memory_dir(memory_dir: Path) -> None: + """Refresh activity for every file in the long-term memory group.""" + for path in memory_dir.glob("*.md"): + _touch(path) + + +def _session_activity_path(session_dir: Path) -> Path: + return session_dir / ".advanced-memory-activity" + + +def _expire_session_dir(session_dir: Path, config: AdvancedMemoryConfig) -> bool: + """Expire all Advanced Memory data belonging to one local session.""" + if not session_dir.exists() or config.session_ttl_seconds is None: + return False + activity_path = _session_activity_path(session_dir) + if activity_path.exists(): + expired = _is_expired(activity_path, config.session_ttl_seconds) + else: + files = [path for path in session_dir.rglob("*") if path.is_file()] + expired = bool(files) and time.time() - max(path.stat().st_mtime + for path in files) >= config.session_ttl_seconds + if expired: + shutil.rmtree(session_dir, ignore_errors=True) + return expired + + +def _refresh_session_dir(session_dir: Path) -> None: + _touch(_session_activity_path(session_dir)) + + class LongTermMemoryStore: """Manage MEMORY.md and its detail files in the same directory.""" @@ -73,8 +126,9 @@ async def read_index(self) -> str: def _read_index_sync(self) -> str: """Synchronously read MEMORY.md within configured limits.""" - if not self.index_path.exists(): + if _expire_memory_dir(self._paths.memory_dir, self._config) or not self.index_path.exists(): return "" + _refresh_memory_dir(self._paths.memory_dir) with self.index_path.open("r", encoding=self._config.encoding) as index_file: lines: list[str] = [] used_bytes = 0 @@ -99,27 +153,29 @@ async def write_index(self, entries: list[MemoryIndexEntry]) -> None: 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) + _refresh_memory_dir(self._paths.memory_dir) 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) + return await asyncio.to_thread(self._read_topic_sync, path) + + def _read_topic_sync(self, path: Path) -> str | None: + if _expire_memory_dir(self._paths.memory_dir, self._config) or not path.exists(): + return None + _refresh_memory_dir(self._paths.memory_dir) + return path.read_text(encoding=self._config.encoding) 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) + return await asyncio.to_thread(self._read_frontmatter_sync, path) - def _read_frontmatter(self, path: Path) -> str | None: + def _read_frontmatter_sync(self, path: Path) -> str | None: """Synchronously read a topic's bounded frontmatter block.""" - if not path.exists(): + if _expire_memory_dir(self._paths.memory_dir, self._config) or not path.exists(): return None + _refresh_memory_dir(self._paths.memory_dir) lines: list[str] = [] with path.open(encoding=self._config.encoding) as file: for line in file: @@ -132,22 +188,25 @@ 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, - ) + await asyncio.to_thread(self._write_topic_sync, path, document.to_markdown()) return path + def _write_topic_sync(self, path: Path, content: str) -> None: + _expire_memory_dir(self._paths.memory_dir, self._config) + _atomic_write_text(path, content, encoding=self._config.encoding) + _refresh_memory_dir(self._paths.memory_dir) + 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 _expire_memory_dir(self._paths.memory_dir, self._config): + return [] if not self._paths.memory_dir.exists(): return [] + _refresh_memory_dir(self._paths.memory_dir) 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, @@ -165,25 +224,33 @@ def __init__(self, config: AdvancedMemoryConfig, paths: AdvancedMemoryPaths | No 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) + return await asyncio.to_thread(self._read_sync, session_id, path) - def _read_sync(self, path: Path) -> str | None: + def _read_sync(self, session_id: str, path: Path) -> str | None: """Synchronously read session memory.""" - if not path.exists(): + session_dir = self._paths.session_dir(session_id) + if _expire_session_dir(session_dir, self._config) or not path.exists(): return None + _refresh_session_dir(session_dir) 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, + self._write_sync, + session_id, path, document.to_markdown(), - encoding=self._config.encoding, ) return path + def _write_sync(self, session_id: str, path: Path, content: str) -> None: + session_dir = self._paths.session_dir(session_id) + _expire_session_dir(session_dir, self._config) + _atomic_write_text(path, content, encoding=self._config.encoding) + _refresh_session_dir(session_dir) + class ToolResultStore: """Persist complete tool results that exceed the context budget.""" @@ -197,24 +264,32 @@ async def write(self, session_id: str, result_id: str, serialized_result: str) - """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, + self._write_sync, + session_id, 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) + return await asyncio.to_thread(self._read_sync, session_id, path) - def _read_sync(self, path: Path) -> str | None: + def _read_sync(self, session_id: str, path: Path) -> str | None: """Synchronously read an optional complete tool-result file.""" - if not path.exists(): + session_dir = self._paths.session_dir(session_id) + if _expire_session_dir(session_dir, self._config) or not path.exists(): return None + _refresh_session_dir(session_dir) return path.read_text(encoding=self._config.encoding) + def _write_sync(self, session_id: str, path: Path, content: str) -> None: + session_dir = self._paths.session_dir(session_id) + _expire_session_dir(session_dir, self._config) + _atomic_write_text(path, content, encoding=self._config.encoding) + _refresh_session_dir(session_dir) + class TranscriptStore: """Store complete per-session records as append-only JSONL.""" @@ -237,9 +312,11 @@ async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: def _append_sync(self, path: Path, serialized: str) -> None: """Synchronously append one transcript line under the write lock.""" + _expire_session_dir(path.parent, self._config) path.parent.mkdir(parents=True, exist_ok=True) with self._write_lock: self._append_serialized_unlocked(path, serialized) + _refresh_session_dir(path.parent) def _append_serialized_unlocked(self, path: Path, serialized: str) -> None: """Append one serialized line while the caller holds the lock.""" @@ -282,9 +359,13 @@ def _append_unique_sync( 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: + if _expire_session_dir(path.parent, self._config): + for cache_key in list(self._seen_unique_values): + if cache_key[0] == path: + self._seen_unique_values.pop(cache_key, None) + path.parent.mkdir(parents=True, exist_ok=True) + cache_key = (path, unique_key) seen_values = self._seen_unique_values.get(cache_key) if seen_values is None: seen_values = self._load_unique_values_unlocked(path, unique_key) @@ -293,6 +374,7 @@ def _append_unique_sync( return False self._append_serialized_unlocked(path, serialized) seen_values.add(unique_value) + _refresh_session_dir(path.parent) return True def _load_unique_values_unlocked(self, path: Path, unique_key: str) -> set[str]: @@ -317,8 +399,11 @@ async def read_all(self, session_id: str) -> list[dict[str, Any]]: 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 _expire_session_dir(path.parent, self._config): + return [] if not path.exists(): return [] + _refresh_session_dir(path.parent) 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): @@ -329,3 +414,75 @@ def _read_all_sync(self, path: Path) -> list[dict[str, Any]]: raise ValueError(f"Transcript line {line_number} is not a JSON object") records.append(parsed) return records + + +class LocalAdvancedMemoryCleanup: + """Periodically remove expired local Advanced Memory data.""" + + def __init__(self, config: AdvancedMemoryConfig) -> 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: + return + if self._config.memory_ttl_seconds is None and self._config.session_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] + session_roots = [root / self._config.session_dir_name] + tenants_root = root / "tenants" + if tenants_root.exists(): + for app_dir in tenants_root.iterdir(): + if app_dir.is_dir(): + for user_dir in app_dir.iterdir(): + if user_dir.is_dir(): + memory_dirs.append(user_dir / self._config.memory_dir_name) + session_roots.append(user_dir / self._config.session_dir_name) + for memory_dir in memory_dirs: + _expire_memory_dir(memory_dir, self._config) + for session_root in session_roots: + if session_root.exists(): + for session_dir in session_root.iterdir(): + if session_dir.is_dir(): + _expire_session_dir(session_dir, self._config) + + async def _run(self) -> None: + if self._stop_event is None: + return + ttls = [ + ttl for ttl in ( + self._config.memory_ttl_seconds, + self._config.session_ttl_seconds, + ) if ttl is not None + ] + interval = min(ttls) if ttls else 60 + try: + while not self._stop_event.is_set(): + try: + await asyncio.wait_for(self._stop_event.wait(), timeout=interval) + break + except asyncio.TimeoutError: + await self.cleanup_once() + except asyncio.CancelledError: + raise + + async def close(self) -> None: + if self._task is not None: + await self.cleanup_once() + 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/advanced_memory/_storage_backend.py b/trpc_agent_sdk/advanced_memory/_storage_backend.py new file mode 100644 index 000000000..4e10a2f0c --- /dev/null +++ b/trpc_agent_sdk/advanced_memory/_storage_backend.py @@ -0,0 +1,30 @@ +"""Storage boundary for Advanced Memory tenant namespaces. + +Backends expose logical records rather than filesystem paths so a future Redis +implementation can preserve the same tenant and session semantics. +""" + +from __future__ import annotations + +from typing import Protocol + +from ._paths import MemoryScope +from ._runtime import ScopedAdvancedMemoryRuntime + + +class AdvancedMemoryStorageBackend(Protocol): + """Create storage views isolated to an application user.""" + + def for_scope(self, scope: MemoryScope) -> ScopedAdvancedMemoryRuntime: + """Return the tenant-bound storage view.""" + + +class LocalAdvancedMemoryStorageBackend: + """Adapt the file-backed runtime to the storage backend boundary.""" + + def __init__(self, runtime: object) -> None: + self._runtime = runtime + + def for_scope(self, scope: MemoryScope) -> ScopedAdvancedMemoryRuntime: + """Return a file-backed scope without exposing local path mechanics.""" + return self._runtime.for_scope(scope.app_name, scope.user_id) diff --git a/trpc_agent_sdk/advanced_memory/_tool_result_budget.py b/trpc_agent_sdk/advanced_memory/_tool_result_budget.py index 73a710aa2..6bfc9d6d8 100644 --- a/trpc_agent_sdk/advanced_memory/_tool_result_budget.py +++ b/trpc_agent_sdk/advanced_memory/_tool_result_budget.py @@ -8,6 +8,7 @@ from __future__ import annotations import asyncio +import copy import hashlib import json from dataclasses import dataclass @@ -120,6 +121,7 @@ def __init__(self, memory_runtime: AdvancedMemoryRuntime) -> None: 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: @@ -128,15 +130,17 @@ def runtime(self) -> AdvancedMemoryRuntime: 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) + 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) @@ -163,7 +167,7 @@ async def _load_state(self, session_id: str) -> ToolResultBudgetState: replacements=replacements, result_hashes=result_hashes, ) - self._states[session_id] = state + self._states[state_key] = state return state def _collect_candidates(self, request: "LlmRequest") -> list[list[ToolResultCandidate]]: @@ -206,7 +210,11 @@ def _build_replacement( 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) + persisted_path = (Path(f"advanced-memory://{self._runtime.config.redis_key_prefix}/" + f"{self._runtime.scope.app_name}/{self._runtime.scope.user_id}/{session_id}/" + f"tool/{candidate.result_id}") if hasattr(self._runtime, "scope") + and self._runtime.config.storage_backend == "redis" else self._runtime.paths.tool_result_path( + session_id, candidate.result_id)) preview, truncated = _preview_text( candidate.serialized_result, self._runtime.config.tool_result_preview_chars, @@ -217,7 +225,7 @@ def _build_replacement( "schema_version": TOOL_RESULT_REPLACEMENT_SCHEMA_VERSION, }, "persisted_output": { - "message": "The tool result exceeded the context budget; the complete content was saved to disk.", + "message": "The tool result exceeded the context budget; the complete content was persisted.", "path": str(persisted_path), "original_chars": candidate.original_size, "preview": preview, @@ -281,11 +289,17 @@ async def _persist_replacement( ) -> None: """Persist the full result before appending its replacement record.""" candidate = replacement.candidate - await self._runtime.tool_results.write( + persisted_path = await self._runtime.tool_results.write( session_id, candidate.result_id, candidate.serialized_result, ) + persisted_path_text = str(persisted_path).replace( + "advanced-memory:/", + "advanced-memory://", + 1, + ) + replacement.replacement_response["persisted_output"]["path"] = persisted_path_text await self._runtime.transcripts.append_unique( session_id, { @@ -296,7 +310,7 @@ async def _persist_replacement( "tool_name": candidate.tool_name, "original_chars": candidate.original_size, "original_sha256": tool_result_sha256(candidate.serialized_result), - "persisted_path": str(replacement.persisted_path), + "persisted_path": persisted_path_text, "replacement_response": replacement.replacement_response, }, unique_key="decision_id", @@ -322,11 +336,31 @@ async def _persist_seen_decision( unique_key="decision_id", ) - async def apply(self, request: "LlmRequest", *, session_id: str) -> ToolResultBudgetResult: + async def apply( + self, + request: "LlmRequest", + *, + session_id: str, + 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"): + await self._runtime.initialize() + return await self._apply_scoped(request, session_id) + 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) -> 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) @@ -390,7 +424,7 @@ 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 diff --git a/trpc_agent_sdk/memory/_advanced_memory_service.py b/trpc_agent_sdk/memory/_advanced_memory_service.py index 8cc2c97f8..5ad1dbdd5 100644 --- a/trpc_agent_sdk/memory/_advanced_memory_service.py +++ b/trpc_agent_sdk/memory/_advanced_memory_service.py @@ -134,4 +134,4 @@ async def close(self) -> None: Advanced Memory stores are file-backed and do not own an external connection. The wrapped session service is closed by Runner. """ - return None + await self._runtime.close() diff --git a/trpc_agent_sdk/sessions/_advanced_memory_session_service.py b/trpc_agent_sdk/sessions/_advanced_memory_session_service.py index feb14fb14..8ad94f4d3 100644 --- a/trpc_agent_sdk/sessions/_advanced_memory_session_service.py +++ b/trpc_agent_sdk/sessions/_advanced_memory_session_service.py @@ -51,17 +51,24 @@ 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" + def _scoped_runtime(self, app_name: str, user_id: str) -> AdvancedMemoryRuntime: + return self._runtime.for_scope(app_name, user_id) - @property - def _state_path(self) -> Path: - return self._runtime.paths.session_root_dir / "_state.json" + def _metadata_path(self, app_name: str, user_id: str, session_id: str) -> Path: + return self._scoped_runtime(app_name, user_id).paths.session_dir(session_id) / "session.json" + + def _app_state_path(self, app_name: str, user_id: str) -> Path: + """Return state shared by every user of one app.""" + return self._scoped_runtime(app_name, user_id).paths.tenant_root_dir.parent / "_state.json" + + def _user_state_path(self, app_name: str, user_id: str) -> Path: + """Return state private to one application user.""" + return self._scoped_runtime(app_name, user_id).paths.tenant_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) + path = self._metadata_path(session.app_name, session.user_id, session.id) await asyncio.to_thread(self._write_json, path, payload, self._runtime.config.encoding) @staticmethod @@ -81,8 +88,8 @@ def _write_json(path: Path, payload: dict[str, Any], encoding: str) -> None: pass raise - async def _read_session(self, session_id: str) -> Session | None: - path = self._metadata_path(session_id) + async def _read_session(self, app_name: str, user_id: str, session_id: str) -> Session | None: + path = self._metadata_path(app_name, user_id, session_id) if not path.exists(): return None payload = await asyncio.to_thread(path.read_text, encoding=self._runtime.config.encoding) @@ -120,10 +127,10 @@ async def _cleanup_loop(self) -> None: 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(): + tenants_root = self._runtime.config.root_dir / "tenants" + if not tenants_root.exists(): return - for metadata_path in root.glob("*/session.json"): + for metadata_path in tenants_root.glob(f"*/*/{self._runtime.config.session_dir_name}/*/session.json"): try: if metadata_path.stat().st_mtime < cutoff: shutil.rmtree(metadata_path.parent, ignore_errors=True) @@ -142,24 +149,35 @@ async def _stop_cleanup_task(self) -> None: 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) + async def _read_global_state(self, app_name: str, user_id: str) -> dict[str, dict[str, Any]]: + + async def read(path: Path) -> dict[str, Any]: + if not path.exists(): + return {} + payload = await asyncio.to_thread(path.read_text, encoding=self._runtime.config.encoding) + return dict(json.loads(payload)) + return { - "app": dict(parsed.get("app", {})), - "user": dict(parsed.get("user", {})), + "app": await read(self._app_state_path(app_name, user_id)), + "user": await read(self._user_state_path(app_name, user_id)), } - 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 _write_global_state(self, app_name: str, user_id: str, state: dict[str, dict[str, Any]]) -> None: + await asyncio.to_thread( + self._write_json, + self._app_state_path(app_name, user_id), + state["app"], + self._runtime.config.encoding, + ) + await asyncio.to_thread( + self._write_json, + self._user_state_path(app_name, user_id), + state["user"], + self._runtime.config.encoding, + ) async def _restore_events(self, session: Session) -> Session: - records = await self._runtime.transcripts.read_all(session.id) + records = await self._scoped_runtime(session.app_name, session.user_id).transcripts.read_all(session.id) events: list[Event] = [] for record in records: event_payload = record.get("event") @@ -191,24 +209,19 @@ async def create_session( 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._scoped_runtime(app_name, user_id).initialize() + await self._read_session(app_name, user_id, resolved_id) + global_state = await self._read_global_state(app_name, user_id) + global_state["app"].update(state_delta.app_state_delta) + global_state["user"].update(state_delta.user_state_delta) + await self._write_global_state(app_name, user_id, 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() - }) + session.state.update({f"app:{key}": value for key, value in global_state["app"].items()}) + session.state.update({f"user:{key}": value for key, value in global_state["user"].items()}) return session async def get_session( @@ -221,12 +234,12 @@ async def get_session( ) -> 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: + session = await self._read_session(app_name, user_id, session_id) + if session is None: 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}", {}) + global_state = await self._read_global_state(app_name, user_id) + app_state = global_state["app"] + user_state = global_state["user"] session.state = merge_state( extract_state_delta(session.state), need_copy=True, @@ -242,10 +255,17 @@ async def list_sessions( user_id: Optional[str] = None, ) -> ListSessionsResponse: self._start_cleanup_task() - if not self._runtime.paths.session_root_dir.exists(): + tenants_root = self._runtime.config.root_dir / "tenants" + if not tenants_root.exists(): return ListSessionsResponse() sessions: list[Session] = [] - for path in await asyncio.to_thread(lambda: list(self._runtime.paths.session_root_dir.glob("*/session.json"))): + if user_id is not None: + root = self._scoped_runtime(app_name, user_id).paths.session_root_dir + session_glob = "*/session.json" + else: + root = tenants_root + session_glob = f"*/*/{self._runtime.config.session_dir_name}/*/session.json" + for path in await asyncio.to_thread(lambda: list(root.glob(session_glob))): try: session = await asyncio.to_thread(lambda path=path: Session.model_validate( json.loads(path.read_text(encoding=self._runtime.config.encoding)))) @@ -266,7 +286,7 @@ async def delete_session(self, *, app_name: str, user_id: str, session_id: str) ) if session is not None: async with self._lock: - await asyncio.to_thread(shutil.rmtree, self._runtime.paths.session_dir(session_id), True) + await asyncio.to_thread(shutil.rmtree, self._metadata_path(app_name, user_id, session_id).parent, True) async def append_event(self, session: Session, event: Event) -> Event: self._start_cleanup_task() @@ -275,22 +295,22 @@ async def append_event(self, session: Session, event: Event) -> 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) + global_state = await self._read_global_state(session.app_name, session.user_id) + global_state["app"].update(state_delta.app_state_delta) + global_state["user"].update(state_delta.user_state_delta) + await self._write_global_state(session.app_name, session.user_id, 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) + runtime = self._scoped_runtime(session.app_name, session.user_id) + records = await 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( + await runtime.transcripts.append_unique( session.id, record, unique_key="event_id", @@ -315,6 +335,7 @@ async def get_session_summary(self, session: Session) -> str | None: async def close(self) -> None: await self._stop_cleanup_task() + await self._runtime.close() class AdvancedMemorySessionService(BaseSessionService): @@ -331,6 +352,9 @@ def __init__( 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) + if self._runtime.config.storage_backend == "redis": + raise ValueError("AdvancedMemorySessionService is file-backed; use RedisSessionService with " + "AdvancedMemoryService when AdvancedMemoryConfig.storage_backend='redis'") self._preload_memory_model = preload_memory_model self._backend = _AdvancedMemorySessionBackend(self._runtime, session_config=session_config) self._integration: Any | None = None diff --git a/trpc_agent_sdk/sessions/_sql_session_service.py b/trpc_agent_sdk/sessions/_sql_session_service.py index 4333ffeb1..2d6d9d3ad 100644 --- a/trpc_agent_sdk/sessions/_sql_session_service.py +++ b/trpc_agent_sdk/sessions/_sql_session_service.py @@ -397,6 +397,10 @@ def __init__(self, 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 @@ -704,6 +708,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 +722,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 +739,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/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..f302c6fdb 100644 --- a/trpc_agent_sdk/tools/_advanced_memory_tool.py +++ b/trpc_agent_sdk/tools/_advanced_memory_tool.py @@ -43,9 +43,9 @@ 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), @@ -66,6 +66,23 @@ def owns_tool(self, tool: Any) -> bool: 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, filename: str, @@ -74,6 +91,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 +105,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,8 +119,8 @@ 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, @@ -110,9 +129,9 @@ async def save_memory( "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,11 +152,12 @@ async def read_memory(self, filename: str) -> dict: "update this memory if it is outdated or incorrect."), } - async def list_memory_index(self) -> dict: + async def list_memory_index(self, tool_context: Any | None = None) -> dict: """Return the current long-term memory index and its disk path.""" + 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": str(runtime.paths.memory_index_path), + "index": await runtime.long_term_memory.read_index(), } From e955a5919acc484367f80e24f98d79f308927dfb Mon Sep 17 00:00:00 2001 From: congkechen Date: Wed, 9 Sep 2026 13:14:02 +0800 Subject: [PATCH 2/6] =?UTF-8?q?feature:=20=E4=BC=98=E5=8C=96=E6=96=87?= =?UTF-8?q?=E4=BB=B6=E5=AD=98=E5=82=A8=E8=B7=AF=E5=BE=84=20/=20=E5=8F=AF?= =?UTF-8?q?=E9=80=89=E4=BF=9D=E7=95=99=E5=8E=9F=E5=A7=8B=E8=AE=B0=E5=BD=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../README.md | 7 +- .../.env | 3 +- .../README.md | 4 +- .../README.md | 6 +- .../test_advanced_memory_session_service.py | 29 ++++++++ .../test_advanced_memory_tools.py | 31 +++++++++ tests/advanced_memory/test_redis_stores.py | 23 ++++++- tests/advanced_memory/test_storage.py | 46 ++++++++++--- .../advanced_memory/_autocompact.py | 4 +- trpc_agent_sdk/advanced_memory/_config.py | 1 + .../advanced_memory/_memory_context.py | 4 +- trpc_agent_sdk/advanced_memory/_paths.py | 67 +++++++++++++++++++ .../advanced_memory/_redis_stores.py | 15 +++-- trpc_agent_sdk/advanced_memory/_sql_stores.py | 21 ++++-- trpc_agent_sdk/advanced_memory/_storage.py | 15 ++++- .../advanced_memory/_tool_result_budget.py | 18 +++-- .../_advanced_memory_session_service.py | 15 ++++- trpc_agent_sdk/tools/_advanced_memory_tool.py | 11 ++- 18 files changed, 273 insertions(+), 47 deletions(-) diff --git a/examples/memory_service_with_advanced_memory/README.md b/examples/memory_service_with_advanced_memory/README.md index 17b534f21..fa3d0bcb4 100644 --- a/examples/memory_service_with_advanced_memory/README.md +++ b/examples/memory_service_with_advanced_memory/README.md @@ -136,8 +136,10 @@ python3 run_agent.py `M_TTL` 和 `SESSION_TTL` 未配置时不会自动删除数据。Session 的后台清理检查间隔由示例内部设置,不需要单独配置。 -本示例提供的 `.env` 默认使用 `M_TTL=120` 和 `SESSION_TTL=60`,方便直接观察 -过期清理;如果不希望自动删除,将这两个值留空即可。 +默认情况下,Session TTL 过期会保留 transcript,便于审计;只有将 +`session_ttl_delete_transcripts=True` 时,transcript 才会随 Session TTL 一起删除。 + +本示例提供的 `.env` 默认使用 `M_TTL=120` 和 `SESSION_TTL=60`,方便直接观察过期清理;如果不希望自动删除,将这两个值留空即可。 `.env` 中留空的变量不会覆盖默认值;如果同时在 Python 中传入`model_context_window_tokens` 或 `max_output_tokens`,Python 显式配置优先。 @@ -165,6 +167,7 @@ session_service = AdvancedMemorySessionService( # TTL(单位:秒;None 表示不过期) memory_ttl_seconds=120, # 长期记忆 TTL(秒) session_ttl_seconds=60, # 会话记忆 TTL(秒) + session_ttl_delete_transcripts=False, # Session TTL 是否删除 transcript # 长期记忆 memory_index_max_lines=200, # 注入 prompt 的索引最大行数 diff --git a/examples/memory_service_with_advanced_memory_redis/.env b/examples/memory_service_with_advanced_memory_redis/.env index ed021e751..3982337a7 100644 --- a/examples/memory_service_with_advanced_memory_redis/.env +++ b/examples/memory_service_with_advanced_memory_redis/.env @@ -1,10 +1,9 @@ -REDIS_URL=redis://localhost:6379/0 +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= - # 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= diff --git a/examples/memory_service_with_advanced_memory_redis/README.md b/examples/memory_service_with_advanced_memory_redis/README.md index c2181ce46..c65167b29 100644 --- a/examples/memory_service_with_advanced_memory_redis/README.md +++ b/examples/memory_service_with_advanced_memory_redis/README.md @@ -223,7 +223,7 @@ runner = Runner( ```text user: Do you remember my name? 🔧 tool call: list_memory_index({}) -📊 Tool Result: {'index_path': '/data/workspace/trpc-agent-python-am-service/examples/memory_service_with_advanced_memory_redis/tenants/advanced-memory-redis-demo/redis-demo-user/MEMORY/MEMORY.md', 'index': ''} +📊 Tool Result: {'index_path': 'advanced-memory://redis/advanced-memory-redis-demo:v1:{advanced-memory-redis-demo:redis-demo-user}:memory:index', 'index': ''} 🤖 Assistant: I checked my long-term memory, but I'm afraid I don't have anything saved yet — the memory index is currently empty, so I don't know your name. If you'd like, just tell me your name (and anything else you'd like me to remember about you), and I'll save it so I can recall it in future conversations! @@ -232,7 +232,7 @@ If you'd like, just tell me your name (and anything else you'd like me to rememb 📝 user: Do you remember my favorite color? 🔧 tool call: list_memory_index({}) -📊 Tool Result: {'index_path': '/data/workspace/trpc-agent-python-am-service/examples/memory_service_with_advanced_memory_redis/tenants/advanced-memory-redis-demo/redis-demo-user/MEMORY/MEMORY.md', 'index': ''} +📊 Tool Result: {'index_path': 'advanced-memory://redis/advanced-memory-redis-demo:v1:{advanced-memory-redis-demo:redis-demo-user}:memory:index', 'index': ''} 🤖 Assistant: I checked my long-term memory, but I don't have anything saved about your favorite color yet — my memory index is currently empty. If you'd like, tell me your favorite color and I'll remember it for future conversations. 💬 diff --git a/examples/memory_service_with_advanced_memory_sql/README.md b/examples/memory_service_with_advanced_memory_sql/README.md index 18540b19f..6eb268dd8 100644 --- a/examples/memory_service_with_advanced_memory_sql/README.md +++ b/examples/memory_service_with_advanced_memory_sql/README.md @@ -136,7 +136,7 @@ runner = Runner( ----- Runner A, query 1 ----- 📝 user: Do you remember my name? 🔧 tool call: list_memory_index({}) -📊 Tool Result: {'index_path': '/data/workspace/trpc-agent-python-am-service/examples/memory_service_with_advanced_memory_sql/tenants/advanced-memory-sql-demo/sql-demo-user/MEMORY/MEMORY.md', 'index': ''} +📊 Tool Result: {'index_path': 'advanced-memory://sql/advanced-memory-sql-demo/sql-demo-user/memory/index', 'index': ''} 🤖 Assistant: I checked my long-term memory index, and it's currently empty — I don't have any saved memories about you yet, so I don't remember your name. If you'd like, tell me your name (or anything else you'd like me to remember about you), and I'll save it to my memory so I can remember it across future conversations. @@ -144,7 +144,7 @@ If you'd like, tell me your name (or anything else you'd like me to remember abo ----- Runner A, query 2 ----- 📝 user: Do you remember my favorite color? 🔧 tool call: list_memory_index({}) -📊 Tool Result: {'index_path': '/data/workspace/trpc-agent-python-am-service/examples/memory_service_with_advanced_memory_sql/tenants/advanced-memory-sql-demo/sql-demo-user/MEMORY/MEMORY.md', 'index': ''} +📊 Tool Result: {'index_path': 'advanced-memory://sql/advanced-memory-sql-demo/sql-demo-user/memory/index', 'index': ''} 🤖 Assistant: I checked my long-term memory index, and it's currently empty — I don't have any saved memories yet, so I don't remember your favorite color. If you tell me your favorite color (or any other preferences you'd like me to keep in mind), I can save it for future conversations! @@ -163,7 +163,7 @@ If you tell me your favorite color (or any other preferences you'd like me to ke 📝 user: Do you remember my name? 🔧 tool call: list_memory_index({}) 🔧 tool call: read_memory({'filename': 'user_identity.md'}) -📊 Tool Result: {'index_path': '/data/workspace/trpc-agent-python-am-service/examples/memory_service_with_advanced_memory_sql/tenants/advanced-memory-sql-demo/sql-demo-user/MEMORY/MEMORY.md', 'index': '- [User identity](user_identity.md):Alice, favorite color blue.\n'} +📊 Tool Result: {'index_path': 'advanced-memory://sql/advanced-memory-sql-demo/sql-demo-user/memory/index', 'index': '- [User identity](user_identity.md):Alice, favorite color blue.\n'} 📊 Tool Result: {'found': True, 'filename': 'user_identity.md', 'content': "---\nname: User identity\ndescription: The user's name and basic identity.\ntype: user\nupdated_at: 2026-09-08T07:26:18.283583+00:00\n---\n---\nname: User identity\ndescription: The user's name and basic identity.\ntype: user\nupdated_at: 2026-09-08T07:26:06.708029+00:00\n---\nThe user's name is Alice. She introduced herself in conversation. Her favorite color is blue.\n", 'updated_at': '2026-09-08T07:26:18.283583+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** — from my long-term memory, that's how you introduced yourself. I also have noted that your favorite color is blue. 😊 diff --git a/tests/advanced_memory/test_advanced_memory_session_service.py b/tests/advanced_memory/test_advanced_memory_session_service.py index 25fbb51aa..cd07c59e9 100644 --- a/tests/advanced_memory/test_advanced_memory_session_service.py +++ b/tests/advanced_memory/test_advanced_memory_session_service.py @@ -150,6 +150,35 @@ async def test_ttl_cleanup_removes_expired_persistent_sessions(tmp_path: Path) - await service.close() +async def test_ttl_cleanup_preserves_transcript_by_default(tmp_path: Path) -> None: + """Keep the transcript when session metadata expires.""" + 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="preserve-transcript", + ) + await service.append_event(session, _event("event-1", "hello")) + transcript_path = service.runtime.for_session(session).paths.transcript_path(session.id) + + await asyncio.sleep(1.1) + + assert await service.get_session( + app_name="demo-app", + user_id="demo-user", + session_id=session.id, + ) is None + assert transcript_path.exists() + 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)) diff --git a/tests/advanced_memory/test_advanced_memory_tools.py b/tests/advanced_memory/test_advanced_memory_tools.py index 112511699..cf3347222 100644 --- a/tests/advanced_memory/test_advanced_memory_tools.py +++ b/tests/advanced_memory/test_advanced_memory_tools.py @@ -3,10 +3,13 @@ from __future__ import annotations from pathlib import Path +from types import SimpleNamespace +from unittest.mock import AsyncMock import pytest from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig +from trpc_agent_sdk.advanced_memory import AdvancedMemoryPaths from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime from trpc_agent_sdk.tools import AdvancedMemoryTools from trpc_agent_sdk.tools import create_advanced_memory_tools @@ -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 = AdvancedMemoryConfig( + 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_redis_stores.py b/tests/advanced_memory/test_redis_stores.py index 8d0e9f8a3..8718a9849 100644 --- a/tests/advanced_memory/test_redis_stores.py +++ b/tests/advanced_memory/test_redis_stores.py @@ -67,7 +67,7 @@ async def test_session_writes_refresh_all_session_keys() -> None: @pytest.mark.asyncio async def test_ttl_refresh_includes_previously_tracked_keys() -> None: - store = _store(RedisSessionMemoryStore) + store = _store(RedisSessionMemoryStore, session_ttl_delete_transcripts=True) session_base = store._session_base("session-1") old_key = f"{session_base}:transcript" store._command = AsyncMock(side_effect=[ @@ -85,6 +85,27 @@ async def test_ttl_refresh_includes_previously_tracked_keys() -> None: assert ("expire", f"{session_base}:summary", 60) in commands +@pytest.mark.asyncio +async def test_ttl_refresh_preserves_transcript_by_default() -> None: + store = _store(RedisSessionMemoryStore) + session_base = store._session_base("session-1") + old_key = f"{session_base}:transcript" + old_seen_key = f"{old_key}:seen:event_id" + store._command = AsyncMock(side_effect=[ + None, # SADD + [old_key.encode(), old_seen_key.encode()], # SMEMBERS + None, # EXPIRE current key + None, # EXPIRE registry + ]) + + await store._refresh_session_ttl("session-1", f"{session_base}:summary") + + commands = [call.args for call in store._command.await_args_list] + assert ("expire", old_key, 60) not in commands + assert ("expire", old_seen_key, 60) not in commands + assert ("expire", f"{session_base}:summary", 60) in commands + + @pytest.mark.asyncio async def test_memory_write_lock_releases_with_token_check() -> None: store = _store(RedisLongTermMemoryStore) diff --git a/tests/advanced_memory/test_storage.py b/tests/advanced_memory/test_storage.py index 99ac2b024..4de28c71d 100644 --- a/tests/advanced_memory/test_storage.py +++ b/tests/advanced_memory/test_storage.py @@ -236,11 +236,13 @@ async def test_transcript_append_unique_uses_persisted_ids(tmp_path: Path) -> No async def test_transcript_unique_cache_is_reset_after_session_ttl(tmp_path: Path, ) -> None: - """Allow a reused session ID to append after local TTL expiration.""" - runtime = AdvancedMemoryRuntime.create(_enabled_config( - tmp_path, - session_ttl_seconds=1, - )) + """Allow a reused session ID to append after transcript deletion.""" + runtime = AdvancedMemoryRuntime.create( + _enabled_config( + tmp_path, + session_ttl_seconds=1, + session_ttl_delete_transcripts=True, + )) transcript = runtime.transcripts await transcript.append_unique( "session-a", @@ -327,11 +329,13 @@ async def test_memory_index_is_truncated_when_read_over_byte_budget(tmp_path: Pa async def test_local_ttl_expires_memory_and_session_groups(tmp_path: Path) -> None: """Expire local memory groups after their last activity.""" - runtime = AdvancedMemoryRuntime.create(_enabled_config( - tmp_path, - memory_ttl_seconds=1, - session_ttl_seconds=1, - )) + runtime = AdvancedMemoryRuntime.create( + _enabled_config( + tmp_path, + memory_ttl_seconds=1, + session_ttl_seconds=1, + session_ttl_delete_transcripts=True, + )) scoped = runtime.for_scope("app", "user") await scoped.initialize() await scoped.long_term_memory.write_index([ @@ -356,6 +360,28 @@ async def test_local_ttl_expires_memory_and_session_groups(tmp_path: Path) -> No await runtime.close() +async def test_local_session_ttl_preserves_transcripts_by_default(tmp_path: Path) -> None: + """Keep local transcripts when session TTL cleanup uses its default.""" + runtime = AdvancedMemoryRuntime.create(_enabled_config( + tmp_path, + session_ttl_seconds=1, + )) + scoped = runtime.for_scope("app", "user") + await scoped.initialize() + await scoped.session_memory.write("session", SessionMemoryDocument(session_title="Session")) + await scoped.transcripts.append("session", {"event_id": "event"}) + + activity_path = scoped.paths.session_dir("session") / ".advanced-memory-activity" + os.utime(activity_path, (1.0, 1.0)) + + assert await scoped.session_memory.read("session") is None + assert scoped.paths.transcript_path("session").exists() + records = await scoped.transcripts.read_all("session") + assert len(records) == 1 + assert records[0]["event_id"] == "event" + await runtime.close() + + 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)) diff --git a/trpc_agent_sdk/advanced_memory/_autocompact.py b/trpc_agent_sdk/advanced_memory/_autocompact.py index 6c80f1ac2..070607009 100644 --- a/trpc_agent_sdk/advanced_memory/_autocompact.py +++ b/trpc_agent_sdk/advanced_memory/_autocompact.py @@ -292,9 +292,9 @@ def _summary_with_recovery_path(self, summary: str, session_id: str) -> str: """Append recovery paths for the full transcript and session memory.""" return (f"{summary.rstrip()}\n\n" "For exact content from before compaction, read the complete transcript: " - f"{self._runtime.paths.transcript_path(session_id)}\n" + f"{self._runtime.paths.storage_reference('transcript', session_id=session_id)}\n" "Current session memory: " - f"{self._runtime.paths.session_memory_path(session_id)}") + f"{self._runtime.paths.storage_reference('session_memory', session_id=session_id)}") def _find_signature_index( self, diff --git a/trpc_agent_sdk/advanced_memory/_config.py b/trpc_agent_sdk/advanced_memory/_config.py index 5be754e9e..35881cb64 100644 --- a/trpc_agent_sdk/advanced_memory/_config.py +++ b/trpc_agent_sdk/advanced_memory/_config.py @@ -175,6 +175,7 @@ class AdvancedMemoryConfig: preload_memory_max_topics: int = 5 preload_memory_max_chars: int = 50_000 preload_memory_candidate_limit: int = 200 + session_ttl_delete_transcripts: bool = False def __post_init__(self) -> None: """Validate the configuration and normalize the root directory.""" diff --git a/trpc_agent_sdk/advanced_memory/_memory_context.py b/trpc_agent_sdk/advanced_memory/_memory_context.py index a3ca94b89..1c13bef6e 100644 --- a/trpc_agent_sdk/advanced_memory/_memory_context.py +++ b/trpc_agent_sdk/advanced_memory/_memory_context.py @@ -79,9 +79,9 @@ async def apply(self, request: "LlmRequest", ctx: "InvocationContext | None" = N "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: " - f"{runtime.paths.memory_dir if config.storage_backend == 'local' else 'Redis'}\n" + f"{runtime.paths.memory_dir if config.storage_backend == 'local' else config.storage_backend.upper()}\n" f"Index file: " - f"{runtime.paths.memory_index_path if config.storage_backend == 'local' else 'Redis memory index'}\n" + f"{runtime.paths.storage_reference('memory_index')}\n" f"\n{index.rstrip()}\n\n" f"") request.append_instructions([instruction]) diff --git a/trpc_agent_sdk/advanced_memory/_paths.py b/trpc_agent_sdk/advanced_memory/_paths.py index 4cf5eaca3..f68646127 100644 --- a/trpc_agent_sdk/advanced_memory/_paths.py +++ b/trpc_agent_sdk/advanced_memory/_paths.py @@ -129,6 +129,73 @@ def tool_result_path(self, session_id: str, result_id: str) -> Path: 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 storage_reference( + self, + resource: str, + *, + session_id: str | None = None, + topic_name: str | None = None, + result_id: str | None = None, + ) -> str: + """Return a model-visible reference for a stored Advanced Memory resource.""" + if resource == "memory_index": + local_path = self.memory_index_path + elif resource == "memory_topic": + if topic_name is None: + raise ValueError("topic_name is required for a memory topic reference") + local_path = self.memory_topic_path(topic_name) + elif resource == "transcript": + if session_id is None: + raise ValueError("session_id is required for a transcript reference") + local_path = self.transcript_path(session_id) + elif resource == "session_memory": + if session_id is None: + raise ValueError("session_id is required for a session memory reference") + local_path = self.session_memory_path(session_id) + elif resource == "tool_result": + if session_id is None or result_id is None: + raise ValueError("session_id and result_id are required for a tool result reference") + local_path = self.tool_result_path(session_id, result_id) + else: + raise ValueError(f"Unknown Advanced Memory resource: {resource}") + if self.config.storage_backend == "local": + return str(local_path) + if self.scope is None: + raise ValueError("A scoped path is required for non-local memory storage") + + app_component = self.tenant_root_dir.parent.name + user_component = self.tenant_root_dir.name + if self.config.storage_backend == "redis": + user_base = f"{self.config.redis_key_prefix}:{{{app_component}:{user_component}}}" + if resource == "memory_index": + key = f"{user_base}:memory:index" + elif resource == "memory_topic": + key = f"{user_base}:memory:topic:{local_path.name}" + else: + safe_session_id = self.session_dir(session_id or "").name + session_base = f"{self.config.redis_key_prefix}:{{{app_component}:{user_component}:{safe_session_id}}}" + if resource == "transcript": + key = f"{session_base}:transcript" + elif resource == "session_memory": + key = f"{session_base}:summary" + else: + key = f"{session_base}:tool:{result_id}" + return f"advanced-memory://redis/{key}" + + app_name = self.scope.app_name + user_id = self.scope.user_id + if resource == "memory_index": + suffix = "memory/index" + elif resource == "memory_topic": + suffix = f"memory/topic/{local_path.name}" + elif resource == "transcript": + suffix = f"{session_id}/transcript" + elif resource == "session_memory": + suffix = f"{session_id}/summary" + else: + suffix = f"{session_id}/tool/{self.tool_result_path(session_id or '', result_id or '').stem}" + return f"advanced-memory://sql/{app_name}/{user_id}/{suffix}" + def ensure_base_directories(self) -> None: """Create the long-term and session memory directories.""" self.memory_dir.mkdir(parents=True, exist_ok=True) diff --git a/trpc_agent_sdk/advanced_memory/_redis_stores.py b/trpc_agent_sdk/advanced_memory/_redis_stores.py index 3bf2bd763..997a64761 100644 --- a/trpc_agent_sdk/advanced_memory/_redis_stores.py +++ b/trpc_agent_sdk/advanced_memory/_redis_stores.py @@ -104,10 +104,11 @@ async def _memory_write_lock(self): ) async def _refresh_ttl_group( - self, - registry: str, - keys: list[str], - ttl: int | None, + 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: @@ -118,15 +119,19 @@ async def _refresh_ttl_group( tracked_keys = {self._text(value) for value in tracked} tracked_keys.update(keys) for key in tracked_keys: - if key: + if key and not key.startswith(skip_prefixes): await self._command("expire", key, ttl) await self._command("expire", registry, ttl) async def _refresh_session_ttl(self, session_id: str, *keys: str) -> None: + skip_prefixes: tuple[str, ...] = () + if not self._config.session_ttl_delete_transcripts: + skip_prefixes = (f"{self._session_base(session_id)}:transcript", ) await self._refresh_ttl_group( self._session_registry(session_id), list(keys), self._config.session_ttl_seconds, + skip_prefixes=skip_prefixes, ) async def _refresh_memory_ttl(self, *keys: str) -> None: diff --git a/trpc_agent_sdk/advanced_memory/_sql_stores.py b/trpc_agent_sdk/advanced_memory/_sql_stores.py index a49a7d6cc..46af4de11 100644 --- a/trpc_agent_sdk/advanced_memory/_sql_stores.py +++ b/trpc_agent_sdk/advanced_memory/_sql_stores.py @@ -157,10 +157,14 @@ async def _refresh_session_scope(self, db: Any, session_id: str) -> None: return tables = ( (SqlSessionMemory, (self._app_name, self._user_id, session_id)), - (SqlTranscript, (self._app_name, self._user_id, session_id)), - (SqlTranscriptSeen, (self._app_name, self._user_id, session_id)), (SqlToolResult, (self._app_name, self._user_id, session_id)), ) + if self._config.session_ttl_delete_transcripts: + tables = ( + (SqlTranscript, (self._app_name, self._user_id, session_id)), + (SqlTranscriptSeen, (self._app_name, self._user_id, session_id)), + *tables, + ) for model, key in tables: rows = await self._storage.query( db, @@ -404,7 +408,8 @@ async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: session_id=session_id, record_id=uuid.uuid4().hex, payload=json.dumps(payload, ensure_ascii=False), - expires_at=self._expiry(self._config.session_ttl_seconds), + expires_at=(self._expiry(self._config.session_ttl_seconds) + if self._config.session_ttl_delete_transcripts else None), )) await self._refresh_session_scope(db, session_id) await self._storage.commit(db) @@ -450,7 +455,8 @@ async def append_unique( session_id=seen_key[2], unique_key=seen_key[3], unique_value=seen_key[4], - expires_at=self._expiry(self._config.session_ttl_seconds), + expires_at=(self._expiry(self._config.session_ttl_seconds) + if self._config.session_ttl_delete_transcripts else None), )) await self._storage.add( db, @@ -460,7 +466,8 @@ async def append_unique( session_id=session_id, record_id=uuid.uuid4().hex, payload=json.dumps(payload, ensure_ascii=False), - expires_at=self._expiry(self._config.session_ttl_seconds), + expires_at=(self._expiry(self._config.session_ttl_seconds) + if self._config.session_ttl_delete_transcripts else None), )) await self._refresh_session_scope(db, session_id) await self._storage.commit(db) @@ -514,7 +521,9 @@ async def start(self) -> None: 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: + models = self._models if self._config.session_ttl_delete_transcripts else tuple( + model for model in self._models if model is not SqlTranscript) + for model in models: await self._storage.delete( db, SqlKey(key=tuple(), storage_cls=model), diff --git a/trpc_agent_sdk/advanced_memory/_storage.py b/trpc_agent_sdk/advanced_memory/_storage.py index 0a8a011de..c390c025f 100644 --- a/trpc_agent_sdk/advanced_memory/_storage.py +++ b/trpc_agent_sdk/advanced_memory/_storage.py @@ -89,7 +89,17 @@ def _expire_session_dir(session_dir: Path, config: AdvancedMemoryConfig) -> bool expired = bool(files) and time.time() - max(path.stat().st_mtime for path in files) >= config.session_ttl_seconds if expired: - shutil.rmtree(session_dir, ignore_errors=True) + if config.session_ttl_delete_transcripts: + shutil.rmtree(session_dir, ignore_errors=True) + else: + transcript_path = session_dir / config.transcript_name + for child in session_dir.iterdir(): + if child == transcript_path: + continue + if child.is_dir(): + shutil.rmtree(child, ignore_errors=True) + else: + child.unlink(missing_ok=True) return expired @@ -399,7 +409,8 @@ async def read_all(self, session_id: str) -> list[dict[str, Any]]: 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 _expire_session_dir(path.parent, self._config): + expired = _expire_session_dir(path.parent, self._config) + if expired and self._config.session_ttl_delete_transcripts: return [] if not path.exists(): return [] diff --git a/trpc_agent_sdk/advanced_memory/_tool_result_budget.py b/trpc_agent_sdk/advanced_memory/_tool_result_budget.py index 6bfc9d6d8..7181584e5 100644 --- a/trpc_agent_sdk/advanced_memory/_tool_result_budget.py +++ b/trpc_agent_sdk/advanced_memory/_tool_result_budget.py @@ -210,11 +210,17 @@ def _build_replacement( candidate: ToolResultCandidate, ) -> ToolResultReplacement: """Build a deterministic storage path and model-visible preview.""" - persisted_path = (Path(f"advanced-memory://{self._runtime.config.redis_key_prefix}/" - f"{self._runtime.scope.app_name}/{self._runtime.scope.user_id}/{session_id}/" - f"tool/{candidate.result_id}") if hasattr(self._runtime, "scope") - and self._runtime.config.storage_backend == "redis" else self._runtime.paths.tool_result_path( - session_id, candidate.result_id)) + persisted_path = Path( + self._runtime.paths.storage_reference( + "tool_result", + session_id=session_id, + result_id=candidate.result_id, + )) + persisted_path_text = str(persisted_path).replace( + "advanced-memory:/", + "advanced-memory://", + 1, + ) preview, truncated = _preview_text( candidate.serialized_result, self._runtime.config.tool_result_preview_chars, @@ -226,7 +232,7 @@ def _build_replacement( }, "persisted_output": { "message": "The tool result exceeded the context budget; the complete content was persisted.", - "path": str(persisted_path), + "path": persisted_path_text, "original_chars": candidate.original_size, "preview": preview, "truncated": truncated, diff --git a/trpc_agent_sdk/sessions/_advanced_memory_session_service.py b/trpc_agent_sdk/sessions/_advanced_memory_session_service.py index 8ad94f4d3..c134fb545 100644 --- a/trpc_agent_sdk/sessions/_advanced_memory_session_service.py +++ b/trpc_agent_sdk/sessions/_advanced_memory_session_service.py @@ -70,6 +70,7 @@ async def _write_session(self, session: Session) -> None: payload["state"] = extract_state_delta(session.state).session_state path = self._metadata_path(session.app_name, session.user_id, session.id) await asyncio.to_thread(self._write_json, path, payload, self._runtime.config.encoding) + await asyncio.to_thread(path.parent.joinpath(".advanced-memory-activity").touch, exist_ok=True) @staticmethod def _write_json(path: Path, payload: dict[str, Any], encoding: str) -> None: @@ -94,6 +95,7 @@ async def _read_session(self, app_name: str, user_id: str, session_id: str) -> S return None payload = await asyncio.to_thread(path.read_text, encoding=self._runtime.config.encoding) await asyncio.to_thread(path.touch) + await asyncio.to_thread(path.parent.joinpath(".advanced-memory-activity").touch, exist_ok=True) return Session.model_validate(json.loads(payload)) def _start_cleanup_task(self) -> None: @@ -133,7 +135,18 @@ def _cleanup_expired_sessions(self) -> None: for metadata_path in tenants_root.glob(f"*/*/{self._runtime.config.session_dir_name}/*/session.json"): try: if metadata_path.stat().st_mtime < cutoff: - shutil.rmtree(metadata_path.parent, ignore_errors=True) + session_dir = metadata_path.parent + if self._runtime.config.session_ttl_delete_transcripts: + shutil.rmtree(session_dir, ignore_errors=True) + else: + transcript_path = session_dir / self._runtime.config.transcript_name + for child in session_dir.iterdir(): + if child == transcript_path: + continue + if child.is_dir(): + shutil.rmtree(child, ignore_errors=True) + else: + child.unlink(missing_ok=True) except FileNotFoundError: continue diff --git a/trpc_agent_sdk/tools/_advanced_memory_tool.py b/trpc_agent_sdk/tools/_advanced_memory_tool.py index f302c6fdb..8d6fdbdbe 100644 --- a/trpc_agent_sdk/tools/_advanced_memory_tool.py +++ b/trpc_agent_sdk/tools/_advanced_memory_tool.py @@ -28,6 +28,11 @@ _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] = [] @@ -124,7 +129,7 @@ async def save_memory( 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, } @@ -153,10 +158,10 @@ async def read_memory(self, filename: str, tool_context: Any | None = None) -> d } async def list_memory_index(self, tool_context: Any | None = None) -> dict: - """Return the current long-term memory index and its disk path.""" + """Return the current long-term memory index and its storage reference.""" runtime = self._runtime_for_context(tool_context) return { - "index_path": str(runtime.paths.memory_index_path), + "index_path": _memory_index_reference(runtime), "index": await runtime.long_term_memory.read_index(), } From 2219dbfed2f3c14a874d1367feffac032664cefe Mon Sep 17 00:00:00 2001 From: congkechen Date: Thu, 10 Sep 2026 10:57:38 +0800 Subject: [PATCH 3/6] =?UTF-8?q?feature:=20advanced=20memory=20=E8=AE=B0?= =?UTF-8?q?=E5=BF=86=E4=B8=8E=E4=B8=8A=E4=B8=8B=E6=96=87=E9=83=A8=E5=88=86?= =?UTF-8?q?=E8=A7=A3=E8=80=A6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.zh_CN.md | 2 +- .../README.md | 295 ++------- .../run_agent.py | 46 +- .../.env | 6 +- .../README.md | 112 +--- .../run_agent.py | 31 +- .../.env | 4 - .../README.md | 48 +- .../run_agent.py | 27 +- .../.env | 10 + .../README.md | 109 ++++ .../agent/__init__.py | 5 + .../agent/agent.py | 39 ++ .../agent/config.py | 19 + .../agent/prompts.py | 10 + .../agent/tools.py | 11 + .../run_agent.py | 120 ++++ .../.env | 11 + .../README.md | 108 ++++ .../agent/__init__.py | 5 + .../agent/agent.py | 39 ++ .../agent/config.py | 19 + .../agent/prompts.py | 10 + .../agent/tools.py | 11 + .../run_agent.py | 117 ++++ .../test_advanced_memory_session_service.py | 255 -------- .../test_advanced_memory_tools.py | 6 +- tests/advanced_memory/test_memory_context.py | 98 +-- tests/advanced_memory/test_preload_memory.py | 8 +- tests/advanced_memory/test_redis_stores.py | 47 +- tests/advanced_memory/test_sql_stores.py | 49 +- tests/advanced_memory/test_storage.py | 40 +- .../compact}/test_autocompact.py | 121 +++- .../test_context_compression_integration.py | 452 ++++++++++++++ .../compact}/test_coordination.py | 2 +- .../compact}/test_history_snip.py | 28 +- .../compact}/test_microcompact.py | 16 +- .../compact}/test_session_memory_extractor.py | 18 +- .../compact/test_session_memory_state.py | 160 +++++ .../compact}/test_token_budget.py | 14 +- .../compact}/test_tool_result_budget.py | 14 +- .../test_transcript_session_service.py | 31 +- .../session_memory_summary_diff_report.json | 12 +- .../test_in_memory_session_service.py | 46 ++ tests/sessions/test_redis_session_service.py | 84 +++ tests/sessions/test_sql_session_service.py | 44 ++ trpc_agent_sdk/abc/_session_service.py | 13 + trpc_agent_sdk/advanced_memory/__init__.py | 119 +--- .../advanced_memory/_integration.py | 186 ++---- .../advanced_memory/_memory_context.py | 4 +- .../advanced_memory/_preload_memory.py | 7 +- .../advanced_memory/_redis_stores.py | 305 +--------- trpc_agent_sdk/advanced_memory/_sql_stores.py | 570 +----------------- trpc_agent_sdk/advanced_memory/_storage.py | 500 +-------------- .../advanced_memory/_storage_backend.py | 4 +- .../evaluation/_eval_session_service.py | 33 + trpc_agent_sdk/memory/__init__.py | 8 +- .../memory/_advanced_memory_service.py | 68 +-- trpc_agent_sdk/runners.py | 7 +- trpc_agent_sdk/sessions/__init__.py | 22 +- .../_advanced_memory_session_service.py | 440 -------------- .../sessions/_base_session_service.py | 73 ++- .../sessions/_in_memory_session_service.py | 41 +- .../sessions/_redis_session_service.py | 112 +++- trpc_agent_sdk/sessions/_session.py | 50 ++ .../sessions/_sql_session_service.py | 59 +- trpc_agent_sdk/sessions/compact/__init__.py | 116 ++++ .../compact}/_autocompact.py | 249 ++++++-- .../sessions/compact/_base_config.py | 29 + .../sessions/compact/_base_manager.py | 57 ++ .../compact}/_callbacks.py | 0 .../compact}/_config.py | 20 +- .../compact}/_coordination.py | 0 .../compact}/_formats.py | 54 ++ .../compact}/_history_snip.py | 0 .../sessions/compact/_integration.py | 155 +++++ trpc_agent_sdk/sessions/compact/_manager.py | 102 ++++ .../compact}/_microcompact.py | 0 .../compact}/_paths.py | 12 +- .../sessions/compact/_redis_stores.py | 297 +++++++++ .../compact}/_runtime.py | 52 +- .../compact}/_session_memory.py | 148 ++++- .../compact}/_session_service.py | 20 + .../sessions/compact/_sql_stores.py | 528 ++++++++++++++++ trpc_agent_sdk/sessions/compact/_storage.py | 499 +++++++++++++++ .../compact}/_token_budget.py | 4 + .../compact}/_tool_result_budget.py | 0 .../compact}/_transcript.py | 0 trpc_agent_sdk/tools/_advanced_memory_tool.py | 12 +- 89 files changed, 4683 insertions(+), 3051 deletions(-) create mode 100644 examples/session_service_with_advanced_memory_redis/.env create mode 100644 examples/session_service_with_advanced_memory_redis/README.md create mode 100644 examples/session_service_with_advanced_memory_redis/agent/__init__.py create mode 100644 examples/session_service_with_advanced_memory_redis/agent/agent.py create mode 100644 examples/session_service_with_advanced_memory_redis/agent/config.py create mode 100644 examples/session_service_with_advanced_memory_redis/agent/prompts.py create mode 100644 examples/session_service_with_advanced_memory_redis/agent/tools.py create mode 100644 examples/session_service_with_advanced_memory_redis/run_agent.py create mode 100644 examples/session_service_with_advanced_memory_sql/.env create mode 100644 examples/session_service_with_advanced_memory_sql/README.md create mode 100644 examples/session_service_with_advanced_memory_sql/agent/__init__.py create mode 100644 examples/session_service_with_advanced_memory_sql/agent/agent.py create mode 100644 examples/session_service_with_advanced_memory_sql/agent/config.py create mode 100644 examples/session_service_with_advanced_memory_sql/agent/prompts.py create mode 100644 examples/session_service_with_advanced_memory_sql/agent/tools.py create mode 100644 examples/session_service_with_advanced_memory_sql/run_agent.py delete mode 100644 tests/advanced_memory/test_advanced_memory_session_service.py rename tests/{advanced_memory => sessions/compact}/test_autocompact.py (79%) create mode 100644 tests/sessions/compact/test_context_compression_integration.py rename tests/{advanced_memory => sessions/compact}/test_coordination.py (93%) rename tests/{advanced_memory => sessions/compact}/test_history_snip.py (90%) rename tests/{advanced_memory => sessions/compact}/test_microcompact.py (92%) rename tests/{advanced_memory => sessions/compact}/test_session_memory_extractor.py (97%) create mode 100644 tests/sessions/compact/test_session_memory_state.py rename tests/{advanced_memory => sessions/compact}/test_token_budget.py (89%) rename tests/{advanced_memory => sessions/compact}/test_tool_result_budget.py (96%) rename tests/{advanced_memory => sessions/compact}/test_transcript_session_service.py (78%) delete mode 100644 trpc_agent_sdk/sessions/_advanced_memory_session_service.py create mode 100644 trpc_agent_sdk/sessions/compact/__init__.py rename trpc_agent_sdk/{advanced_memory => sessions/compact}/_autocompact.py (75%) create mode 100644 trpc_agent_sdk/sessions/compact/_base_config.py create mode 100644 trpc_agent_sdk/sessions/compact/_base_manager.py rename trpc_agent_sdk/{advanced_memory => sessions/compact}/_callbacks.py (100%) rename trpc_agent_sdk/{advanced_memory => sessions/compact}/_config.py (94%) rename trpc_agent_sdk/{advanced_memory => sessions/compact}/_coordination.py (100%) rename trpc_agent_sdk/{advanced_memory => sessions/compact}/_formats.py (75%) rename trpc_agent_sdk/{advanced_memory => sessions/compact}/_history_snip.py (100%) create mode 100644 trpc_agent_sdk/sessions/compact/_integration.py create mode 100644 trpc_agent_sdk/sessions/compact/_manager.py rename trpc_agent_sdk/{advanced_memory => sessions/compact}/_microcompact.py (100%) rename trpc_agent_sdk/{advanced_memory => sessions/compact}/_paths.py (96%) create mode 100644 trpc_agent_sdk/sessions/compact/_redis_stores.py rename trpc_agent_sdk/{advanced_memory => sessions/compact}/_runtime.py (87%) rename trpc_agent_sdk/{advanced_memory => sessions/compact}/_session_memory.py (82%) rename trpc_agent_sdk/{advanced_memory => sessions/compact}/_session_service.py (91%) create mode 100644 trpc_agent_sdk/sessions/compact/_sql_stores.py create mode 100644 trpc_agent_sdk/sessions/compact/_storage.py rename trpc_agent_sdk/{advanced_memory => sessions/compact}/_token_budget.py (98%) rename trpc_agent_sdk/{advanced_memory => sessions/compact}/_tool_result_budget.py (100%) rename trpc_agent_sdk/{advanced_memory => sessions/compact}/_transcript.py (100%) 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/README.md b/examples/memory_service_with_advanced_memory/README.md index fa3d0bcb4..dd58e491a 100644 --- a/examples/memory_service_with_advanced_memory/README.md +++ b/examples/memory_service_with_advanced_memory/README.md @@ -1,277 +1,66 @@ -# Advanced Memory +# Standard SessionService + Advanced Compact + Advanced Memory -## Advanced Memory 简介 +本示例使用统一后的组合方式: -`Advanced Memory` 是一套面向 Agent 的本地化记忆与上下文管理机制,重点增强 Agent 在长期信息沉淀和超长对话处理方面的能力: - -- **更强的长期记忆能力**:支持将对话中的稳定事实、用户偏好和重要经验主动沉淀为可组织、可更新、可跨 Session 使用的长期记忆,而不是简单堆积历史消息。 -- **分层记忆管理**:分别管理原始对话、Session 级记忆和跨 Session 长期记忆,让不同类型的信息以合适的粒度参与后续推理。 -- **上下文管理**:根据上下文规模、信息类型和使用情况,对历史消息、工具结果及记忆内容进行统一治理,在保留关键信息的同时控制模型输入规模。 -- **上下文压缩**:支持对历史上下文和工具结果进行渐进式裁剪、压缩和摘要,降低长对话导致的上下文膨胀以及超出模型窗口限制的风险。 -- **结构化记忆提取**:从持续增长的对话中提取结构化信息,形成更稳定、更易维护的Session Memory,提升后续对话对历史信息的利用效率。 -- **本地化持久存储**:记忆和上下文数据以本地文件形式持久化,存储位置、数据边界和组织方式清晰可控,适合本地开发、调试、迁移和审计。 - -本示例演示如何使用 `AdvancedMemorySessionService`。它把 Session 持久化和Advanced Memory 上下文管理整合到一个 SessionService 中,用户不需要显式调用`setup_advanced_memory()`,也不需要再创建 `InMemorySessionService`。 - -**Advanced Memory 在 Redis 存储:** -[Redis `run_agent.py`](../memory_service_with_advanced_memory_redis/run_agent.py) - -**Advanced Memory 在 SQL 存储:** -[SQL `run_agent.py`](../memory_service_with_advanced_memory_sql/run_agent.py) - -## 示例流程 - -脚本使用同一个 Runner 执行多个 Session: +```text +InMemorySessionService +└── AdvancedSessionCompactManager + ├── Session Memory + ├── Tool Result Budget + ├── History Snip + ├── Microcompact + └── AutoCompact + +AdvancedMemoryService +├── save_memory +├── read_memory +├── list_memory_index +└── long-term memory injection +``` -1. `session-1` 连续输入多轮 Python 开发偏好。 -2. 当累计上下文和工具调用达到配置阈值后,系统会提取 session memory,并写入 - `session_memory.md`。 -3. `session-1` 请求总结已经学习到的开发偏好。 -4. `session-2` 查询长期记忆,验证不同 Session 共享同一个 `MEMORY/`。 +不再使用独立的 Advanced SessionService。Session 的创建、Event 保存和状态管理始终 +由标准 `InMemorySessionService`、`RedisSessionService` 或 `SqlSessionService` +负责;Advanced Compact 通过 `BaseSessionCompactManager` 生命周期接入。 -## 使用方式 +## 核心组装 ```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 +config = AdvancedCompactConfig( + root_dir=Path(__file__).resolve().parent, +) -session_service = AdvancedMemorySessionService( - config=AdvancedMemoryConfig( - root_dir=Path(__file__).resolve().parent, - memory_ttl_seconds=120, - session_ttl_seconds=60, - memory_focus_instruction=( - "特别关注并主动记住用户长期稳定的兴趣爱好、" - "编程语言偏好、开发习惯和测试习惯。" - ), - ) +session_service = InMemorySessionService( + session_config=SessionServiceConfig( + store_historical_events=True, + ), +) +compact_manager = setup_advanced_session_compact( + agent, + session_service, + config, ) +memory_service = AdvancedMemoryService(runtime=compact_manager.runtime) runner = Runner( app_name="advanced_memory_demo", agent=agent, session_service=session_service, - defer_post_turn_processing=True, # True 时开启,后台线程异步执行子 Agent 摘要 + memory_service=memory_service, ) ``` -`Runner` 检测到 `AdvancedMemorySessionService` 后会自动完成 Advanced Memory -绑定,包括: - -- transcript 持久化 -- session memory 提取 -- 长期记忆 tools:`save_memory`、`read_memory`、`list_memory_index` -- `HistorySnip` -- `Microcompact` -- `AutoCompact` -- `ToolResultBudget` - -`AdvancedMemoryConfig` 默认已经启用这些能力,本示例直接使用默认配置。 - -## 不同存储后端的 SessionService 选择 - -`AdvancedMemorySessionService` 是本地文件版 SessionService。使用 Redis 或 SQL 时,不要继续使用它,否则可能形成 Session 数据与 Advanced Memory 数据分开存储的混合模式。 - -推荐组合: - -- local:`AdvancedMemorySessionService` -- Redis:`RedisSessionService` + `AdvancedMemoryService` -- SQL:`SqlSessionService` + `AdvancedMemoryService` - -Redis 和 SQL 的完整示例分别见: - -- [Advanced Memory Redis 示例](../memory_service_with_advanced_memory_redis/README.md) -- [Advanced Memory SQL 示例](../memory_service_with_advanced_memory_sql/README.md) - -## 数据目录 - -运行后,数据默认写入当前示例目录: - -```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 -``` - -其中: +Session Compact 与 Advanced Memory 可以共享一个 Runtime;Runtime 的 `close()` +支持幂等调用,因此两个 Service 的正常关闭流程不会造成重复释放错误。 -- `session.json` 保存 Session 元数据和状态,不保存完整 Events。 -- `transcript.jsonl` 是追加写入的原始事件日志,可用于恢复 Session。 -- `session_memory.md` 是根据 transcript 提取的结构化摘要。 -- `MEMORY/` 保存跨 Session 使用的长期记忆。 +也可以直接构造实现了 `BaseSessionCompactManager` 的自定义 Manager,并通过 +`session_compact_manager=` 注入标准 SessionService。 ## 运行 -先在本目录创建 `.env`,然后填写模型配置: +在 `.env` 中配置模型,然后执行: ```bash -cd examples/memory_service_with_advanced_memory -python3 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) -- `M_TTL`(可选,长期 memory 过期时间,单位为秒) -- `SESSION_TTL`(可选,session 相关数据过期时间,单位为秒) - -`M_TTL` 和 `SESSION_TTL` 未配置时不会自动删除数据。Session 的后台清理检查间隔由示例内部设置,不需要单独配置。 - -默认情况下,Session TTL 过期会保留 transcript,便于审计;只有将 -`session_ttl_delete_transcripts=True` 时,transcript 才会随 Session TTL 一起删除。 - -本示例提供的 `.env` 默认使用 `M_TTL=120` 和 `SESSION_TTL=60`,方便直接观察过期清理;如果不希望自动删除,将这两个值留空即可。 - -`.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` 即可**;其中 TTL 和记忆重点使用本示例的演示值。 - -```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 - - # TTL(单位:秒;None 表示不过期) - memory_ttl_seconds=120, # 长期记忆 TTL(秒) - session_ttl_seconds=60, # 会话记忆 TTL(秒) - session_ttl_delete_transcripts=False, # Session TTL 是否删除 transcript - - # 长期记忆 - memory_index_max_lines=200, # 注入 prompt 的索引最大行数 - memory_index_max_bytes=25_000, # 注入 prompt 的索引最大字节数 - long_term_memory_injection_enabled=True, # 是否注入 MEMORY.md - memory_focus_instruction=( # 可选:重点记忆要求 - "特别关注并主动记住用户长期稳定的兴趣爱好、" - "编程语言偏好、开发习惯和测试习惯。" - ), - - # 工具结果 - 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 数 - ), -) -``` - -`memory_focus_instruction` 可以传入应用级的自定义记忆偏好,例如: - -```python -memory_focus_instruction="特别关注用户长期稳定的兴趣爱好和开发习惯。" -``` - -它会追加到长期记忆的 system instruction 中,提示模型优先关注这些内容。 - -本示例还会把同一个 `SESSION_TTL` 传给 `SessionServiceConfig`,用于清理`session.json` 和 Session 目录;`cleanup_interval_seconds=5` 只是内部检查频率,不是另一个需要用户配置的 TTL: - -```python -session_config = SessionServiceConfig( - ttl=SessionServiceConfig.create_ttl_config( - enable=True, - ttl_seconds=60, # SESSION_TTL - cleanup_interval_seconds=5, # 内部检查频率 - ) -) +python run_agent.py ``` -`preload_memory_model` 不是 `AdvancedMemoryConfig` 字段,而是 -`AdvancedMemorySessionService` 的可选参数,用于指定轻量筛选模型: - -```python -session_service = AdvancedMemorySessionService( - config=AdvancedMemoryConfig(preload_memory_enabled=True), - preload_memory_model=small_model, # 不传时复用主 Agent 的模型 -) -``` +示例会在两个 Session 中使用同一用户,验证用户级长期记忆可以跨 Session 使用。 diff --git a/examples/memory_service_with_advanced_memory/run_agent.py b/examples/memory_service_with_advanced_memory/run_agent.py index 0f2d564dd..17e8298cb 100644 --- a/examples/memory_service_with_advanced_memory/run_agent.py +++ b/examples/memory_service_with_advanced_memory/run_agent.py @@ -12,9 +12,11 @@ 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.advanced_memory import AdvancedCompactConfig +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 setup_advanced_session_compact from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import Part @@ -23,25 +25,34 @@ load_dotenv(Path(__file__).with_name(".env")) -def create_session_service() -> AdvancedMemorySessionService: - """Create the persistent Advanced Memory session service.""" +def create_services(agent) -> tuple[InMemorySessionService, AdvancedMemoryService]: + """Create standard Session storage with Advanced Compact and Memory.""" memory_ttl = os.getenv("M_TTL") session_ttl = os.getenv("SESSION_TTL") session_ttl_seconds = int(session_ttl) if session_ttl else 0 - return AdvancedMemorySessionService( - config=AdvancedMemoryConfig( - root_dir=Path(__file__).resolve().parent, - memory_ttl_seconds=int(memory_ttl) if memory_ttl else None, - session_ttl_seconds=session_ttl_seconds or None, - memory_focus_instruction=("特别关注并主动记住用户长期稳定的兴趣爱好、" - "编程语言偏好、开发习惯和测试习惯。"), + config = AdvancedCompactConfig( + root_dir=Path(__file__).resolve().parent, + memory_ttl_seconds=int(memory_ttl) if memory_ttl else None, + session_ttl_seconds=session_ttl_seconds or None, + memory_focus_instruction=("特别关注并主动记住用户长期稳定的兴趣爱好、" + "编程语言偏好、开发习惯和测试习惯。"), + ) + session_service = InMemorySessionService( + session_config=SessionServiceConfig( + ttl=SessionServiceConfig.create_ttl_config( + enable=bool(session_ttl), + ttl_seconds=session_ttl_seconds, + cleanup_interval_seconds=5, + ), + store_historical_events=True, ), - session_config=SessionServiceConfig(ttl=SessionServiceConfig.create_ttl_config( - enable=bool(session_ttl), - ttl_seconds=session_ttl_seconds, - cleanup_interval_seconds=5, - ), ), ) + compact_manager = setup_advanced_session_compact( + agent, + session_service, + config, + ) + return session_service, AdvancedMemoryService(runtime=compact_manager.runtime) async def run_turn(runner, *, user_id: str, session_id: str, prompt: str) -> None: @@ -67,13 +78,14 @@ async def run_turn(runner, *, user_id: str, session_id: str, prompt: str) -> Non async def main() -> None: """Run two independent sessions sharing Advanced Memory.""" agent = create_agent() - session_service = create_session_service() + session_service, memory_service = create_services(agent) from trpc_agent_sdk.runners import Runner runner = Runner( app_name="advanced_memory_demo", agent=agent, session_service=session_service, + memory_service=memory_service, ) memory_ttl = os.getenv("M_TTL") memory_ttl_seconds = int(memory_ttl) if memory_ttl else 0 diff --git a/examples/memory_service_with_advanced_memory_redis/.env b/examples/memory_service_with_advanced_memory_redis/.env index 3982337a7..6a46edc83 100644 --- a/examples/memory_service_with_advanced_memory_redis/.env +++ b/examples/memory_service_with_advanced_memory_redis/.env @@ -4,7 +4,5 @@ REDIS_URL= 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= \ No newline at end of file + +M_TTL=120 \ 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 index c65167b29..c97328e13 100644 --- a/examples/memory_service_with_advanced_memory_redis/README.md +++ b/examples/memory_service_with_advanced_memory_redis/README.md @@ -2,20 +2,20 @@ 本示例演示如何将 Advanced Memory 的本地文件存储切换为 Redis,并验证: -- Redis:`RedisSessionService` + `AdvancedMemoryService` +- Redis:`AdvancedMemoryService(storage_backend="redis")` - 长期 memory 可以跨 Python 进程持久化; - 同一用户在不同 `session_id` 中可以读取自己的长期 memory; - session 相关数据和长期 memory 可以分别设置 TTL; - Redis 中的 Markdown、Stream 和索引数据如何组织。 -示例使用两个服务: +本示例只关注长期 Memory 的 Redis 持久化: ```text -RedisSessionService -└── 保存 Session、app state、user state +AdvancedMemoryService +└── Redis 保存长期 memory index 和 topic -AdvancedMemoryService(storage_backend="redis") -└── 保存长期 memory、session memory、transcript、tool result +Runner +└── InMemorySessionService(仅用于运行示例) ``` ## 环境要求 @@ -115,17 +115,13 @@ REDIS_URL=redis://localhost:6379/0 # 长期 memory 的 TTL,单位为秒 M_TTL=120 -# 所有 session 相关内容的 TTL,单位为秒 -SESSION_TTL=60 ``` TTL 规则: - `M_TTL` 管理用户级长期 memory 的全部 Redis key; -- `SESSION_TTL` 管理 session memory、transcript、tool result、去重 key; -- `SESSION_TTL` 也传给 `RedisSessionService`,用于 Session 和 state; - TTL 会在访问或写入时刷新,是“最后一次活动后过期”; -- 两个 TTL 必须设置为大于 0 的整数。 +- `M_TTL` 必须设置为大于 0 的整数。 更多 Advanced Memory 配置请参考[Advanced Memory README](../memory_service_with_advanced_memory/README.md)。 @@ -181,41 +177,27 @@ Redis 版本最核心的构建过程可以简化为三步: redis_url = "redis://:password@localhost:6379/0" memory_service = AdvancedMemoryService( - AdvancedMemoryConfig( + AdvancedCompactConfig( storage_backend="redis", redis_url=redis_url, memory_ttl_seconds=120, # from M_TTL; omit to disable expiration - session_ttl_seconds=60, # from SESSION_TTL; omit to disable expiration ) ) -session_config = SessionServiceConfig( - ttl=SessionServiceConfig.create_ttl_config( - enable=True, - ttl_seconds=60, # same value as SESSION_TTL - cleanup_interval_seconds=60, - ) -) -session_service = RedisSessionService( - db_url=redis_url, - is_async=True, - session_config=session_config, -) - runner = Runner( app_name="advanced-memory-redis-demo", agent=create_agent(), - session_service=session_service, + session_service=InMemorySessionService(), memory_service=memory_service, ) ``` 其中: -- 用户只需要配置 `M_TTL` 和 `SESSION_TTL` 两个 TTL; -- `AdvancedMemoryService` 负责长期 memory、session memory、transcript 和 tool result; -- `RedisSessionService` 负责框架 Session、app state 和 user state; -- `Runner` 将 Agent、Session Service 和 Memory Service 组合起来; +- 用户只需要配置长期 Memory 的 `M_TTL`; +- `AdvancedMemoryService` 只负责长期 memory; +- Session Service 的 Redis 高级压缩接入请看 + [`session_service_with_advanced_memory_redis`](../session_service_with_advanced_memory_redis/); - 运行请求时通过 `user_id` 和 `session_id` 指定用户及会话。 ## 运行结果(实测) @@ -297,14 +279,6 @@ TTL "advanced-memory-redis-demo:v1:{advanced-memory-redis-demo:redis-demo-user}: 预期接近 `120`。 -session transcript: - -```redis -TTL "advanced-memory-redis-demo:v1:{advanced-memory-redis-demo:redis-demo-user:redis-write-session}:transcript" -``` - -预期接近 `60`。 - TTL 含义: ```text @@ -313,16 +287,6 @@ TTL 含义: 大于 0 剩余秒数 ``` -观察 session key: - -```bash -docker exec advanced-memory-redis redis-cli --scan \ - --pattern 'advanced-memory-redis-demo:v1:*:summary' - -docker exec advanced-memory-redis redis-cli --scan \ - --pattern 'advanced-memory-redis-demo:v1:*:transcript*' -``` - ## 清理测试数据 只删除本示例的 Advanced Memory key: @@ -376,53 +340,3 @@ memory TTL registry: ``` 它记录该用户的所有长期 memory key,用于统一刷新 `M_TTL`。 - -### session memory - -本地文件概念: - -```text -SESSION/{session_id}/session_memory.md -``` - -Redis 映射: - -```text -{prefix}:{app:user:session}:summary -``` - -类型是 Redis String,内容是 Markdown。 - -### transcript - -本地文件概念: - -```text -SESSION/{session_id}/transcript.jsonl -``` - -Redis 映射: - -```text -{prefix}:{app:user:session}:transcript -``` - -类型是 Redis Stream,每条记录保存一份 JSON 数据。 - -### transcript 去重和 tool result - -```text -{prefix}:{app:user:session}:transcript:seen:{unique_key} -{prefix}:{app:user:session}:tool:{result_id} -``` - -去重 key 使用 Set,tool result 使用 String。 - -session TTL registry: - -```text -{prefix}:{app:user:session}:keys -``` - -它记录该 session 下的 summary、transcript、tool result 等 key,用于统一刷新 -`SESSION_TTL`,避免同一个 session 的不同内容出现 TTL 不一致。 diff --git a/examples/memory_service_with_advanced_memory_redis/run_agent.py b/examples/memory_service_with_advanced_memory_redis/run_agent.py index 5aafab3db..93dce8789 100644 --- a/examples/memory_service_with_advanced_memory_redis/run_agent.py +++ b/examples/memory_service_with_advanced_memory_redis/run_agent.py @@ -14,10 +14,10 @@ from dotenv import load_dotenv from agent.agent import create_agent -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig +from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig from trpc_agent_sdk.memory import AdvancedMemoryService from trpc_agent_sdk.runners import Runner -from trpc_agent_sdk.sessions import RedisSessionService, SessionServiceConfig +from trpc_agent_sdk.sessions import InMemorySessionService from trpc_agent_sdk.types import Content, Part load_dotenv(Path(__file__).with_name(".env")) @@ -61,34 +61,17 @@ def build_redis_url_from_environment() -> str: def create_advanced_memory_service(redis_url: str) -> AdvancedMemoryService: - """Create Advanced Memory backed by the configured Redis instance.""" + """Create the long-term Advanced Memory service backed by Redis.""" memory_ttl = os.getenv("M_TTL") - session_ttl = os.getenv("SESSION_TTL") - config = AdvancedMemoryConfig( + config = AdvancedCompactConfig( storage_backend="redis", redis_url=redis_url, redis_key_prefix="advanced-memory-redis-demo:v1", memory_ttl_seconds=int(memory_ttl) if memory_ttl else None, - session_ttl_seconds=int(session_ttl) if session_ttl else None, ) return AdvancedMemoryService(config) -def create_redis_session_service(redis_url: str) -> RedisSessionService: - """Create session storage with the Advanced Memory session TTL.""" - session_ttl = os.getenv("SESSION_TTL") - ttl_seconds = int(session_ttl) if session_ttl else 0 - return RedisSessionService( - db_url=redis_url, - is_async=True, - session_config=SessionServiceConfig(ttl=SessionServiceConfig.create_ttl_config( - enable=bool(session_ttl), - ttl_seconds=ttl_seconds, - cleanup_interval_seconds=ttl_seconds, - ), ), - ) - - 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}") @@ -112,13 +95,11 @@ 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() - memory_service = create_advanced_memory_service(redis_url) - session_service = create_redis_session_service(redis_url) runner = Runner( app_name=app_name, agent=create_agent(), - session_service=session_service, - memory_service=memory_service, + session_service=InMemorySessionService(), + memory_service=create_advanced_memory_service(redis_url), ) try: queries = RUNNER_A_QUERIES if phase == "write" else RUNNER_B_QUERIES diff --git a/examples/memory_service_with_advanced_memory_sql/.env b/examples/memory_service_with_advanced_memory_sql/.env index a617a519c..81dbccf4a 100644 --- a/examples/memory_service_with_advanced_memory_sql/.env +++ b/examples/memory_service_with_advanced_memory_sql/.env @@ -3,9 +3,6 @@ TRPC_AGENT_API_KEY= TRPC_AGENT_BASE_URL= TRPC_AGENT_MODEL_NAME= -TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS= -TRPC_AGENT_MAX_OUTPUT_TOKENS= - # 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 @@ -14,4 +11,3 @@ TRPC_AGENT_MAX_OUTPUT_TOKENS= SQL_URL= SQL_IS_ASYNC=true M_TTL=120 -SESSION_TTL=60 diff --git a/examples/memory_service_with_advanced_memory_sql/README.md b/examples/memory_service_with_advanced_memory_sql/README.md index 6eb268dd8..87dbac3f4 100644 --- a/examples/memory_service_with_advanced_memory_sql/README.md +++ b/examples/memory_service_with_advanced_memory_sql/README.md @@ -2,14 +2,14 @@ 本示例使用 SQL 保存 Advanced Memory,并验证同一用户的长期 memory 可以跨 Python 进程和不同 session 读取。 -- SQL:`SqlSessionService` + `AdvancedMemoryService` +- SQL:`AdvancedMemoryService(storage_backend="sql")` ```text -SqlSessionService -└── Session、app state、user state +AdvancedMemoryService +└── SQL 保存长期 memory index 和 topic -AdvancedMemoryService(storage_backend="sql") -└── 长期 memory、session memory、transcript、tool result +Runner +└── InMemorySessionService(仅用于运行示例) ``` ## 配置 @@ -37,8 +37,7 @@ TRPC_AGENT_BASE_URL=your-base-url TRPC_AGENT_MODEL_NAME=your-model-name ``` -`M_TTL` 默认控制长期 memory 的过期时间,`SESSION_TTL` 控制 session 相关内容的过期时间, -单位都是秒。 +`M_TTL` 控制长期 memory 的过期时间,单位为秒。 更多 Advanced Memory 配置请参考[Advanced Memory README](../memory_service_with_advanced_memory/README.md)。 @@ -88,42 +87,28 @@ SQL 版本最核心的构建过程可以简化为三步: sql_url = "mysql+aiomysql://user:password@localhost:3306/trpc_agent_advanced_memory" memory_service = AdvancedMemoryService( - AdvancedMemoryConfig( + AdvancedCompactConfig( storage_backend="sql", sql_url=sql_url, sql_is_async=True, memory_ttl_seconds=120, # from M_TTL; omit to disable expiration - session_ttl_seconds=60, # from SESSION_TTL; omit to disable expiration ) ) -session_config = SessionServiceConfig( - ttl=SessionServiceConfig.create_ttl_config( - enable=True, - ttl_seconds=60, # same value as SESSION_TTL - cleanup_interval_seconds=60, - ) -) -session_service = SqlSessionService( - db_url=sql_url, - is_async=True, - session_config=session_config, -) - runner = Runner( app_name="advanced-memory-sql-demo", agent=create_agent(), - session_service=session_service, + session_service=InMemorySessionService(), memory_service=memory_service, ) ``` 其中: -- 用户只需要配置 `M_TTL` 和 `SESSION_TTL` 两个 TTL; -- `AdvancedMemoryService` 负责长期 memory、session memory、transcript 和 tool result; -- `SqlSessionService` 负责框架 Session、app state 和 user state; -- `Runner` 将 Agent、Session Service 和 Memory Service 组合起来; +- 用户只需要配置长期 Memory 的 `M_TTL`; +- `AdvancedMemoryService` 只负责长期 memory; +- Session Service 的 SQL 高级压缩接入请看 + [`session_service_with_advanced_memory_sql`](../session_service_with_advanced_memory_sql/); - 运行请求时通过 `user_id` 和 `session_id` 指定用户及会话; - 多个节点只要使用相同的 SQL 数据库、`app_name` 和 `user_id`,就能访问同一份长期 memory。 @@ -183,12 +168,7 @@ Advanced Memory 使用独立的表,不复用原始 `SqlMemoryService` 的 `mem ```text advanced_memory_indexes advanced_memory_topics -advanced_memory_session_memory -advanced_memory_transcripts -advanced_memory_transcript_seen -advanced_memory_tool_results ``` -Markdown 内容保存在 `TEXT` 字段;transcript 保存 JSON 字符串; -`expires_at` 用于 SQL TTL。SQL 后端在读取时过滤过期数据,并在访问或写入时刷新 -同一用户或同一 session 下相关记录的过期时间。 +Markdown 内容保存在 `TEXT` 字段,`expires_at` 用于 Memory TTL。 +SQL 后端在读取时过滤过期数据,并在访问或写入时刷新同一用户的长期 Memory。 diff --git a/examples/memory_service_with_advanced_memory_sql/run_agent.py b/examples/memory_service_with_advanced_memory_sql/run_agent.py index 7c570be0c..6fdd3d2f1 100644 --- a/examples/memory_service_with_advanced_memory_sql/run_agent.py +++ b/examples/memory_service_with_advanced_memory_sql/run_agent.py @@ -14,10 +14,10 @@ from dotenv import load_dotenv from agent.agent import create_agent -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig +from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig from trpc_agent_sdk.memory import AdvancedMemoryService from trpc_agent_sdk.runners import Runner -from trpc_agent_sdk.sessions import SessionServiceConfig, SqlSessionService +from trpc_agent_sdk.sessions import InMemorySessionService from trpc_agent_sdk.types import Content, Part load_dotenv(Path(__file__).with_name(".env")) @@ -58,41 +58,24 @@ def sql_is_async() -> bool: def create_advanced_memory_service(sql_url: str) -> AdvancedMemoryService: - """Create Advanced Memory backed by SQL.""" + """Create the long-term Advanced Memory service backed by SQL.""" memory_ttl = os.getenv("M_TTL") - session_ttl = os.getenv("SESSION_TTL") - config = AdvancedMemoryConfig( + config = AdvancedCompactConfig( storage_backend="sql", sql_url=sql_url, sql_is_async=sql_is_async(), memory_ttl_seconds=int(memory_ttl) if memory_ttl else None, - session_ttl_seconds=int(session_ttl) if session_ttl else None, ) return AdvancedMemoryService(config) -def create_sql_session_service(sql_url: str) -> SqlSessionService: - """Create the SQL-backed framework session service.""" - session_ttl = os.getenv("SESSION_TTL") - ttl_seconds = int(session_ttl) if session_ttl else 0 - return SqlSessionService( - db_url=sql_url, - is_async=sql_is_async(), - session_config=SessionServiceConfig(ttl=SessionServiceConfig.create_ttl_config( - enable=bool(session_ttl), - ttl_seconds=ttl_seconds, - cleanup_interval_seconds=ttl_seconds, - ), ), - ) - - 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=create_sql_session_service(sql_url), + session_service=InMemorySessionService(), memory_service=create_advanced_memory_service(sql_url), ) try: 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..dff135499 --- /dev/null +++ b/examples/session_service_with_advanced_memory_redis/README.md @@ -0,0 +1,109 @@ +# Redis SessionService + Session Compact + +本示例只演示如何在已有 `RedisSessionService` 上增加: + +- Tool Result Budget +- History Snip +- Microcompact +- AutoCompact +- AutoCompact 触发时生成的 Session Memory + +压缩能力来自 `trpc_agent_sdk.sessions.compact`,长期记忆仍属于独立的 +`trpc_agent_sdk.advanced_memory`。 + +## 组装关系 + +```text +AdvancedCompactConfig + ↓ Runner 自动创建 +RedisSessionService +├── AdvancedSessionCompactManager +├── events: summary + recent Events +├── historical_events: 被压缩的原始 Events +└── state["_trpc_agent:summary"] + +AdvancedMemoryRuntime +├── 精简 compression transcript +└── 完整 Tool Result 旁路存储 +``` + +核心调用: + +```python +session_config = SessionServiceConfig( + store_historical_events=True, +) +compact_config = AdvancedCompactConfig( + redis_key_prefix="session-compression-demo:v1", + 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_config=compact_config, +) + +runner = Runner( + app_name=app_name, + agent=agent, + session_service=session_service, +) +``` + +`Runner` 会读取 `session_compact_config`,自动从 `RedisSessionService` 获取 URL 和 +异步模式,创建 `AdvancedSessionCompactManager` 并通过基类接口注入。 +用户不需要手动调用 `setup_advanced_session_compact`,也不需要直接创建 Manager。 + +## 兼容已有 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 和完整 Tool Result 都能跨进程恢复。 + +运行结束会输出 `Active Events`、`Historical Events`、活跃窗口是否以 summary +开头,以及 Session Memory state 是否存在,方便直接确认压缩是否触发。 + +## 存储职责 + +- `RedisSessionService`:Session、活跃 Events、historical Events、state 和 Session Memory。 +- Advanced Memory Redis stores:压缩重放记录和完整 Tool Result。 +- Redis transcript 不保存 `kind=event`,也不保存 `session-memory-checkpoint`。 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..4ae9f332d --- /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.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( + redis_key_prefix="session-compression-demo:v1", + 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() + session_config = SessionServiceConfig(store_historical_events=True) + session_service = RedisSessionService( + db_url=redis_url(), + is_async=True, + session_config=session_config, + session_compact_config=compact_config, + ) + 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..0809508e9 --- /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=root +MYSQL_PASSWORD= +MYSQL_HOST=127.0.0.1 +MYSQL_PORT=3306 +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..b50276df7 --- /dev/null +++ b/examples/session_service_with_advanced_memory_sql/README.md @@ -0,0 +1,108 @@ +# SQL SessionService + Session Compact + +本示例只演示如何在已有 `SqlSessionService` 上增加: + +- Tool Result Budget +- History Snip +- Microcompact +- AutoCompact +- AutoCompact 触发时生成的 Session Memory + +压缩能力来自 `trpc_agent_sdk.sessions.compact`,长期记忆仍属于独立的 +`trpc_agent_sdk.advanced_memory`。SQL 表结构不变,但活跃/历史 Event +会按原 Session 语义重新分区。 + +## 组装关系 + +```text +AdvancedCompactConfig + ↓ Runner 自动创建 +SqlSessionService +├── AdvancedSessionCompactManager +├── events: summary + recent Events +├── sessions.historical_events: 被压缩的原始 Events +└── sessions.state["_trpc_agent:summary"] + +AdvancedMemoryRuntime +├── advanced_memory_transcripts +├── advanced_memory_transcript_seen +└── advanced_memory_tool_results +``` + +核心调用: + +```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_config=compact_config, +) + +runner = Runner( + app_name=app_name, + agent=agent, + session_service=session_service, +) +``` + +`Runner` 会读取 `session_compact_config`,自动从 `SqlSessionService` 获取 URL 和异步 +模式,创建 `AdvancedSessionCompactManager` 并通过基类接口注入。 +用户不需要手动调用 `setup_advanced_session_compact`,也不需要直接创建 Manager。 + +## 兼容已有 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 和 Session Memory。 +- Advanced Memory SQL stores:压缩重放记录和完整 Tool Result。 +- 不再创建 `advanced_memory_session_memory` 表。 +- Advanced Memory transcript 不保存 `kind=event` 或 + `session-memory-checkpoint`。 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..7c4274bb2 --- /dev/null +++ b/examples/session_service_with_advanced_memory_sql/run_agent.py @@ -0,0 +1,117 @@ +#!/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.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() + session_config = SessionServiceConfig(store_historical_events=True) + session_service = SqlSessionService( + db_url=sql_url(), + is_async=False, + session_config=session_config, + session_compact_config=compact_config, + ) + 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 cd07c59e9..000000000 --- a/tests/advanced_memory/test_advanced_memory_session_service.py +++ /dev/null @@ -1,255 +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.for_session(session).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_same_session_id_is_isolated_between_users(tmp_path: Path) -> None: - """Allow matching IDs because each user owns a separate session directory.""" - service = AdvancedMemorySessionService(config=_config(tmp_path)) - first = await service.create_session( - app_name="demo-app", - user_id="user-a", - session_id="shared-session", - ) - second = await service.create_session( - app_name="demo-app", - user_id="user-b", - session_id="shared-session", - ) - await service.append_event(first, _event("event-a", "for user a")) - await service.append_event(second, _event("event-b", "for user b")) - - assert (await service.get_session(app_name="demo-app", user_id="user-a", - session_id="shared-session")).events[0].id == "event-a" - assert (await service.get_session(app_name="demo-app", user_id="user-b", - session_id="shared-session")).events[0].id == "event-b" - assert service.runtime.for_session(first).paths.session_dir( - first.id) != service.runtime.for_session(second).paths.session_dir(second.id) - - -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.for_session(session).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.for_session(session).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_ttl_cleanup_preserves_transcript_by_default(tmp_path: Path) -> None: - """Keep the transcript when session metadata expires.""" - 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="preserve-transcript", - ) - await service.append_event(session, _event("event-1", "hello")) - transcript_path = service.runtime.for_session(session).paths.transcript_path(session.id) - - await asyncio.sleep(1.1) - - assert await service.get_session( - app_name="demo-app", - user_id="demo-user", - session_id=session.id, - ) is None - assert transcript_path.exists() - 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.for_scope("demo-app", "demo-user").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 cf3347222..750777d1b 100644 --- a/tests/advanced_memory/test_advanced_memory_tools.py +++ b/tests/advanced_memory/test_advanced_memory_tools.py @@ -8,7 +8,7 @@ import pytest -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig +from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig from trpc_agent_sdk.advanced_memory import AdvancedMemoryPaths from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime from trpc_agent_sdk.tools import AdvancedMemoryTools @@ -17,7 +17,7 @@ def _runtime(tmp_path: Path) -> AdvancedMemoryRuntime: """Create a test runtime with long-term memory enabled.""" - return AdvancedMemoryRuntime.create(AdvancedMemoryConfig( + return AdvancedMemoryRuntime.create(AdvancedCompactConfig( enabled=True, root_dir=tmp_path, )).for_scope("demo-app", "demo-user") @@ -80,7 +80,7 @@ async def test_list_memory_index_reports_backend_storage_reference( expected_prefix: str, ) -> None: """Avoid exposing a local filesystem path for external memory stores.""" - config = AdvancedMemoryConfig( + config = AdvancedCompactConfig( 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, diff --git a/tests/advanced_memory/test_memory_context.py b/tests/advanced_memory/test_memory_context.py index 78c48a1a3..703ab3a96 100644 --- a/tests/advanced_memory/test_memory_context.py +++ b/tests/advanced_memory/test_memory_context.py @@ -7,21 +7,22 @@ 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 AdvancedCompactConfig 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.advanced_memory import setup_long_term_memory +from trpc_agent_sdk.memory import AdvancedMemoryService from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.sessions.compact import AutoCompactCallback +from trpc_agent_sdk.sessions.compact import HistorySnipCallback +from trpc_agent_sdk.sessions.compact import MicrocompactCallback +from trpc_agent_sdk.sessions.compact import setup_context_compression +from trpc_agent_sdk.sessions.compact import ToolResultBudgetCallback +from trpc_agent_sdk.sessions.compact._callbacks import install_staged_callback from trpc_agent_sdk.sessions import InMemorySessionService +from trpc_agent_sdk.sessions import SessionServiceConfig class FakeSummaryGenerator: @@ -35,7 +36,7 @@ async def generate(self, history: str, ctx) -> str: def _runtime(tmp_path: Path) -> AdvancedMemoryRuntime: """Create a test runtime with long-term memory injection enabled.""" - return AdvancedMemoryRuntime.create(AdvancedMemoryConfig( + return AdvancedMemoryRuntime.create(AdvancedCompactConfig( enabled=True, root_dir=tmp_path, )) @@ -107,7 +108,7 @@ async def test_long_term_memory_index_is_injected_once(tmp_path: Path) -> 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( - AdvancedMemoryConfig( + AdvancedCompactConfig( enabled=True, root_dir=tmp_path, memory_focus_instruction="重点记住用户长期稳定的兴趣爱好。", @@ -122,49 +123,47 @@ async def test_custom_memory_focus_is_injected_into_system_instruction(tmp_path: assert "重点记住用户长期稳定的兴趣爱好。" 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.""" +async def test_context_setup_installs_four_compaction_stages(tmp_path: Path) -> None: + """Ensure Session compact setup installs only the four compact stages.""" runtime = _runtime(tmp_path) agent = SimpleNamespace(before_model_callback=None) + session_service = InMemorySessionService( + session_config=SessionServiceConfig(store_historical_events=True), + ) - components = setup_context_management( + setup_context_compression( agent, + session_service, runtime, FakeSummaryGenerator(), ) - 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) + assert session_service.session_compact_manager.runtime is 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) + await session_service.close() -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_explicit_memory_and_compact_setup_compose(tmp_path: Path, ) -> None: + """Ensure long-term memory and Session compact are composed explicitly.""" runtime = _runtime(tmp_path) agent = SimpleNamespace(before_model_callback=None, tools=[]) - - first = setup_advanced_memory( - agent, - InMemorySessionService(), - runtime, - FakeSummaryGenerator(), + session_service = InMemorySessionService( + session_config=SessionServiceConfig(store_historical_events=True), ) - second = setup_advanced_memory( + long_term = setup_long_term_memory(agent, runtime) + compact = setup_context_compression( agent, - first.session_service, + session_service, runtime, FakeSummaryGenerator(), ) - 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 compact is session_service + assert session_service.session_compact_manager is not None + assert long_term.tools is not None assert len(agent.before_model_callback) == 5 tool_names = {tool.name for tool in agent.tools} assert tool_names == { @@ -174,9 +173,34 @@ async def test_full_setup_wraps_session_service_and_is_idempotent(tmp_path: Path } +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=[]) + + bound = memory_service.bind(agent, session_service) + + assert bound is session_service + assert len(agent.before_model_callback) == 1 + assert isinstance( + agent.before_model_callback[0], + LongTermMemoryContextCallback, + ) + assert {tool.name + for tool in agent.tools} == { + "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(AdvancedCompactConfig(enabled=False, root_dir=tmp_path)) request = LlmRequest(model="test-model") applied = await LongTermMemoryContext(runtime).apply(request) diff --git a/tests/advanced_memory/test_preload_memory.py b/tests/advanced_memory/test_preload_memory.py index 602421263..2f4a3f054 100644 --- a/tests/advanced_memory/test_preload_memory.py +++ b/tests/advanced_memory/test_preload_memory.py @@ -5,7 +5,7 @@ from pathlib import Path from types import SimpleNamespace -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig +from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime from trpc_agent_sdk.advanced_memory import MemoryDocument from trpc_agent_sdk.advanced_memory import MemoryPreloader @@ -32,7 +32,7 @@ async def select(self, query, candidates, ctx, *, limit): async def test_preloader_injects_selected_topic_with_budget(tmp_path: Path) -> None: """Ensure selected topic content is rendered and bounded.""" runtime = AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( + AdvancedCompactConfig( enabled=True, root_dir=tmp_path, preload_memory_enabled=True, @@ -68,7 +68,7 @@ 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( + AdvancedCompactConfig( enabled=True, root_dir=tmp_path, preload_memory_enabled=True, @@ -103,7 +103,7 @@ 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( + AdvancedCompactConfig( enabled=True, root_dir=tmp_path, preload_memory_enabled=True, diff --git a/tests/advanced_memory/test_redis_stores.py b/tests/advanced_memory/test_redis_stores.py index 8718a9849..b681bbc9f 100644 --- a/tests/advanced_memory/test_redis_stores.py +++ b/tests/advanced_memory/test_redis_stores.py @@ -7,16 +7,16 @@ import pytest -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig +from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig from trpc_agent_sdk.advanced_memory import AdvancedMemoryPaths from trpc_agent_sdk.advanced_memory import MemoryIndexEntry -from trpc_agent_sdk.advanced_memory import SessionMemoryDocument from trpc_agent_sdk.advanced_memory._redis_stores import RedisLongTermMemoryStore -from trpc_agent_sdk.advanced_memory._redis_stores import RedisSessionMemoryStore +from trpc_agent_sdk.sessions.compact._redis_stores import RedisToolResultStore +from trpc_agent_sdk.sessions.compact._redis_stores import RedisTranscriptStore def _store(store_type: type, **overrides: object): - config = AdvancedMemoryConfig( + config = AdvancedCompactConfig( storage_backend="redis", redis_url="redis://localhost:6379/0", root_dir=Path("/tmp/advanced-memory-redis-tests"), @@ -53,21 +53,22 @@ async def test_memory_writes_refresh_all_memory_keys() -> None: @pytest.mark.asyncio async def test_session_writes_refresh_all_session_keys() -> None: - store = _store(RedisSessionMemoryStore) + store = _store(RedisToolResultStore) - await store.write("session-1", SessionMemoryDocument(session_title="Test session")) + await store.write("session-1", "result-1", "complete result") session_base = store._session_base("session-1") commands = [call.args for call in store._command.await_args_list] - assert any(command[0] == "set" and command[1] == f"{session_base}:summary" for command in commands) - assert ("sadd", f"{session_base}:keys", f"{session_base}:summary") in commands - assert ("expire", f"{session_base}:summary", 60) in commands + tool_key = f"{session_base}:tool:result-1" + assert any(command[0] == "set" and command[1] == tool_key for command in commands) + assert ("sadd", f"{session_base}:keys", tool_key) in commands + assert ("expire", tool_key, 60) in commands assert ("expire", f"{session_base}:keys", 60) in commands @pytest.mark.asyncio async def test_ttl_refresh_includes_previously_tracked_keys() -> None: - store = _store(RedisSessionMemoryStore, session_ttl_delete_transcripts=True) + store = _store(RedisToolResultStore, session_ttl_delete_transcripts=True) session_base = store._session_base("session-1") old_key = f"{session_base}:transcript" store._command = AsyncMock(side_effect=[ @@ -78,16 +79,17 @@ async def test_ttl_refresh_includes_previously_tracked_keys() -> None: None, # EXPIRE registry ]) - await store._refresh_session_ttl("session-1", f"{session_base}:summary") + current_key = f"{session_base}:tool:result-1" + await store._refresh_session_ttl("session-1", current_key) commands = [call.args for call in store._command.await_args_list] assert ("expire", old_key, 60) in commands - assert ("expire", f"{session_base}:summary", 60) in commands + assert ("expire", current_key, 60) in commands @pytest.mark.asyncio async def test_ttl_refresh_preserves_transcript_by_default() -> None: - store = _store(RedisSessionMemoryStore) + store = _store(RedisToolResultStore) session_base = store._session_base("session-1") old_key = f"{session_base}:transcript" old_seen_key = f"{old_key}:seen:event_id" @@ -98,12 +100,27 @@ async def test_ttl_refresh_preserves_transcript_by_default() -> None: None, # EXPIRE registry ]) - await store._refresh_session_ttl("session-1", f"{session_base}:summary") + current_key = f"{session_base}:tool:result-1" + await store._refresh_session_ttl("session-1", current_key) commands = [call.args for call in store._command.await_args_list] assert ("expire", old_key, 60) not in commands assert ("expire", old_seen_key, 60) not in commands - assert ("expire", f"{session_base}:summary", 60) in commands + assert ("expire", current_key, 60) in commands + + +@pytest.mark.asyncio +async def test_transcript_rejects_event_copies() -> None: + store = _store(RedisTranscriptStore) + + with pytest.raises(ValueError, match="context-compression"): + await store.append( + "session-1", + { + "kind": "event", + "event_id": "event-1" + }, + ) @pytest.mark.asyncio diff --git a/tests/advanced_memory/test_sql_stores.py b/tests/advanced_memory/test_sql_stores.py index 8b96e5a7b..b29745c55 100644 --- a/tests/advanced_memory/test_sql_stores.py +++ b/tests/advanced_memory/test_sql_stores.py @@ -4,19 +4,20 @@ from pathlib import Path +import pytest + from trpc_agent_sdk.advanced_memory import ( - AdvancedMemoryConfig, + AdvancedCompactConfig, AdvancedMemoryRuntime, MemoryDocument, MemoryIndexEntry, MemoryType, - SessionMemoryDocument, ) def _runtime(tmp_path: Path) -> AdvancedMemoryRuntime: return AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( + AdvancedCompactConfig( storage_backend="sql", sql_url=f"sqlite:///{tmp_path / 'advanced-memory.db'}", sql_is_async=False, @@ -42,31 +43,59 @@ async def test_sql_stores_round_trip_and_deduplicate(tmp_path: Path) -> None: content="A user profile", ), ) - await scoped.session_memory.write("session", SessionMemoryDocument(session_title="Test")) await scoped.tool_results.write("session", "result", '{"ok": true}') - await scoped.transcripts.append("session", {"event_id": "one"}) + await scoped.transcripts.append( + "session", + { + "kind": "autocompact-failure", + "attempt_id": "one" + }, + ) _, first = await scoped.transcripts.append_unique( "session", - {"event_id": "two"}, - unique_key="event_id", + { + "kind": "history-snip", + "snip_id": "two" + }, + unique_key="snip_id", ) _, second = await scoped.transcripts.append_unique( "session", - {"event_id": "two"}, - unique_key="event_id", + { + "kind": "history-snip", + "snip_id": "two" + }, + unique_key="snip_id", ) assert first is True assert second is False assert "profile.md" in await scoped.long_term_memory.read_index() assert await scoped.long_term_memory.read_topic("profile") - assert await scoped.session_memory.read("session") + assert scoped.session_memory is None assert await scoped.tool_results.read("session", "result") == '{"ok": true}' assert len(await scoped.transcripts.read_all("session")) == 2 await root.close() +async def test_sql_transcript_rejects_event_copies(tmp_path: Path) -> None: + root = _runtime(tmp_path) + scoped = root.for_scope("app", "user") + await scoped.initialize() + + with pytest.raises(ValueError, match="context-compression"): + await scoped.transcripts.append( + "session", + { + "kind": "event", + "event_id": "event-1" + }, + ) + + await root.close() + + async def test_sql_stores_isolate_users(tmp_path: Path) -> None: root = _runtime(tmp_path) first = root.for_scope("app", "first") diff --git a/tests/advanced_memory/test_storage.py b/tests/advanced_memory/test_storage.py index 4de28c71d..e4f025cd5 100644 --- a/tests/advanced_memory/test_storage.py +++ b/tests/advanced_memory/test_storage.py @@ -12,21 +12,21 @@ import pytest -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig +from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig 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 +from trpc_agent_sdk.sessions.compact import SESSION_MEMORY_SECTIONS +from trpc_agent_sdk.sessions.compact import SessionMemoryDocument -def _enabled_config(tmp_path: Path, **overrides: object) -> AdvancedMemoryConfig: +def _enabled_config(tmp_path: Path, **overrides: object) -> AdvancedCompactConfig: """Create an enabled configuration rooted at the test directory.""" - return AdvancedMemoryConfig(enabled=True, root_dir=tmp_path, **overrides) + return AdvancedCompactConfig(enabled=True, root_dir=tmp_path, **overrides) def test_config_reads_context_window_from_environment(monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -34,7 +34,7 @@ def test_config_reads_context_window_from_environment(monkeypatch: pytest.Monkey monkeypatch.setenv("TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS", "128000") monkeypatch.setenv("TRPC_AGENT_MAX_OUTPUT_TOKENS", "8192") - config = AdvancedMemoryConfig() + config = AdvancedCompactConfig() assert config.model_context_window_tokens == 128_000 assert config.max_output_tokens == 8_192 @@ -45,7 +45,7 @@ def test_config_rejects_invalid_context_window_environment(monkeypatch: pytest.M monkeypatch.setenv("TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS", "not-a-number") with pytest.raises(ValueError, match="TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS"): - AdvancedMemoryConfig() + AdvancedCompactConfig() def test_config_rejects_invalid_max_output_tokens_environment(monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -53,12 +53,21 @@ def test_config_rejects_invalid_max_output_tokens_environment(monkeypatch: pytes monkeypatch.setenv("TRPC_AGENT_MAX_OUTPUT_TOKENS", "-1") with pytest.raises(ValueError, match="TRPC_AGENT_MAX_OUTPUT_TOKENS"): - AdvancedMemoryConfig() + AdvancedCompactConfig() + + +def test_config_rejects_unknown_storage_backend(tmp_path: Path) -> None: + """Prevent misspelled external backends from silently using local files.""" + with pytest.raises(ValueError, match="storage_backend must be one of"): + AdvancedCompactConfig( + root_dir=tmp_path, + storage_backend="redisx", # type: ignore[arg-type] + ) 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)) + runtime = AdvancedMemoryRuntime.create(AdvancedCompactConfig(enabled=False, root_dir=tmp_path)) initialized = await runtime.initialize() @@ -67,6 +76,15 @@ async def test_disabled_runtime_does_not_create_directories(tmp_path: Path) -> N assert not (tmp_path / "SESSION").exists() +async def test_runtime_close_is_idempotent(tmp_path: Path) -> None: + """Allow a shared Runtime to be closed by more than one service owner.""" + runtime = AdvancedMemoryRuntime.create(_enabled_config(tmp_path)) + await runtime.initialize() + + await runtime.close() + await runtime.close() + + 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)) @@ -310,7 +328,7 @@ def slow_append(path: Path, serialized: str) -> None: 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( + config = AdvancedCompactConfig( enabled=True, root_dir=tmp_path, memory_index_max_bytes=80, @@ -400,7 +418,7 @@ def test_paths_sanitize_external_identifiers(tmp_path: Path) -> None: 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") + AdvancedCompactConfig(root_dir=tmp_path, memory_dir_name="../MEMORY") def test_memory_freshness_uses_expected_buckets() -> None: diff --git a/tests/advanced_memory/test_autocompact.py b/tests/sessions/compact/test_autocompact.py similarity index 79% rename from tests/advanced_memory/test_autocompact.py rename to tests/sessions/compact/test_autocompact.py index 782faacab..03eaba20b 100644 --- a/tests/advanced_memory/test_autocompact.py +++ b/tests/sessions/compact/test_autocompact.py @@ -5,21 +5,22 @@ 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.sessions.compact import AutoCompact +from trpc_agent_sdk.sessions.compact import AutoCompactCallback +from trpc_agent_sdk.sessions.compact import AdvancedCompactConfig +from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime +from trpc_agent_sdk.sessions.compact import HistorySnipCallback +from trpc_agent_sdk.sessions.compact import MicrocompactCallback +from trpc_agent_sdk.sessions.compact import SessionMemoryDocument +from trpc_agent_sdk.sessions.compact import setup_autocompact +from trpc_agent_sdk.sessions.compact import setup_history_snip +from trpc_agent_sdk.sessions.compact import setup_microcompact +from trpc_agent_sdk.sessions.compact import setup_tool_result_budget +from trpc_agent_sdk.sessions.compact 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.sessions import SessionServiceConfig from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import Part @@ -52,7 +53,7 @@ def _runtime( ) -> AdvancedMemoryRuntime: """Create an isolated runtime with small automatic-compaction limits.""" return AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( + AdvancedCompactConfig( enabled=enabled, root_dir=tmp_path, autocompact_trigger_chars=trigger, @@ -114,10 +115,68 @@ async def test_legacy_compact_replaces_old_prefix_and_keeps_recent(tmp_path: Pat assert len(generator.histories) == 1 +async def test_compact_persists_summary_and_archives_replaced_events(tmp_path: Path) -> None: + """Ensure AutoCompact writes the compressed window through SessionService.""" + runtime = _runtime(tmp_path) + service = InMemorySessionService( + session_config=SessionServiceConfig(store_historical_events=True), + ) + session = await service.create_session( + app_name="demo-app", + user_id="demo-user", + session_id="session-a", + ) + request = _request(5) + for index, content in enumerate(request.contents): + await service.append_event( + session, + Event( + id=f"event-{index}", + invocation_id="invocation-1", + author="user" if index % 2 == 0 else "agent", + content=content.model_copy(deep=True), + ), + ) + ctx = SimpleNamespace( + session_id=session.id, + app_name=session.app_name, + session=session, + session_service=service, + agent=SimpleNamespace(model="fake-model"), + ) + + result = await AutoCompact(runtime, FakeSummaryGenerator()).apply( + request, + session_id=session.id, + ctx=ctx, + force=True, + ) + + assert result.compacted + restored = await service.get_session( + app_name=session.app_name, + user_id=session.user_id, + session_id=session.id, + ) + assert restored is not None + assert restored.events[0].is_summary_event() + assert [event.id for event in restored.events[1:]] == ["event-3", "event-4"] + assert [event.id for event in restored.historical_events] == [ + "event-0", + "event-1", + "event-2", + ] + assert not restored.compact_events( + Event(author="system", content=Content(parts=[Part.from_text(text="duplicate")])), + "event-2", + compaction_id=restored.events[0].custom_metadata["session_compaction_id"], + ) + + 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( + AdvancedCompactConfig( enabled=True, root_dir=tmp_path, autocompact_trigger_chars=100_000, @@ -142,6 +201,40 @@ async def test_token_budget_triggers_autocompact_and_records_diagnostics(tmp_pat assert records[-1]["request_tokens_before"] == result.request_tokens_before +async def test_token_reduction_uses_consistent_full_request_estimates(tmp_path: Path) -> None: + """Do not compare a usage-based before value with an estimated after value.""" + runtime = AdvancedMemoryRuntime.create( + AdvancedCompactConfig( + enabled=True, + root_dir=tmp_path, + model_context_window_tokens=20_000, + max_output_tokens=100, + token_warning_ratio=0.4, + token_autocompact_ratio=0.5, + autocompact_keep_recent_contents=2, + )).for_scope("demo-app", "demo-user") + request = _request(5) + ctx = _ctx() + ctx.session.events = [ + SimpleNamespace( + content=request.contents[0].model_copy(deep=True), + usage_metadata=SimpleNamespace(total_token_count=12_000), + custom_metadata={}, + ), + ] + + result = await AutoCompact(runtime, FakeSummaryGenerator()).apply( + request, + session_id="session-a", + ctx=ctx, + ) + + assert result.compacted + assert result.request_tokens_after < result.request_tokens_before + assert result.request_tokens_before < 12_000 + assert result.token_source == "estimated" + + 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( diff --git a/tests/sessions/compact/test_context_compression_integration.py b/tests/sessions/compact/test_context_compression_integration.py new file mode 100644 index 000000000..6ac51305d --- /dev/null +++ b/tests/sessions/compact/test_context_compression_integration.py @@ -0,0 +1,452 @@ +"""Tests for request compression over an unchanged SessionService.""" + +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from trpc_agent_sdk.evaluation._eval_session_service import EvalSessionService +from trpc_agent_sdk.sessions.compact import AdvancedCompactConfig +from trpc_agent_sdk.sessions.compact import BaseSessionCompactConfig +from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime +from trpc_agent_sdk.sessions.compact import BaseSessionCompactManager +from trpc_agent_sdk.sessions.compact import AutoCompactCallback +from trpc_agent_sdk.sessions.compact import HistorySnipCallback +from trpc_agent_sdk.sessions.compact import MicrocompactCallback +from trpc_agent_sdk.sessions.compact import SESSION_MEMORY_STATE_KEY +from trpc_agent_sdk.sessions.compact import SessionMemoryDocument +from trpc_agent_sdk.sessions.compact import ToolResultBudget +from trpc_agent_sdk.sessions.compact import ToolResultBudgetCallback +from trpc_agent_sdk.sessions.compact import setup_advanced_session_compact +from trpc_agent_sdk.sessions.compact import setup_context_compression +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 import SqlSessionService +from trpc_agent_sdk.types import Content +from trpc_agent_sdk.types import FunctionResponse +from trpc_agent_sdk.types import Part + + +class FakeSummaryGenerator: + """Return a deterministic autocompact summary.""" + + async def generate(self, history: str, ctx) -> str: + del history, ctx + return "summary" + + +class FakeSessionMemoryGenerator: + """Return deterministic structured Session Memory.""" + + async def generate(self, extraction_input, ctx) -> SessionMemoryDocument: + del ctx + return SessionMemoryDocument( + session_title="Post-turn memory", + current_state=f"Processed {extraction_input.last_event_id}", + ) + + +class DummySummarizerManager: + """Provide the BaseSessionService attachment protocol.""" + + def set_session_service(self, service) -> None: + self.service = service + + +def _runtime(tmp_path: Path) -> AdvancedMemoryRuntime: + return AdvancedMemoryRuntime.create( + AdvancedCompactConfig( + root_dir=tmp_path, + tool_result_max_chars=200, + tool_results_per_message_max_chars=5_000, + tool_result_preview_chars=40, + )) + + +def _session_service() -> InMemorySessionService: + return InMemorySessionService( + session_config=SessionServiceConfig(store_historical_events=True), + ) + + +async def test_session_service_accepts_base_compact_manager(tmp_path: Path) -> None: + """Inject the Advanced manager through the common manager contract.""" + agent = SimpleNamespace(before_model_callback=None) + service = InMemorySessionService( + session_config=SessionServiceConfig(store_historical_events=True), + ) + manager = setup_advanced_session_compact( + agent, + service, + AdvancedCompactConfig(root_dir=tmp_path), + session_memory_generator=FakeSessionMemoryGenerator(), + ) + + assert isinstance(service.session_compact_manager, BaseSessionCompactManager) + assert service.session_compact_manager is manager + await service.close() + + +def test_advanced_config_implements_compact_config_contract() -> None: + """Concrete strategies must be selectable through the config base class.""" + assert issubclass(AdvancedCompactConfig, BaseSessionCompactConfig) + + +async def test_advanced_setup_infers_sql_backend_from_session_service( + tmp_path: Path, +) -> None: + """Use the SessionService as the single source of backend settings.""" + database_url = f"sqlite:///{tmp_path / 'compact.db'}" + service = SqlSessionService( + db_url=database_url, + is_async=False, + session_config=SessionServiceConfig(store_historical_events=True), + ) + manager = setup_advanced_session_compact( + SimpleNamespace(before_model_callback=None), + service, + AdvancedCompactConfig(root_dir=tmp_path), + session_memory_generator=FakeSessionMemoryGenerator(), + ) + + assert manager.runtime.config.storage_backend == "sql" + assert manager.runtime.config.sql_url == database_url + assert manager.runtime.config.sql_is_async is False + await service.close() + + +@pytest.mark.asyncio +async def test_runner_auto_installs_compact_from_session_config(tmp_path: Path) -> None: + """Let Runner create the manager from the declarative SessionService config.""" + from trpc_agent_sdk.runners import Runner + + agent = SimpleNamespace( + name="compact-agent", + tools=[], + before_model_callback=None, + get_subagents=lambda: [], + ) + service = InMemorySessionService( + session_config=SessionServiceConfig(store_historical_events=True), + session_compact_config=AdvancedCompactConfig(root_dir=tmp_path), + ) + + runner = Runner( + app_name="compact-test", + agent=agent, + session_service=service, + enable_post_turn_processing=False, + ) + + assert service.session_compact_manager is not None + assert service.session_compact_manager.runtime.config.root_dir == tmp_path.resolve() + await runner.close() + + +def _tool_event(output: str) -> Event: + return Event( + id="event-1", + invocation_id="invocation-1", + author="user", + content=Content(parts=[ + Part(function_response=FunctionResponse( + id="result-1", + name="demo_tool", + response={"output": output}, + )) + ]), + ) + + +async def test_setup_attaches_manager_to_original_service(tmp_path: Path) -> None: + """Install only the four request callbacks over the original service.""" + runtime = _runtime(tmp_path) + delegate = _session_service() + agent = SimpleNamespace(before_model_callback=None) + + service = setup_context_compression( + agent, + delegate, + runtime, + FakeSummaryGenerator(), + ) + + assert service is delegate + assert service.session_compact_manager is not None + assert service.session_compact_manager.runtime is runtime + assert [type(callback) for callback in agent.before_model_callback] == [ + ToolResultBudgetCallback, + HistorySnipCallback, + MicrocompactCallback, + AutoCompactCallback, + ] + + +async def test_setup_rejects_original_session_summarizer(tmp_path: Path) -> None: + """Prevent two independent mechanisms from writing summary Events.""" + delegate = InMemorySessionService( + summarizer_manager=DummySummarizerManager(), + session_config=SessionServiceConfig(store_historical_events=True), + ) + agent = SimpleNamespace(before_model_callback=None) + + with pytest.raises(ValueError, match="mutually exclusive"): + setup_context_compression( + agent, + delegate, + _runtime(tmp_path), + FakeSummaryGenerator(), + ) + await delegate.close() + + +async def test_manager_keeps_events_in_original_service_only(tmp_path: Path) -> None: + """Read and append Events without a second Event transcript.""" + runtime = _runtime(tmp_path) + delegate = _session_service() + session = await delegate.create_session( + app_name="demo-app", + user_id="demo-user", + session_id="legacy-session", + ) + old_event = Event( + id="old-event", + invocation_id="invocation-1", + author="user", + content=Content(parts=[Part.from_text(text="old event")]), + ) + await delegate.append_event(session, old_event) + + agent = SimpleNamespace(before_model_callback=None) + service = setup_context_compression(agent, delegate, runtime, FakeSummaryGenerator()) + loaded = await service.get_session( + app_name="demo-app", + user_id="demo-user", + session_id=session.id, + ) + assert loaded is not None + await service.append_event(loaded, _tool_event("x" * 500)) + + stored = await delegate.get_session( + app_name="demo-app", + user_id="demo-user", + session_id=session.id, + ) + assert stored is not None + assert [event.id for event in stored.events] == ["old-event", "event-1"] + assert await runtime.for_session(stored).transcripts.read_all(stored.id) == [] + + +async def test_request_replacement_does_not_rewrite_stored_event(tmp_path: Path) -> None: + """Replace a request copy while retaining the complete persisted result.""" + runtime = _runtime(tmp_path) + delegate = _session_service() + agent = SimpleNamespace(before_model_callback=None) + service = setup_context_compression(agent, delegate, runtime, FakeSummaryGenerator()) + session = await service.create_session( + app_name="demo-app", + user_id="demo-user", + session_id="budget-session", + ) + await service.append_event(session, _tool_event("x" * 500)) + request = LlmRequest( + model="test-model", + contents=[session.events[0].content.model_copy(deep=True)], + ) + + result = await ToolResultBudget(runtime.for_session(session)).apply( + request, + session_id=session.id, + ) + + stored = await delegate.get_session( + app_name="demo-app", + user_id="demo-user", + session_id=session.id, + ) + assert result.replaced_count == 1 + assert "persisted_output" in request.contents[0].parts[0].function_response.response + assert stored is not None + assert stored.events[0].content.parts[0].function_response.response == { + "output": "x" * 500 + } + records = await runtime.for_session(stored).transcripts.read_all(stored.id) + assert all(record.get("kind") != "event" for record in records) + + +async def test_setup_is_idempotent_and_validates_runtime_first(tmp_path: Path) -> None: + """Reuse one manager and reject a different runtime without changing callbacks.""" + runtime = _runtime(tmp_path / "one") + agent = SimpleNamespace(before_model_callback=None) + service = setup_context_compression( + agent, + _session_service(), + runtime, + FakeSummaryGenerator(), + ) + repeated = setup_context_compression(agent, service, runtime, FakeSummaryGenerator()) + assert repeated is service + assert len(agent.before_model_callback) == 4 + + clean_agent = SimpleNamespace(before_model_callback=None) + with pytest.raises(ValueError, match="another runtime"): + setup_context_compression( + clean_agent, + service, + _runtime(tmp_path / "two"), + FakeSummaryGenerator(), + ) + assert clean_agent.before_model_callback is None + + +async def test_compact_manager_is_mutually_exclusive_with_native_summarizer(tmp_path: Path) -> None: + """Prevent adding the native summarizer after compact setup.""" + service = _session_service() + setup_context_compression( + SimpleNamespace(before_model_callback=None), + service, + _runtime(tmp_path), + FakeSummaryGenerator(), + ) + + with pytest.raises(ValueError, match="mutually exclusive"): + service.set_summarizer_manager(DummySummarizerManager()) + + +async def test_original_service_delete_cleans_compact_side_data(tmp_path: Path) -> None: + """Run compact cleanup through the original SessionService lifecycle.""" + runtime = _runtime(tmp_path) + service = _session_service() + setup_context_compression( + SimpleNamespace(before_model_callback=None), + service, + runtime, + FakeSummaryGenerator(), + ) + session = await service.create_session( + app_name="demo-app", + user_id="demo-user", + session_id="delete-me", + ) + scoped = runtime.for_session(session) + await scoped.transcripts.append(session.id, {"kind": "test-record"}) + + await service.delete_session( + app_name=session.app_name, + user_id=session.user_id, + session_id=session.id, + ) + + assert await scoped.transcripts.read_all(session.id) == [] + + +async def test_eval_session_service_forwards_compact_manager(tmp_path: Path) -> None: + """Keep evaluation wrappers on the inner service's compact lifecycle.""" + inner = _session_service() + service = EvalSessionService(inner) + runtime = _runtime(tmp_path) + + configured = setup_context_compression( + SimpleNamespace(before_model_callback=None), + service, + runtime, + FakeSummaryGenerator(), + ) + + assert configured is service + assert service.session_compact_manager is inner.session_compact_manager + assert service.session_compact_manager.runtime is runtime + + +async def test_sql_delegate_keeps_its_existing_event_storage(tmp_path: Path) -> None: + """Ensure manager composition works with the SQL SessionService.""" + runtime = _runtime(tmp_path / "advanced") + delegate = SqlSessionService( + db_url=f"sqlite:///{tmp_path / 'sessions.db'}", + is_async=False, + ) + session = await delegate.create_session( + app_name="demo-app", + user_id="demo-user", + session_id="sql-session", + ) + event = Event( + id="sql-event", + invocation_id="invocation-1", + author="user", + content=Content(parts=[Part.from_text(text="stored by SQL")]), + ) + await delegate.append_event(session, event) + + service = setup_context_compression( + SimpleNamespace(before_model_callback=None), + delegate, + runtime, + FakeSummaryGenerator(), + ) + loaded = await service.get_session( + app_name="demo-app", + user_id="demo-user", + session_id=session.id, + ) + + assert loaded is not None + assert [item.id for item in loaded.events] == ["sql-event"] + assert await runtime.for_session(loaded).transcripts.read_all(loaded.id) == [] + await service.close() + await runtime.close() + + +async def test_post_turn_hook_updates_session_memory_state(tmp_path: Path) -> None: + """Ensure the existing Runner summary hook updates Session Memory.""" + database = tmp_path / "post-turn.db" + runtime = AdvancedMemoryRuntime.create( + AdvancedCompactConfig( + storage_backend="sql", + sql_url=f"sqlite:///{database}", + sql_is_async=False, + session_memory_initial_chars=1, + session_memory_update_chars=1, + ), + ) + delegate = SqlSessionService( + db_url=f"sqlite:///{database}", + is_async=False, + ) + agent = SimpleNamespace(before_model_callback=None) + service = setup_context_compression( + agent, + delegate, + runtime, + FakeSummaryGenerator(), + session_memory_generator=FakeSessionMemoryGenerator(), + ) + session = await service.create_session( + app_name="demo-app", + user_id="demo-user", + session_id="post-turn", + ) + await service.append_event(session, _tool_event("post-turn content")) + ctx = SimpleNamespace( + session=session, + session_service=service, + agent=SimpleNamespace(model="fake-model"), + ) + + await service.create_session_summary(session, ctx=ctx) + + assert SESSION_MEMORY_STATE_KEY in session.state + loaded = await service.get_session( + app_name=session.app_name, + user_id=session.user_id, + session_id=session.id, + ) + assert loaded is not None + assert SESSION_MEMORY_STATE_KEY in loaded.state + summary = await service.get_session_summary(loaded) + assert summary is not None + assert "Post-turn memory" in summary + await service.close() + await runtime.close() 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/advanced_memory/test_history_snip.py b/tests/sessions/compact/test_history_snip.py similarity index 90% rename from tests/advanced_memory/test_history_snip.py rename to tests/sessions/compact/test_history_snip.py index 9d13a966c..7c39c88d8 100644 --- a/tests/advanced_memory/test_history_snip.py +++ b/tests/sessions/compact/test_history_snip.py @@ -5,17 +5,17 @@ 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.sessions.compact import AdvancedCompactConfig +from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime +from trpc_agent_sdk.sessions.compact import HistorySnip +from trpc_agent_sdk.sessions.compact import HistorySnipCallback +from trpc_agent_sdk.sessions.compact import Microcompact +from trpc_agent_sdk.sessions.compact import MicrocompactCallback +from trpc_agent_sdk.sessions.compact import setup_history_snip +from trpc_agent_sdk.sessions.compact import setup_microcompact +from trpc_agent_sdk.sessions.compact import setup_tool_result_budget +from trpc_agent_sdk.sessions.compact import ToolResultBudgetCallback +from trpc_agent_sdk.sessions.compact import ToolResultBudget from trpc_agent_sdk.models import LlmRequest from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import FunctionResponse @@ -33,7 +33,7 @@ def _runtime( ) -> AdvancedMemoryRuntime: """Create an isolated runtime with small history-snip limits.""" return AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( + AdvancedCompactConfig( enabled=enabled, root_dir=tmp_path, tool_result_max_chars=5_000, @@ -85,7 +85,7 @@ async def test_token_budget_triggers_snip_without_character_pressure(tmp_path: P """Ensure a configured model window triggers cleanup by token warning.""" request, _ = _request(4, output_size=1_000) runtime = AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( + AdvancedCompactConfig( enabled=True, root_dir=tmp_path, tool_result_max_chars=10_000, @@ -159,7 +159,7 @@ async def test_snipped_results_are_reapplied_after_restart(tmp_path: Path) -> No 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( + AdvancedCompactConfig( enabled=True, root_dir=tmp_path, tool_result_max_chars=200, diff --git a/tests/advanced_memory/test_microcompact.py b/tests/sessions/compact/test_microcompact.py similarity index 92% rename from tests/advanced_memory/test_microcompact.py rename to tests/sessions/compact/test_microcompact.py index 91887898e..4b76d961d 100644 --- a/tests/advanced_memory/test_microcompact.py +++ b/tests/sessions/compact/test_microcompact.py @@ -5,13 +5,13 @@ 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.sessions.compact import AdvancedCompactConfig +from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime +from trpc_agent_sdk.sessions.compact import Microcompact +from trpc_agent_sdk.sessions.compact import MicrocompactCallback +from trpc_agent_sdk.sessions.compact import setup_microcompact +from trpc_agent_sdk.sessions.compact import setup_tool_result_budget +from trpc_agent_sdk.sessions.compact import ToolResultBudgetCallback from trpc_agent_sdk.models import LlmRequest from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import FunctionResponse @@ -29,7 +29,7 @@ def _runtime( ) -> AdvancedMemoryRuntime: """Create an isolated runtime with small mechanical-compaction limits.""" return AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( + AdvancedCompactConfig( enabled=enabled, root_dir=tmp_path, tool_result_max_chars=1_000, diff --git a/tests/advanced_memory/test_session_memory_extractor.py b/tests/sessions/compact/test_session_memory_extractor.py similarity index 97% rename from tests/advanced_memory/test_session_memory_extractor.py rename to tests/sessions/compact/test_session_memory_extractor.py index 7b8e9273a..5ebdf4dc0 100644 --- a/tests/advanced_memory/test_session_memory_extractor.py +++ b/tests/sessions/compact/test_session_memory_extractor.py @@ -7,13 +7,13 @@ 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.sessions.compact import AdvancedCompactConfig +from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime +from trpc_agent_sdk.sessions.compact import ForkedSessionMemoryGenerator +from trpc_agent_sdk.sessions.compact import SessionMemoryExtractionInput +from trpc_agent_sdk.sessions.compact import SessionMemoryDocument +from trpc_agent_sdk.sessions.compact import SessionMemoryExtractor +from trpc_agent_sdk.sessions.compact import TranscriptSessionService from trpc_agent_sdk.events import Event from trpc_agent_sdk.models import LLMModel from trpc_agent_sdk.models import LlmResponse @@ -94,7 +94,7 @@ def _runtime( ) -> AdvancedMemoryRuntime: """Create an isolated runtime with small extraction limits.""" return AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( + AdvancedCompactConfig( enabled=True, root_dir=tmp_path, session_memory_initial_chars=initial_chars, @@ -163,7 +163,7 @@ async def test_first_extraction_writes_document_and_checkpoint(tmp_path: Path) - 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( + AdvancedCompactConfig( enabled=True, root_dir=tmp_path, session_memory_initial_chars=100_000, diff --git a/tests/sessions/compact/test_session_memory_state.py b/tests/sessions/compact/test_session_memory_state.py new file mode 100644 index 000000000..ee0b60492 --- /dev/null +++ b/tests/sessions/compact/test_session_memory_state.py @@ -0,0 +1,160 @@ +"""Session-state persistence tests for Redis/SQL Advanced Memory.""" + +from pathlib import Path +from types import SimpleNamespace + +from trpc_agent_sdk.sessions.compact import AdvancedCompactConfig +from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime +from trpc_agent_sdk.sessions.compact import AutoCompact +from trpc_agent_sdk.sessions.compact import SessionMemoryDocument +from trpc_agent_sdk.sessions.compact import SessionMemoryExtractor +from trpc_agent_sdk.sessions.compact._formats import SESSION_MEMORY_STATE_KEY +from trpc_agent_sdk.sessions.compact._formats import parse_session_memory_state +from trpc_agent_sdk.events import Event +from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.sessions import SqlSessionService +from trpc_agent_sdk.types import Content +from trpc_agent_sdk.types import Part + + +class _Generator: + + def __init__(self) -> None: + self.inputs = [] + + async def generate(self, extraction_input, ctx) -> SessionMemoryDocument: + del ctx + self.inputs.append(extraction_input) + return SessionMemoryDocument( + session_title="State-backed session", + current_state=f"Processed {extraction_input.last_event_id}", + ) + + +class _LegacyGenerator: + + async def generate(self, history, ctx) -> str: + del history, ctx + return "legacy" + + +def _event(event_id: str, text: str) -> Event: + return Event( + id=event_id, + invocation_id="invocation", + author="agent", + content=Content(role="model", parts=[Part.from_text(text=text)]), + ) + + +async def test_sql_session_memory_is_persisted_in_session_state(tmp_path: Path, ) -> None: + database = tmp_path / "state-memory.db" + runtime = AdvancedMemoryRuntime.create( + AdvancedCompactConfig( + storage_backend="sql", + sql_url=f"sqlite:///{database}", + sql_is_async=False, + session_memory_initial_chars=1, + session_memory_update_chars=1, + )) + service = SqlSessionService(db_url=f"sqlite:///{database}", is_async=False) + session = await service.create_session( + app_name="app", + user_id="user", + session_id="session", + ) + await service.append_event(session, _event("event-1", "x" * 2_000)) + generator = _Generator() + extractor = SessionMemoryExtractor( + runtime, + generator, + session_service=service, + ) + ctx = SimpleNamespace( + session=session, + agent=SimpleNamespace(model="test-model"), + ) + + result = await extractor.extract_if_needed(session, ctx, force=True) + + loaded = await service.get_session( + app_name="app", + user_id="user", + session_id="session", + ) + assert result.extracted is True + assert loaded is not None + parsed = parse_session_memory_state(loaded.state[SESSION_MEMORY_STATE_KEY]) + assert parsed is not None + document, checkpoint, _ = parsed + assert document.current_state == "Processed event-1" + assert checkpoint["last_event_id"] == "event-1" + assert len(loaded.events) == 1 + assert runtime.for_session(loaded).session_memory is None + assert await runtime.for_session(loaded).transcripts.read_all(loaded.id) == [] + await service.close() + await runtime.close() + + +async def test_autocompact_generates_state_memory_only_when_invoked(tmp_path: Path, ) -> None: + database = tmp_path / "autocompact-state.db" + runtime = AdvancedMemoryRuntime.create( + AdvancedCompactConfig( + storage_backend="sql", + sql_url=f"sqlite:///{database}", + sql_is_async=False, + autocompact_target_chars=20_000, + session_memory_initial_chars=1, + session_memory_update_chars=1, + )) + service = SqlSessionService(db_url=f"sqlite:///{database}", is_async=False) + session = await service.create_session( + app_name="app", + user_id="user", + session_id="session", + ) + for index in range(3): + await service.append_event( + session, + _event(f"event-{index}", f"message-{index}-" + "x" * 3_000), + ) + generator = _Generator() + extractor = SessionMemoryExtractor( + runtime, + generator, + session_service=service, + ) + compressor = AutoCompact(runtime, _LegacyGenerator()) + compressor.attach_session_memory_extractor(extractor) + ctx = SimpleNamespace( + session=session, + session_service=service, + agent=SimpleNamespace(model="test-model"), + ) + request = LlmRequest( + model="test-model", + contents=[event.content.model_copy(deep=True) for event in session.events], + ) + + result = await compressor.apply( + request, + session_id=session.id, + ctx=ctx, + force=True, + ) + + assert result.compacted is True + assert result.source == "session-memory" + assert generator.inputs + assert SESSION_MEMORY_STATE_KEY in session.state + assert session.events[0].is_summary_event() + assert [event.id for event in session.historical_events] == [ + "event-0", + "event-1", + "event-2", + ] + records = await runtime.for_session(session).transcripts.read_all(session.id) + assert [record["kind"] for record in records] == ["autocompact-success"] + assert all(record["kind"] != "event" for record in records) + await service.close() + await runtime.close() diff --git a/tests/advanced_memory/test_token_budget.py b/tests/sessions/compact/test_token_budget.py similarity index 89% rename from tests/advanced_memory/test_token_budget.py rename to tests/sessions/compact/test_token_budget.py index 0431f9515..4e2a29fbb 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,7 +37,7 @@ 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, @@ -63,7 +63,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, root_dir=tmp_path)) estimate = tracker.estimate(request, ctx) @@ -85,7 +85,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, root_dir=tmp_path)).estimate(request, ctx) assert estimate.source == "estimated" assert estimate.tokens < 999_999 @@ -94,7 +94,7 @@ 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, @@ -111,7 +111,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, + budget = TokenContextTracker(AdvancedCompactConfig(enabled=True, root_dir=tmp_path)).budget(_request("compatibility request")) assert not budget.token_mode_enabled diff --git a/tests/advanced_memory/test_tool_result_budget.py b/tests/sessions/compact/test_tool_result_budget.py similarity index 96% rename from tests/advanced_memory/test_tool_result_budget.py rename to tests/sessions/compact/test_tool_result_budget.py index 3fbbd17c3..2a3e806b4 100644 --- a/tests/advanced_memory/test_tool_result_budget.py +++ b/tests/sessions/compact/test_tool_result_budget.py @@ -8,11 +8,11 @@ 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.sessions.compact import AdvancedCompactConfig +from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime +from trpc_agent_sdk.sessions.compact import setup_tool_result_budget +from trpc_agent_sdk.sessions.compact import ToolResultBudget +from trpc_agent_sdk.sessions.compact import ToolResultBudgetCallback from trpc_agent_sdk.models import LlmRequest from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import FunctionResponse @@ -29,7 +29,7 @@ def _runtime( ) -> AdvancedMemoryRuntime: """Create an isolated runtime with small test limits.""" return AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( + AdvancedCompactConfig( enabled=enabled, root_dir=tmp_path, tool_result_max_chars=per_result, @@ -72,7 +72,7 @@ async def test_single_large_result_is_persisted_and_replaced(tmp_path: Path) -> async def test_sql_replacement_reports_sql_storage_path(tmp_path: Path) -> None: """Expose the path returned by the SQL tool-result store.""" - root = AdvancedMemoryRuntime.create(AdvancedMemoryConfig( + root = AdvancedMemoryRuntime.create(AdvancedCompactConfig( enabled=True, storage_backend="sql", sql_url=f"sqlite:///{tmp_path / 'memory.db'}", diff --git a/tests/advanced_memory/test_transcript_session_service.py b/tests/sessions/compact/test_transcript_session_service.py similarity index 78% rename from tests/advanced_memory/test_transcript_session_service.py rename to tests/sessions/compact/test_transcript_session_service.py index 47cf79322..f4a20fe7f 100644 --- a/tests/advanced_memory/test_transcript_session_service.py +++ b/tests/sessions/compact/test_transcript_session_service.py @@ -4,9 +4,11 @@ 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 +import pytest + +from trpc_agent_sdk.sessions.compact import AdvancedCompactConfig +from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime +from trpc_agent_sdk.sessions.compact import TranscriptSessionService from trpc_agent_sdk.events import Event from trpc_agent_sdk.sessions import InMemorySessionService from trpc_agent_sdk.types import Content @@ -35,7 +37,7 @@ async def _session(service: TranscriptSessionService): 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)) + runtime = AdvancedMemoryRuntime.create(AdvancedCompactConfig(enabled=True, root_dir=tmp_path)) service = TranscriptSessionService(InMemorySessionService(), runtime) session = await _session(service) @@ -57,7 +59,7 @@ async def test_append_event_writes_versioned_parent_chain(tmp_path: Path) -> Non 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)) + runtime = AdvancedMemoryRuntime.create(AdvancedCompactConfig(enabled=True, root_dir=tmp_path)) service = TranscriptSessionService(InMemorySessionService(), runtime) session = await _session(service) duplicate = _event("event-1", "hello") @@ -71,7 +73,7 @@ async def test_duplicate_event_id_is_not_written_twice(tmp_path: Path) -> None: 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)) + runtime = AdvancedMemoryRuntime.create(AdvancedCompactConfig(enabled=True, root_dir=tmp_path)) service = TranscriptSessionService(InMemorySessionService(), runtime) session = await _session(service) await service.append_event(session, _event("event-1", "first")) @@ -87,13 +89,13 @@ async def test_old_duplicate_does_not_rewind_parent_chain(tmp_path: Path) -> Non 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)) + runtime = AdvancedMemoryRuntime.create(AdvancedCompactConfig(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_runtime = AdvancedMemoryRuntime.create(AdvancedCompactConfig(enabled=True, root_dir=tmp_path)) second_service = TranscriptSessionService(delegate, second_runtime) await second_service.append_event(session, _event("event-2", "second")) @@ -103,7 +105,7 @@ async def test_new_wrapper_restores_parent_from_existing_transcript(tmp_path: Pa 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)) + runtime = AdvancedMemoryRuntime.create(AdvancedCompactConfig(enabled=False, root_dir=tmp_path)) service = TranscriptSessionService(InMemorySessionService(), runtime) session = await _session(service) @@ -115,9 +117,18 @@ async def test_disabled_runtime_preserves_old_service_without_disk_writes(tmp_pa assert not (tmp_path / "SESSION").exists() +async def test_nested_transcript_wrapper_is_rejected(tmp_path: Path) -> None: + """Ensure a transcript decorator cannot wrap another decorator.""" + runtime = AdvancedMemoryRuntime.create(AdvancedCompactConfig(enabled=True, root_dir=tmp_path)) + inner = TranscriptSessionService(InMemorySessionService(), runtime) + + with pytest.raises(ValueError, match="already wrapped"): + TranscriptSessionService(inner, runtime) + + 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)) + runtime = AdvancedMemoryRuntime.create(AdvancedCompactConfig(enabled=True, root_dir=tmp_path)) service = TranscriptSessionService(InMemorySessionService(), runtime) session = await _session(service) 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/__init__.py b/trpc_agent_sdk/advanced_memory/__init__.py index fe1afb8a5..658252342 100644 --- a/trpc_agent_sdk/advanced_memory/__init__.py +++ b/trpc_agent_sdk/advanced_memory/__init__.py @@ -3,97 +3,40 @@ # 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.""" +"""Optional long-term memory APIs.""" -from ._autocompact import AutoCompact -from ._autocompact import AutoCompactCallback -from ._autocompact import AutoCompactResult -from ._autocompact import content_signature -from ._autocompact import ForkedLegacySummaryGenerator -from ._autocompact import setup_autocompact -from ._config import AdvancedMemoryConfig -from ._formats import MemoryDocument -from ._formats import MemoryIndexEntry -from ._formats import MemoryType -from ._formats import memory_freshness -from ._formats import parse_memory_updated_at -from ._formats import SESSION_MEMORY_SECTION_DESCRIPTIONS -from ._formats import SESSION_MEMORY_SECTIONS -from ._formats import SessionMemoryDocument -from ._history_snip import estimate_request_chars -from ._history_snip import HistorySnip -from ._history_snip import HistorySnipCallback -from ._history_snip import HistorySnipResult -from ._history_snip import setup_history_snip +from trpc_agent_sdk.sessions.compact._config import AdvancedCompactConfig +from trpc_agent_sdk.sessions.compact._formats import MemoryDocument +from trpc_agent_sdk.sessions.compact._formats import MemoryIndexEntry +from trpc_agent_sdk.sessions.compact._formats import MemoryType +from trpc_agent_sdk.sessions.compact._formats import memory_freshness +from trpc_agent_sdk.sessions.compact._formats import parse_memory_updated_at +from trpc_agent_sdk.sessions.compact._paths import AdvancedMemoryPaths +from trpc_agent_sdk.sessions.compact._paths import MemoryScope +from trpc_agent_sdk.sessions.compact._runtime import AdvancedMemoryRuntime +from trpc_agent_sdk.sessions.compact._runtime import ScopedAdvancedMemoryRuntime +from trpc_agent_sdk.sessions.compact._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 ._microcompact import Microcompact -from ._microcompact import MicrocompactCallback -from ._microcompact import MicrocompactResult -from ._microcompact import setup_microcompact -from ._paths import AdvancedMemoryPaths -from ._paths import MemoryScope 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 ._runtime import ScopedAdvancedMemoryRuntime -from ._session_memory import build_session_memory_prompt -from ._session_memory import ForkedSessionMemoryGenerator -from ._session_memory import has_session_memory_content -from ._session_memory import limit_session_memory_document -from ._session_memory import SessionMemoryExtractionInput -from ._session_memory import SessionMemoryExtractionResult -from ._session_memory import SessionMemoryExtractor -from ._session_service import TranscriptSessionService -from ._integration import AdvancedContextManagement -from ._integration import AdvancedMemoryIntegration -from ._integration import setup_advanced_memory -from ._integration import setup_context_management -from ._storage import LongTermMemoryStore -from ._storage import SessionMemoryStore -from ._storage import ToolResultStore -from ._storage import TranscriptStore from ._storage_backend import AdvancedMemoryStorageBackend from ._storage_backend import LocalAdvancedMemoryStorageBackend -from ._tool_result_budget import setup_tool_result_budget -from ._tool_result_budget import ToolResultBudget -from ._tool_result_budget import ToolResultBudgetCallback -from ._tool_result_budget import ToolResultBudgetResult -from ._transcript import TRANSCRIPT_SCHEMA_VERSION -from ._token_budget import ContextBudget -from ._token_budget import ContextTokenEstimate -from ._token_budget import HeuristicTokenEstimator -from ._token_budget import ModelContextWindowResolver -from ._token_budget import TokenContextTracker -from ._token_budget import TokenEstimator __all__ = [ - "AutoCompact", - "AutoCompactCallback", - "AutoCompactResult", "AdvancedMemoryStorageBackend", - "AdvancedMemoryConfig", - "AdvancedContextManagement", - "AdvancedMemoryIntegration", + "AdvancedCompactConfig", + "LongTermMemoryIntegration", "AdvancedMemoryPaths", "AdvancedMemoryRuntime", "ScopedAdvancedMemoryRuntime", - "ContextBudget", - "ContextTokenEstimate", - "build_session_memory_prompt", - "content_signature", - "estimate_request_chars", - "ForkedLegacySummaryGenerator", - "ForkedSessionMemoryGenerator", - "has_session_memory_content", - "HistorySnip", - "HistorySnipCallback", - "HistorySnipResult", - "HeuristicTokenEstimator", "LongTermMemoryStore", "LocalAdvancedMemoryStorageBackend", "LongTermMemoryContext", @@ -109,32 +52,6 @@ "select_relevant_memory_filenames", "memory_freshness", "parse_memory_updated_at", - "Microcompact", - "MicrocompactCallback", - "MicrocompactResult", - "ModelContextWindowResolver", - "SESSION_MEMORY_SECTION_DESCRIPTIONS", - "SESSION_MEMORY_SECTIONS", - "SessionMemoryDocument", - "SessionMemoryExtractionInput", - "SessionMemoryExtractionResult", - "SessionMemoryExtractor", - "SessionMemoryStore", - "TRANSCRIPT_SCHEMA_VERSION", - "ToolResultBudget", - "ToolResultBudgetCallback", - "ToolResultBudgetResult", - "ToolResultStore", - "TokenContextTracker", - "TokenEstimator", - "TranscriptSessionService", - "TranscriptStore", - "setup_autocompact", - "setup_advanced_memory", - "limit_session_memory_document", - "setup_history_snip", - "setup_context_management", "setup_long_term_memory_context", - "setup_microcompact", - "setup_tool_result_budget", + "setup_long_term_memory", ] diff --git a/trpc_agent_sdk/advanced_memory/_integration.py b/trpc_agent_sdk/advanced_memory/_integration.py index e4f2ca3e9..b68d8577b 100644 --- a/trpc_agent_sdk/advanced_memory/_integration.py +++ b/trpc_agent_sdk/advanced_memory/_integration.py @@ -3,7 +3,7 @@ # 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.""" +"""Provide setup entry points for long-term memory.""" from __future__ import annotations @@ -11,47 +11,22 @@ 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 trpc_agent_sdk.sessions.compact._runtime import AdvancedMemoryRuntime + 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.""" +class LongTermMemoryIntegration: + """Aggregate the long-term memory callback and tools.""" - context_management: AdvancedContextManagement - session_memory_extractor: SessionMemoryExtractor - session_service: TranscriptSessionService - long_term_memory_tools: "AdvancedMemoryTools | None" + context: LongTermMemoryContext + tools: "AdvancedMemoryTools | None" def _setup_long_term_memory_tools( @@ -60,22 +35,38 @@ def _setup_long_term_memory_tools( ) -> "AdvancedMemoryTools": """Install the three official memory tools idempotently.""" from trpc_agent_sdk.tools._advanced_memory_tool import ( - ADVANCED_MEMORY_TOOL_NAMES, ) + 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] + 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} + 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") + 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") + 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} + 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") + raise ValueError( + "Advanced Memory tools are only partially installed" + ) return owner tools = AdvancedMemoryTools(memory_runtime) agent.tools.extend(tools.as_tools()) @@ -88,101 +79,58 @@ def _setup_preload_memory_tool( 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: + 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"] + 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") + 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, - ), + preloader = MemoryPreloader( + memory_runtime, + ModelMemoryRelevanceSelector(model), + ) + agent.tools.append( + PreloadMemoryTool( + memory_preloader=preloader.preload, + use_legacy_memory=use_legacy_memory, + ) ) -def setup_advanced_memory( +def setup_long_term_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( + 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, - 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, + 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/advanced_memory/_memory_context.py index 1c13bef6e..73bbfd336 100644 --- a/trpc_agent_sdk/advanced_memory/_memory_context.py +++ b/trpc_agent_sdk/advanced_memory/_memory_context.py @@ -9,8 +9,8 @@ from typing import TYPE_CHECKING -from ._callbacks import install_staged_callback -from ._runtime import AdvancedMemoryRuntime +from trpc_agent_sdk.sessions.compact._callbacks import install_staged_callback +from trpc_agent_sdk.sessions.compact._runtime import AdvancedMemoryRuntime if TYPE_CHECKING: from trpc_agent_sdk.agents import LlmAgent diff --git a/trpc_agent_sdk/advanced_memory/_preload_memory.py b/trpc_agent_sdk/advanced_memory/_preload_memory.py index 88c487753..7866bb603 100644 --- a/trpc_agent_sdk/advanced_memory/_preload_memory.py +++ b/trpc_agent_sdk/advanced_memory/_preload_memory.py @@ -21,13 +21,12 @@ 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.sessions.compact._formats import memory_freshness +from trpc_agent_sdk.sessions.compact._formats import parse_memory_updated_at +from trpc_agent_sdk.sessions.compact._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 diff --git a/trpc_agent_sdk/advanced_memory/_redis_stores.py b/trpc_agent_sdk/advanced_memory/_redis_stores.py index 997a64761..f49479063 100644 --- a/trpc_agent_sdk/advanced_memory/_redis_stores.py +++ b/trpc_agent_sdk/advanced_memory/_redis_stores.py @@ -1,304 +1,5 @@ -"""Redis implementations of the Advanced Memory storage contracts.""" +"""Redis stores owned by long-term Advanced Memory.""" -from __future__ import annotations +from trpc_agent_sdk.sessions.compact._redis_stores import RedisLongTermMemoryStore -import asyncio -import json -from collections.abc import Mapping -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 AdvancedMemoryConfig -from ._formats import MemoryDocument, MemoryIndexEntry, SessionMemoryDocument -from ._paths import AdvancedMemoryPaths - -_APPEND_UNIQUE_SCRIPT = """ -if redis.call('SADD', KEYS[2], ARGV[1]) == 0 then return 0 end -redis.call('XADD', KEYS[1], '*', 'data', ARGV[2]) -return 1 -""" - -_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: AdvancedMemoryConfig, 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}}}" - self._app_base = f"{config.redis_key_prefix}:{{{app_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 _session_base(self, session_id: str) -> str: - safe_session_id = self._paths.session_dir(session_id).name - tenant = f"{self._paths.tenant_root_dir.parent.name}:{self._paths.tenant_root_dir.name}" - return f"{self._config.redis_key_prefix}:{{{tenant}:{safe_session_id}}}" - - def _session_registry(self, session_id: str) -> str: - return f"{self._session_base(session_id)}:keys" - - 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_session_ttl(self, session_id: str, *keys: str) -> None: - skip_prefixes: tuple[str, ...] = () - if not self._config.session_ttl_delete_transcripts: - skip_prefixes = (f"{self._session_base(session_id)}:transcript", ) - await self._refresh_ttl_group( - self._session_registry(session_id), - list(keys), - self._config.session_ttl_seconds, - skip_prefixes=skip_prefixes, - ) - - async def _refresh_memory_ttl(self, *keys: str) -> None: - await self._refresh_ttl_group( - self._memory_registry(), - list(keys), - self._config.memory_ttl_seconds, - ) - - async def delete_session(self, session_id: str) -> None: - """Delete all Advanced Memory keys for one session.""" - session_base = self._session_base(session_id) - registry = self._session_registry(session_id) - keys: set[str] = {registry} - tracked = await self._command("smembers", registry) or [] - keys.update(value for value in (self._text(item) for item in tracked) if value) - - cursor: Any = 0 - pattern = f"{session_base}:*" - while True: - cursor, scanned = await self._command( - "scan", - cursor, - match=pattern, - count=100, - ) - keys.update(value for value in (self._text(item) for item in scanned) if value) - if int(cursor) == 0: - break - if keys: - await self._command("delete", *keys) - - @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() - lines, used_bytes = [], 0 - for line in value.splitlines(keepends=True)[:self._config.memory_index_max_lines]: - size = len(line.encode(self._config.encoding)) - if used_bytes + size > self._config.memory_index_max_bytes: - break - lines.append(line) - used_bytes += size - return "".join(lines) - - 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] - - -class RedisSessionMemoryStore(_RedisStore): - - async def read(self, session_id: str) -> str | None: - key = f"{self._session_base(session_id)}:summary" - value = await self._command("get", key) - await self._refresh_session_ttl(session_id, key) - return self._text(value) - - async def write(self, session_id: str, document: SessionMemoryDocument) -> Path: - key = f"{self._session_base(session_id)}:summary" - await self._command("set", key, document.to_markdown()) - await self._refresh_session_ttl(session_id, key) - return Path(f"advanced-memory://{key}") - - -class RedisToolResultStore(_RedisStore): - - async def write(self, session_id: str, result_id: str, serialized_result: str) -> Path: - key = f"{self._session_base(session_id)}:tool:{result_id}" - await self._command("set", key, serialized_result) - await self._refresh_session_ttl(session_id, key) - return Path(f"advanced-memory://{key}") - - async def read(self, session_id: str, result_id: str) -> str | None: - key = f"{self._session_base(session_id)}:tool:{result_id}" - value = await self._command("get", key) - await self._refresh_session_ttl(session_id, key) - return self._text(value) - - -class RedisTranscriptStore(_RedisStore): - - async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: - payload = dict(record) - payload.setdefault("recorded_at", datetime.now(timezone.utc).isoformat()) - stream = f"{self._session_base(session_id)}:transcript" - await self._command("xadd", stream, {"data": json.dumps(payload)}) - await self._refresh_session_ttl(session_id, stream) - return Path(f"advanced-memory://{stream}") - - async def append_unique(self, session_id: str, record: Mapping[str, Any], *, unique_key: str) -> tuple[Path, bool]: - payload = dict(record) - value = payload.get(unique_key) - if not isinstance(value, str) or not 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()) - stream = f"{self._session_base(session_id)}:transcript" - seen = f"{stream}:seen:{unique_key}" - async with self._storage.create_db_session() as connection: - added = await self._storage.execute_command( - connection, - RedisCommand( - method="eval", - args=(_APPEND_UNIQUE_SCRIPT, 2, stream, seen, value, json.dumps(payload, ensure_ascii=False)), - )) - await self._refresh_session_ttl(session_id, stream, seen) - return Path(f"advanced-memory://{stream}"), bool(added) - - async def read_all(self, session_id: str) -> list[dict[str, Any]]: - stream = f"{self._session_base(session_id)}:transcript" - entries = await self._command("xrange", stream, "-", "+") - await self._refresh_session_ttl(session_id, stream) - records: list[dict[str, Any]] = [] - for _, fields in entries: - value = fields.get(b"data") if isinstance(fields, dict) else None - value = value or fields.get("data") - text = self._text(value) - if text: - records.append(json.loads(text)) - return records +__all__ = ["RedisLongTermMemoryStore"] diff --git a/trpc_agent_sdk/advanced_memory/_sql_stores.py b/trpc_agent_sdk/advanced_memory/_sql_stores.py index 46af4de11..d9173b3ce 100644 --- a/trpc_agent_sdk/advanced_memory/_sql_stores.py +++ b/trpc_agent_sdk/advanced_memory/_sql_stores.py @@ -1,569 +1,5 @@ -"""SQL implementations of the Advanced Memory storage contracts.""" +"""SQL stores owned by long-term Advanced Memory.""" -from __future__ import annotations +from trpc_agent_sdk.sessions.compact._sql_stores import SqlLongTermMemoryStore -import json -import asyncio -import hashlib -import uuid -from datetime import datetime, timedelta, timezone -from dataclasses import replace -from pathlib import Path -from collections.abc import Mapping -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 AdvancedMemoryConfig -from ._formats import MemoryDocument, MemoryIndexEntry, SessionMemoryDocument -from ._paths import AdvancedMemoryPaths - - -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 SqlSessionMemory(AdvancedMemorySqlBase): - __tablename__ = "advanced_memory_session_memory" - - 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) - session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_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 SqlTranscript(AdvancedMemorySqlBase): - __tablename__ = "advanced_memory_transcripts" - - 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) - session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) - record_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) - payload: Mapped[str] = mapped_column(Text) - recorded_at: Mapped[datetime] = mapped_column(PreciseTimestamp, default=func.now()) - expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) - - -class SqlTranscriptSeen(AdvancedMemorySqlBase): - __tablename__ = "advanced_memory_transcript_seen" - - dedupe_id: Mapped[str] = mapped_column(String(64), primary_key=True) - app_name: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) - user_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) - session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) - unique_key: Mapped[str] = mapped_column(String(DEFAULT_MAX_VARCHAR_LENGTH)) - unique_value: Mapped[str] = mapped_column(String(DEFAULT_MAX_VARCHAR_LENGTH)) - expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) - - -class SqlToolResult(AdvancedMemorySqlBase): - __tablename__ = "advanced_memory_tool_results" - - 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) - session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) - result_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_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: AdvancedMemoryConfig, 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 - - async def _refresh_session_scope(self, db: Any, session_id: str) -> None: - expiry = self._expiry(self._config.session_ttl_seconds) - if expiry is None: - return - tables = ( - (SqlSessionMemory, (self._app_name, self._user_id, session_id)), - (SqlToolResult, (self._app_name, self._user_id, session_id)), - ) - if self._config.session_ttl_delete_transcripts: - tables = ( - (SqlTranscript, (self._app_name, self._user_id, session_id)), - (SqlTranscriptSeen, (self._app_name, self._user_id, session_id)), - *tables, - ) - for model, key in tables: - rows = await self._storage.query( - db, - SqlKey(key=key, storage_cls=model), - SqlCondition(filters=[ - getattr(model, "app_name") == self._app_name, - getattr(model, "user_id") == self._user_id, - getattr(model, "session_id") == session_id, - getattr(model, "expires_at").is_(None) | (getattr(model, "expires_at") > self._now()), - ]), - ) - for row in rows: - row.expires_at = expiry - - async def delete_session(self, session_id: str) -> None: - """Delete all Advanced Memory rows for one session.""" - models = ( - SqlSessionMemory, - SqlTranscript, - SqlTranscriptSeen, - SqlToolResult, - ) - filters = { - SqlSessionMemory: [ - SqlSessionMemory.app_name == self._app_name, - SqlSessionMemory.user_id == self._user_id, - SqlSessionMemory.session_id == session_id, - ], - SqlTranscript: [ - SqlTranscript.app_name == self._app_name, - SqlTranscript.user_id == self._user_id, - SqlTranscript.session_id == session_id, - ], - SqlTranscriptSeen: [ - SqlTranscriptSeen.app_name == self._app_name, - SqlTranscriptSeen.user_id == self._user_id, - SqlTranscriptSeen.session_id == session_id, - ], - SqlToolResult: [ - SqlToolResult.app_name == self._app_name, - SqlToolResult.user_id == self._user_id, - SqlToolResult.session_id == session_id, - ], - } - async with self._storage.create_db_session() as db: - for model in models: - await self._storage.delete( - db, - SqlKey(key=tuple(), storage_cls=model), - SqlCondition(filters=filters[model]), - ) - await self._storage.commit(db) - - -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) - await self._storage.commit(db) - content = row.content - lines, used_bytes = [], 0 - for line in content.splitlines(keepends=True)[:self._config.memory_index_max_lines]: - size = len(line.encode(self._config.encoding)) - if used_bytes + size > self._config.memory_index_max_bytes: - break - lines.append(line) - used_bytes += size - return "".join(lines) - - 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 SqlSessionMemoryStore(_SqlStore): - - async def read(self, session_id: str) -> str | None: - async with self._storage.create_db_session() as db: - row = await self._storage.get( - db, SqlKey(key=(self._app_name, self._user_id, session_id), storage_cls=SqlSessionMemory)) - if row is None or self._expired(row.expires_at): - return None - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return row.content - - async def write(self, session_id: str, document: SessionMemoryDocument) -> Path: - async with self._storage.create_db_session() as db: - key = (self._app_name, self._user_id, session_id) - row = await self._storage.get(db, SqlKey(key=key, storage_cls=SqlSessionMemory)) - if row is None: - row = SqlSessionMemory(app_name=key[0], user_id=key[1], session_id=key[2]) - await self._storage.add(db, row) - row.content = document.to_markdown() - row.updated_at = self._now() - row.expires_at = self._expiry(self._config.session_ttl_seconds) - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/summary") - - -class SqlToolResultStore(_SqlStore): - - async def write(self, session_id: str, result_id: str, serialized_result: str) -> Path: - async with self._storage.create_db_session() as db: - key = (self._app_name, self._user_id, session_id, result_id) - row = await self._storage.get(db, SqlKey(key=key, storage_cls=SqlToolResult)) - if row is None: - row = SqlToolResult( - app_name=key[0], - user_id=key[1], - session_id=key[2], - result_id=key[3], - ) - await self._storage.add(db, row) - row.content = serialized_result - row.updated_at = self._now() - row.expires_at = self._expiry(self._config.session_ttl_seconds) - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/tool/{result_id}") - - async def read(self, session_id: str, result_id: str) -> str | None: - async with self._storage.create_db_session() as db: - row = await self._storage.get( - db, - SqlKey(key=(self._app_name, self._user_id, session_id, result_id), storage_cls=SqlToolResult), - ) - if row is None or self._expired(row.expires_at): - return None - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return row.content - - -class SqlTranscriptStore(_SqlStore): - - def _dedupe_id(self, session_id: str, unique_key: str, value: str) -> str: - raw = "\0".join((self._app_name, self._user_id, session_id, unique_key, value)) - return hashlib.sha256(raw.encode("utf-8")).hexdigest() - - async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: - payload = dict(record) - payload.setdefault("recorded_at", self._now().replace(tzinfo=timezone.utc).isoformat()) - async with self._storage.create_db_session() as db: - await self._storage.add( - db, - SqlTranscript( - app_name=self._app_name, - user_id=self._user_id, - session_id=session_id, - record_id=uuid.uuid4().hex, - payload=json.dumps(payload, ensure_ascii=False), - expires_at=(self._expiry(self._config.session_ttl_seconds) - if self._config.session_ttl_delete_transcripts else None), - )) - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript") - - async def append_unique( - self, - session_id: str, - record: Mapping[str, Any], - *, - unique_key: str, - ) -> tuple[Path, bool]: - payload = dict(record) - value = payload.get(unique_key) - if not isinstance(value, str) or not value: - raise ValueError(f"Transcript unique key {unique_key!r} must be a non-empty string") - async with self._storage.create_db_session() as db: - dedupe_id = self._dedupe_id(session_id, unique_key, value) - seen_key = (self._app_name, self._user_id, session_id, unique_key, value) - seen = await self._storage.get( - db, - SqlKey(key=(dedupe_id, ), storage_cls=SqlTranscriptSeen), - ) - if seen is not None and not self._expired(seen.expires_at): - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript"), False - if seen is not None: - await self._storage.delete( - db, - SqlKey(key=(dedupe_id, ), storage_cls=SqlTranscriptSeen), - SqlCondition(filters=[ - SqlTranscriptSeen.dedupe_id == dedupe_id, - ]), - ) - payload.setdefault("recorded_at", self._now().replace(tzinfo=timezone.utc).isoformat()) - await self._storage.add( - db, - SqlTranscriptSeen( - dedupe_id=dedupe_id, - app_name=seen_key[0], - user_id=seen_key[1], - session_id=seen_key[2], - unique_key=seen_key[3], - unique_value=seen_key[4], - expires_at=(self._expiry(self._config.session_ttl_seconds) - if self._config.session_ttl_delete_transcripts else None), - )) - await self._storage.add( - db, - SqlTranscript( - app_name=self._app_name, - user_id=self._user_id, - session_id=session_id, - record_id=uuid.uuid4().hex, - payload=json.dumps(payload, ensure_ascii=False), - expires_at=(self._expiry(self._config.session_ttl_seconds) - if self._config.session_ttl_delete_transcripts else None), - )) - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript"), True - - async def read_all(self, session_id: str) -> list[dict[str, Any]]: - async with self._storage.create_db_session() as db: - rows = await self._storage.query( - db, - SqlKey(key=(self._app_name, self._user_id, session_id), storage_cls=SqlTranscript), - SqlCondition( - filters=[ - SqlTranscript.app_name == self._app_name, - SqlTranscript.user_id == self._user_id, - SqlTranscript.session_id == session_id, - SqlTranscript.expires_at.is_(None) | (SqlTranscript.expires_at > self._now()), - ], - order_func=SqlTranscript.recorded_at.asc, - ), - ) - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return [json.loads(row.payload) for row in rows] - - -class SqlAdvancedMemoryCleanup: - """Periodically remove expired Advanced Memory SQL rows.""" - - _models = ( - SqlMemoryIndex, - SqlMemoryTopic, - SqlSessionMemory, - SqlTranscript, - SqlTranscriptSeen, - SqlToolResult, - ) - - def __init__(self, config: AdvancedMemoryConfig, 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 - and self._config.session_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: - models = self._models if self._config.session_ttl_delete_transcripts else tuple( - model for model in self._models if model is not SqlTranscript) - for model in models: - await self._storage.delete( - db, - SqlKey(key=tuple(), storage_cls=model), - SqlCondition(filters=[model.expires_at.is_not(None), model.expires_at <= now]), - ) - 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", - "SqlSessionMemoryStore", - "SqlToolResultStore", - "SqlTranscriptStore", -] +__all__ = ["SqlLongTermMemoryStore"] diff --git a/trpc_agent_sdk/advanced_memory/_storage.py b/trpc_agent_sdk/advanced_memory/_storage.py index c390c025f..0e0a957ae 100644 --- a/trpc_agent_sdk/advanced_memory/_storage.py +++ b/trpc_agent_sdk/advanced_memory/_storage.py @@ -1,499 +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. -"""Basic disk stores for long-term memory, session memory, and transcripts.""" +"""Local storage owned by long-term Advanced Memory.""" -from __future__ import annotations +from trpc_agent_sdk.sessions.compact._storage import LongTermMemoryStore -import asyncio -import json -import os -import shutil -import tempfile -import threading -import time -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 - - -def _is_expired(path: Path, ttl: int | None) -> bool: - if ttl is None or not path.exists(): - return False - return time.time() - path.stat().st_mtime >= ttl - - -def _touch(path: Path) -> None: - path.parent.mkdir(parents=True, exist_ok=True) - path.touch() - - -def _expire_memory_dir(memory_dir: Path, config: AdvancedMemoryConfig) -> bool: - """Expire the whole long-term memory group using index activity time.""" - index_path = memory_dir / config.memory_index_name - if not _is_expired(index_path, config.memory_ttl_seconds): - return False - for path in memory_dir.glob("*.md"): - path.unlink(missing_ok=True) - return True - - -def _refresh_memory_dir(memory_dir: Path) -> None: - """Refresh activity for every file in the long-term memory group.""" - for path in memory_dir.glob("*.md"): - _touch(path) - - -def _session_activity_path(session_dir: Path) -> Path: - return session_dir / ".advanced-memory-activity" - - -def _expire_session_dir(session_dir: Path, config: AdvancedMemoryConfig) -> bool: - """Expire all Advanced Memory data belonging to one local session.""" - if not session_dir.exists() or config.session_ttl_seconds is None: - return False - activity_path = _session_activity_path(session_dir) - if activity_path.exists(): - expired = _is_expired(activity_path, config.session_ttl_seconds) - else: - files = [path for path in session_dir.rglob("*") if path.is_file()] - expired = bool(files) and time.time() - max(path.stat().st_mtime - for path in files) >= config.session_ttl_seconds - if expired: - if config.session_ttl_delete_transcripts: - shutil.rmtree(session_dir, ignore_errors=True) - else: - transcript_path = session_dir / config.transcript_name - for child in session_dir.iterdir(): - if child == transcript_path: - continue - if child.is_dir(): - shutil.rmtree(child, ignore_errors=True) - else: - child.unlink(missing_ok=True) - return expired - - -def _refresh_session_dir(session_dir: Path) -> None: - _touch(_session_activity_path(session_dir)) - - -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 _expire_memory_dir(self._paths.memory_dir, self._config) or not self.index_path.exists(): - return "" - _refresh_memory_dir(self._paths.memory_dir) - 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) - _refresh_memory_dir(self._paths.memory_dir) - - 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_topic_sync, path) - - def _read_topic_sync(self, path: Path) -> str | None: - if _expire_memory_dir(self._paths.memory_dir, self._config) or not path.exists(): - return None - _refresh_memory_dir(self._paths.memory_dir) - return path.read_text(encoding=self._config.encoding) - - 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_sync, path) - - def _read_frontmatter_sync(self, path: Path) -> str | None: - """Synchronously read a topic's bounded frontmatter block.""" - if _expire_memory_dir(self._paths.memory_dir, self._config) or not path.exists(): - return None - _refresh_memory_dir(self._paths.memory_dir) - 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(self._write_topic_sync, path, document.to_markdown()) - return path - - def _write_topic_sync(self, path: Path, content: str) -> None: - _expire_memory_dir(self._paths.memory_dir, self._config) - _atomic_write_text(path, content, encoding=self._config.encoding) - _refresh_memory_dir(self._paths.memory_dir) - - 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 _expire_memory_dir(self._paths.memory_dir, self._config): - return [] - if not self._paths.memory_dir.exists(): - return [] - _refresh_memory_dir(self._paths.memory_dir) - 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, session_id, path) - - def _read_sync(self, session_id: str, path: Path) -> str | None: - """Synchronously read session memory.""" - session_dir = self._paths.session_dir(session_id) - if _expire_session_dir(session_dir, self._config) or not path.exists(): - return None - _refresh_session_dir(session_dir) - 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( - self._write_sync, - session_id, - path, - document.to_markdown(), - ) - return path - - def _write_sync(self, session_id: str, path: Path, content: str) -> None: - session_dir = self._paths.session_dir(session_id) - _expire_session_dir(session_dir, self._config) - _atomic_write_text(path, content, encoding=self._config.encoding) - _refresh_session_dir(session_dir) - - -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( - self._write_sync, - session_id, - path, - serialized_result, - ) - 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, session_id, path) - - def _read_sync(self, session_id: str, path: Path) -> str | None: - """Synchronously read an optional complete tool-result file.""" - session_dir = self._paths.session_dir(session_id) - if _expire_session_dir(session_dir, self._config) or not path.exists(): - return None - _refresh_session_dir(session_dir) - return path.read_text(encoding=self._config.encoding) - - def _write_sync(self, session_id: str, path: Path, content: str) -> None: - session_dir = self._paths.session_dir(session_id) - _expire_session_dir(session_dir, self._config) - _atomic_write_text(path, content, encoding=self._config.encoding) - _refresh_session_dir(session_dir) - - -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.""" - _expire_session_dir(path.parent, self._config) - path.parent.mkdir(parents=True, exist_ok=True) - with self._write_lock: - self._append_serialized_unlocked(path, serialized) - _refresh_session_dir(path.parent) - - 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.""" - with self._write_lock: - if _expire_session_dir(path.parent, self._config): - for cache_key in list(self._seen_unique_values): - if cache_key[0] == path: - self._seen_unique_values.pop(cache_key, None) - path.parent.mkdir(parents=True, exist_ok=True) - cache_key = (path, unique_key) - 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) - _refresh_session_dir(path.parent) - 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: - expired = _expire_session_dir(path.parent, self._config) - if expired and self._config.session_ttl_delete_transcripts: - return [] - if not path.exists(): - return [] - _refresh_session_dir(path.parent) - 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 - - -class LocalAdvancedMemoryCleanup: - """Periodically remove expired local Advanced Memory data.""" - - def __init__(self, config: AdvancedMemoryConfig) -> 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: - return - if self._config.memory_ttl_seconds is None and self._config.session_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] - session_roots = [root / self._config.session_dir_name] - tenants_root = root / "tenants" - if tenants_root.exists(): - for app_dir in tenants_root.iterdir(): - if app_dir.is_dir(): - for user_dir in app_dir.iterdir(): - if user_dir.is_dir(): - memory_dirs.append(user_dir / self._config.memory_dir_name) - session_roots.append(user_dir / self._config.session_dir_name) - for memory_dir in memory_dirs: - _expire_memory_dir(memory_dir, self._config) - for session_root in session_roots: - if session_root.exists(): - for session_dir in session_root.iterdir(): - if session_dir.is_dir(): - _expire_session_dir(session_dir, self._config) - - async def _run(self) -> None: - if self._stop_event is None: - return - ttls = [ - ttl for ttl in ( - self._config.memory_ttl_seconds, - self._config.session_ttl_seconds, - ) if ttl is not None - ] - interval = min(ttls) if ttls else 60 - try: - while not self._stop_event.is_set(): - try: - await asyncio.wait_for(self._stop_event.wait(), timeout=interval) - break - except asyncio.TimeoutError: - await self.cleanup_once() - except asyncio.CancelledError: - raise - - async def close(self) -> None: - if self._task is not None: - await self.cleanup_once() - 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 +__all__ = ["LongTermMemoryStore"] diff --git a/trpc_agent_sdk/advanced_memory/_storage_backend.py b/trpc_agent_sdk/advanced_memory/_storage_backend.py index 4e10a2f0c..19a5b2729 100644 --- a/trpc_agent_sdk/advanced_memory/_storage_backend.py +++ b/trpc_agent_sdk/advanced_memory/_storage_backend.py @@ -8,8 +8,8 @@ from typing import Protocol -from ._paths import MemoryScope -from ._runtime import ScopedAdvancedMemoryRuntime +from trpc_agent_sdk.sessions.compact._paths import MemoryScope +from trpc_agent_sdk.sessions.compact._runtime import ScopedAdvancedMemoryRuntime class AdvancedMemoryStorageBackend(Protocol): 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..d93e9dabc 100644 --- a/trpc_agent_sdk/memory/__init__.py +++ b/trpc_agent_sdk/memory/__init__.py @@ -27,7 +27,7 @@ __all__ = [ "BaseMemoryService", "MemoryServiceConfig", - "AdvancedMemoryConfig", + "AdvancedCompactConfig", "AdvancedMemoryService", "EventTtl", "InMemoryMemoryService", @@ -43,8 +43,8 @@ def __getattr__(name: str): """Lazily expose Advanced Memory configuration without import cycles.""" - if name == "AdvancedMemoryConfig": - from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig + if name == "AdvancedCompactConfig": + from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig - return AdvancedMemoryConfig + return AdvancedCompactConfig 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 5ad1dbdd5..dc00d31c6 100644 --- a/trpc_agent_sdk/memory/_advanced_memory_service.py +++ b/trpc_agent_sdk/memory/_advanced_memory_service.py @@ -19,51 +19,42 @@ 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 AdvancedCompactConfig from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime + from trpc_agent_sdk.advanced_memory import LongTermMemoryIntegration class AdvancedMemoryService(BaseMemoryService): - """Expose Advanced Memory through the standard Runner memory API. + """Expose user-scoped 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. Session compression is + configured independently with ``setup_context_compression``. """ def __init__( self, - config: AdvancedMemoryConfig | None = None, + config: AdvancedCompactConfig | 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 AdvancedCompactConfig from trpc_agent_sdk.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 AdvancedCompactConfig()) 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) -> AdvancedCompactConfig: """Return the Advanced Memory configuration.""" return self._runtime.config @@ -73,45 +64,34 @@ 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.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 async def store_session( self, session: Session, agent_context: Optional[AgentContext] = None, ) -> None: - """Keep the standard Runner post-turn contract without duplicating work. - - The wrapped session service performs session-memory extraction from - ``create_session_summary`` before Runner reaches this method. - """ + """Long-term Memory is updated explicitly through its tools.""" return None async def search_memory( @@ -129,9 +109,5 @@ async def search_memory( return SearchMemoryResponse() 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. - """ + """Release service-owned local or external storage resources.""" await self._runtime.close() diff --git a/trpc_agent_sdk/runners.py b/trpc_agent_sdk/runners.py index e93023cb5..083c36803 100644 --- a/trpc_agent_sdk/runners.py +++ b/trpc_agent_sdk/runners.py @@ -227,12 +227,13 @@ 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_config = getattr(session_service, "session_compact_config", None) + from trpc_agent_sdk.sessions.compact import BaseSessionCompactConfig + if isinstance(compact_config, BaseSessionCompactConfig): + compact_config.setup(agent, session_service) self.app_name = app_name self.agent = agent self.artifact_service = artifact_service diff --git a/trpc_agent_sdk/sessions/__init__.py b/trpc_agent_sdk/sessions/__init__.py index 1b8d84418..5501aed23 100644 --- a/trpc_agent_sdk/sessions/__init__.py +++ b/trpc_agent_sdk/sessions/__init__.py @@ -53,7 +53,13 @@ "ListSessionsResponse", "State", "BaseSessionService", - "AdvancedMemorySessionService", + "BaseSessionCompactManager", + "BaseSessionCompactConfig", + "AdvancedCompactConfig", + "AdvancedSessionCompactManager", + "AutoCompact", + "setup_advanced_session_compact", + "setup_context_compression", "HistoryRecord", "InMemorySessionService", "SessionWithTTL", @@ -92,8 +98,16 @@ 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", + "BaseSessionCompactConfig", + "setup_advanced_session_compact", + "setup_context_compression", + }: + 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 c134fb545..000000000 --- a/trpc_agent_sdk/sessions/_advanced_memory_session_service.py +++ /dev/null @@ -1,440 +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 _scoped_runtime(self, app_name: str, user_id: str) -> AdvancedMemoryRuntime: - return self._runtime.for_scope(app_name, user_id) - - def _metadata_path(self, app_name: str, user_id: str, session_id: str) -> Path: - return self._scoped_runtime(app_name, user_id).paths.session_dir(session_id) / "session.json" - - def _app_state_path(self, app_name: str, user_id: str) -> Path: - """Return state shared by every user of one app.""" - return self._scoped_runtime(app_name, user_id).paths.tenant_root_dir.parent / "_state.json" - - def _user_state_path(self, app_name: str, user_id: str) -> Path: - """Return state private to one application user.""" - return self._scoped_runtime(app_name, user_id).paths.tenant_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.app_name, session.user_id, session.id) - await asyncio.to_thread(self._write_json, path, payload, self._runtime.config.encoding) - await asyncio.to_thread(path.parent.joinpath(".advanced-memory-activity").touch, exist_ok=True) - - @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, app_name: str, user_id: str, session_id: str) -> Session | None: - path = self._metadata_path(app_name, user_id, 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) - await asyncio.to_thread(path.parent.joinpath(".advanced-memory-activity").touch, exist_ok=True) - 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 - tenants_root = self._runtime.config.root_dir / "tenants" - if not tenants_root.exists(): - return - for metadata_path in tenants_root.glob(f"*/*/{self._runtime.config.session_dir_name}/*/session.json"): - try: - if metadata_path.stat().st_mtime < cutoff: - session_dir = metadata_path.parent - if self._runtime.config.session_ttl_delete_transcripts: - shutil.rmtree(session_dir, ignore_errors=True) - else: - transcript_path = session_dir / self._runtime.config.transcript_name - for child in session_dir.iterdir(): - if child == transcript_path: - continue - if child.is_dir(): - shutil.rmtree(child, ignore_errors=True) - else: - child.unlink(missing_ok=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, app_name: str, user_id: str) -> dict[str, dict[str, Any]]: - - async def read(path: Path) -> dict[str, Any]: - if not path.exists(): - return {} - payload = await asyncio.to_thread(path.read_text, encoding=self._runtime.config.encoding) - return dict(json.loads(payload)) - - return { - "app": await read(self._app_state_path(app_name, user_id)), - "user": await read(self._user_state_path(app_name, user_id)), - } - - async def _write_global_state(self, app_name: str, user_id: str, state: dict[str, dict[str, Any]]) -> None: - await asyncio.to_thread( - self._write_json, - self._app_state_path(app_name, user_id), - state["app"], - self._runtime.config.encoding, - ) - await asyncio.to_thread( - self._write_json, - self._user_state_path(app_name, user_id), - state["user"], - self._runtime.config.encoding, - ) - - async def _restore_events(self, session: Session) -> Session: - records = await self._scoped_runtime(session.app_name, session.user_id).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._scoped_runtime(app_name, user_id).initialize() - await self._read_session(app_name, user_id, resolved_id) - global_state = await self._read_global_state(app_name, user_id) - global_state["app"].update(state_delta.app_state_delta) - global_state["user"].update(state_delta.user_state_delta) - await self._write_global_state(app_name, user_id, 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"].items()}) - session.state.update({f"user:{key}": value for key, value in global_state["user"].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(app_name, user_id, session_id) - if session is None: - return None - global_state = await self._read_global_state(app_name, user_id) - app_state = global_state["app"] - user_state = global_state["user"] - 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() - tenants_root = self._runtime.config.root_dir / "tenants" - if not tenants_root.exists(): - return ListSessionsResponse() - sessions: list[Session] = [] - if user_id is not None: - root = self._scoped_runtime(app_name, user_id).paths.session_root_dir - session_glob = "*/session.json" - else: - root = tenants_root - session_glob = f"*/*/{self._runtime.config.session_dir_name}/*/session.json" - for path in await asyncio.to_thread(lambda: list(root.glob(session_glob))): - 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._metadata_path(app_name, user_id, session_id).parent, 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(session.app_name, session.user_id) - global_state["app"].update(state_delta.app_state_delta) - global_state["user"].update(state_delta.user_state_delta) - await self._write_global_state(session.app_name, session.user_id, 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: - runtime = self._scoped_runtime(session.app_name, session.user_id) - records = await runtime.transcripts.read_all(session.id) - record = build_event_transcript_record( - session, - persisted, - parent_event_id=find_last_event_id(records), - ) - await 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() - await self._runtime.close() - - -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) - if self._runtime.config.storage_backend == "redis": - raise ValueError("AdvancedMemorySessionService is file-backed; use RedisSessionService with " - "AdvancedMemoryService when AdvancedMemoryConfig.storage_backend='redis'") - 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..8827b2f2f 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,10 @@ from ._summarizer_manager import SummarizerSessionManager from ._types import SessionServiceConfig +if TYPE_CHECKING: + from .compact import BaseSessionCompactManager + from .compact import BaseSessionCompactConfig + class BaseSessionService(SessionServiceABC): """Abstract base class for session management services. @@ -45,14 +50,25 @@ class BaseSessionService(SessionServiceABC): def __init__(self, summarizer_manager: Optional[SummarizerSessionManager] = None, - session_config: Optional[SessionServiceConfig] = None): + session_config: Optional[SessionServiceConfig] = None, + session_compact_config: Optional["BaseSessionCompactConfig"] = 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_config: Optional Advanced Compact configuration + session_compact_manager: Optional pluggable Session Compact manager """ + if session_compact_config is not None and session_compact_manager is not None: + raise ValueError( + "Provide either session_compact_config or " + "session_compact_manager, not both" + ) self._summarizer_manager = summarizer_manager + self._session_compact_config = session_compact_config + 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 +76,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 +89,16 @@ def session_config(self) -> SessionServiceConfig: """Get the session service configuration.""" return self._session_config + @property + def session_compact_config(self) -> Optional["BaseSessionCompactConfig"]: + """Return deferred Session Compact configuration, if configured.""" + return self._session_compact_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 +106,31 @@ 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 +223,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,8 +240,25 @@ 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 + async def _delete_session_compact_data( + self, + *, + app_name: str, + user_id: str, + session_id: str, + ) -> None: + """Delete side data owned by the configured compact manager.""" + if self._session_compact_manager: + await self._session_compact_manager.delete_session( + app_name=app_name, + user_id=user_id, + session_id=session_id, + ) + def filter_events(self, session: Session, need_copy: bool = False) -> Session: """Filter events based on the session config. @@ -211,4 +279,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..642e06135 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,10 @@ from ._utils import extract_state_delta from ._utils import merge_state +if TYPE_CHECKING: + from .compact._base_manager import BaseSessionCompactManager + from .compact._base_config import BaseSessionCompactConfig + class SessionWithTTL(BaseModel): """Wrapper for session with TTL support.""" @@ -108,8 +113,15 @@ 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_config: "BaseSessionCompactConfig | None" = None, + session_compact_manager: BaseSessionCompactManager | None = None): + super().__init__( + summarizer_manager=summarizer_manager, + session_config=session_config, + session_compact_config=session_compact_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 +225,13 @@ 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] + await self._delete_session_compact_data( + app_name=app_name, + user_id=user_id, + session_id=session_id, + ) @override async def append_event(self, session: Session, event: Event) -> Event: @@ -294,6 +310,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..650c7188c 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,10 @@ from ._utils import session_key from ._utils import user_state_key +if TYPE_CHECKING: + from .compact._base_manager import BaseSessionCompactManager + from .compact._base_config import BaseSessionCompactConfig + def _session_key_prefix(app_name: str, user_id: Optional[str] = None) -> str: """Generate a Redis key prefix for listing sessions. @@ -54,6 +60,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 +94,34 @@ def __init__(self, summarizer_manager: Optional[SummarizerSessionManager] = None, session_config: Optional[SessionServiceConfig] = None, is_async: bool = False, + session_compact_config: "BaseSessionCompactConfig | None" = 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_config=session_compact_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. @@ -186,6 +220,11 @@ async def delete_session(self, *, app_name: str, user_id: str, session_id: str) async with self._redis_storage.create_db_session() as redis_session: key = session_key(app_name, user_id, session_id) await self._redis_storage.delete(redis_session, key) + await self._delete_session_compact_data( + app_name=app_name, + user_id=user_id, + session_id=session_id, + ) @override async def append_event(self, session: Session, event: Event) -> Event: @@ -251,6 +290,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 +518,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..41335af34 100644 --- a/trpc_agent_sdk/sessions/_session.py +++ b/trpc_agent_sdk/sessions/_session.py @@ -136,3 +136,53 @@ 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 2d6d9d3ad..5cfb3e9f8 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,10 @@ from ._utils import extract_state_delta from ._utils import merge_state +if TYPE_CHECKING: + from .compact._base_manager import BaseSessionCompactManager + from .compact._base_config import BaseSessionCompactConfig + 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,9 +396,18 @@ def __init__(self, summarizer_manager: Optional[SummarizerSessionManager] = None, is_async: bool = False, session_config: Optional[SessionServiceConfig] = None, + session_compact_config: "BaseSessionCompactConfig | None" = 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_config=session_compact_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 @@ -407,6 +421,16 @@ def __init__(self, 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, @@ -533,6 +557,11 @@ async def delete_session(self, app_name: str, user_id: str, session_id: str) -> session_key = SqlKey(key=(app_name, user_id, session_id), storage_cls=StorageSession) await self._sql_storage.delete(sql_session, session_key, conditions) await self._sql_storage.commit(sql_session) + await self._delete_session_compact_data( + app_name=app_name, + user_id=user_id, + session_id=session_id, + ) @override async def append_event(self, session: Session, event: Event) -> Event: @@ -547,7 +576,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 @@ -655,6 +687,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() diff --git a/trpc_agent_sdk/sessions/compact/__init__.py b/trpc_agent_sdk/sessions/compact/__init__.py new file mode 100644 index 000000000..4551f8745 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/__init__.py @@ -0,0 +1,116 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Canonical context-compression package for session management.""" + +from ._autocompact import AutoCompact +from ._autocompact import AutoCompactCallback +from ._autocompact import AutoCompactResult +from ._autocompact import content_signature +from ._autocompact import ForkedLegacySummaryGenerator +from ._autocompact import setup_autocompact +from ._base_manager import BaseSessionCompactManager +from ._base_config import BaseSessionCompactConfig +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 ._integration import setup_advanced_session_compact +from ._integration import setup_context_compression +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 ._paths import MemoryScope +from ._runtime import AdvancedMemoryRuntime +from ._runtime import ScopedAdvancedMemoryRuntime +from ._session_memory import build_session_memory_prompt +from ._session_memory import ForkedSessionMemoryGenerator +from ._session_memory import has_session_memory_content +from ._session_memory import limit_session_memory_document +from ._session_memory import SessionMemoryExtractionInput +from ._session_memory import SessionMemoryExtractionResult +from ._session_memory import SessionMemoryExtractor +from ._session_service import TranscriptSessionService +from ._storage import SessionMemoryStore +from ._storage import ToolResultStore +from ._storage import TranscriptStore +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 ._transcript import TRANSCRIPT_SCHEMA_VERSION + +__all__ = [ + "AdvancedCompactConfig", + "BaseSessionCompactConfig", + "AdvancedMemoryPaths", + "AdvancedMemoryRuntime", + "AutoCompact", + "AutoCompactCallback", + "AutoCompactResult", + "ContextBudget", + "ContextTokenEstimate", + "ForkedLegacySummaryGenerator", + "ForkedSessionMemoryGenerator", + "HeuristicTokenEstimator", + "HistorySnip", + "HistorySnipCallback", + "HistorySnipResult", + "MemoryScope", + "Microcompact", + "MicrocompactCallback", + "MicrocompactResult", + "ModelContextWindowResolver", + "ScopedAdvancedMemoryRuntime", + "SESSION_MEMORY_SECTION_DESCRIPTIONS", + "SESSION_MEMORY_SECTIONS", + "SESSION_MEMORY_STATE_KEY", + "SessionMemoryDocument", + "SessionMemoryExtractionInput", + "SessionMemoryExtractionResult", + "SessionMemoryExtractor", + "SessionMemoryStore", + "BaseSessionCompactManager", + "AdvancedSessionCompactManager", + "TokenContextTracker", + "TokenEstimator", + "ToolResultBudget", + "ToolResultBudgetCallback", + "ToolResultBudgetResult", + "ToolResultStore", + "TRANSCRIPT_SCHEMA_VERSION", + "TranscriptSessionService", + "TranscriptStore", + "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_advanced_session_compact", + "setup_context_compression", + "setup_history_snip", + "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 75% rename from trpc_agent_sdk/advanced_memory/_autocompact.py rename to trpc_agent_sdk/sessions/compact/_autocompact.py index 070607009..045d3052d 100644 --- a/trpc_agent_sdk/advanced_memory/_autocompact.py +++ b/trpc_agent_sdk/sessions/compact/_autocompact.py @@ -19,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 @@ -27,7 +28,9 @@ 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 ._token_budget import TokenContextTracker @@ -36,6 +39,7 @@ 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 = ( @@ -69,6 +73,8 @@ class AutoCompactRecord: boundary_occurrence: int summary: str source: str + boundary_event_id: str | None = None + compaction_id: str | None = None @dataclass @@ -227,12 +233,14 @@ def __init__( 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"] = {} @@ -242,6 +250,15 @@ def runtime(self) -> AdvancedMemoryRuntime: """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.""" key = self._runtime.session_key(session_id) if hasattr(self._runtime, "session_key") else session_id @@ -273,6 +290,10 @@ async def _load_state(self, session_id: str) -> AutoCompactState: occurrence, summary, source, + record.get("boundary_event_id") + if isinstance(record.get("boundary_event_id"), str) else None, + record.get("compaction_id") + if isinstance(record.get("compaction_id"), str) else None, ) failures = 0 elif record.get("kind") == "autocompact-failure": @@ -290,6 +311,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.""" + if self._runtime.config.storage_backend in {"redis", "sql"}: + return (f"{summary.rstrip()}\n\n" + "For exact content from before compaction, read the original " + "SessionService Events. Current session memory is stored in " + f"session.state[{SESSION_MEMORY_STATE_KEY!r}].") return (f"{summary.rstrip()}\n\n" "For exact content from before compaction, read the complete transcript: " f"{self._runtime.paths.storage_reference('transcript', session_id=session_id)}\n" @@ -311,6 +337,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], @@ -376,21 +413,44 @@ def _apply_record(self, request: "LlmRequest", record: AutoCompactRecord) -> boo async def _latest_session_memory_record( self, session_id: str, - ) -> tuple[str, str] | None: + ctx: "InvocationContext", + ) -> tuple[str, str, int, str] | None: """Read session memory and its checkpoint Event for model-free compaction.""" + if self._runtime.config.storage_backend in {"redis", "sql"}: + parsed = parse_session_memory_state(ctx.session.state.get(SESSION_MEMORY_STATE_KEY)) + if parsed is None: + return None + document, checkpoint, _ = parsed + signature = checkpoint.get("boundary_signature") + occurrence = checkpoint.get("boundary_occurrence") + event_id = checkpoint.get("last_event_id") + if (not isinstance(signature, str) or not isinstance(occurrence, int) or occurrence <= 0 + or not isinstance(event_id, str)): + return None + memory = document.to_markdown() + if memory.strip() == SessionMemoryDocument().to_markdown().strip(): + return None + return memory, signature, occurrence, event_id async with self._runtime.coordination.guard( session_id, timeout=self._runtime.config.session_memory_wait_timeout_seconds, ) as acquired: if not acquired: return None + if self._runtime.session_memory is None: + 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"] + boundary = self._event_content_signature( + records, + record["last_event_id"], + ) + if boundary is not None: + return memory, boundary[0], boundary[1], record["last_event_id"] return None def _event_content_signature( @@ -423,6 +483,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]) @@ -439,7 +500,97 @@ 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. setup_context_compression always + # 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.""" @@ -486,14 +637,16 @@ async def _persist_success( token_source: str | None = None, ) -> None: """Persist a successful compaction and reset the circuit-breaker count.""" + compaction_id = record.compaction_id or f"autocompact:{uuid.uuid4().hex}" await self._runtime.transcripts.append( session_id, { "schema_version": AUTOCOMPACT_SCHEMA_VERSION, "kind": "autocompact-success", - "compaction_id": f"autocompact:{uuid.uuid4().hex}", + "compaction_id": compaction_id, "boundary_signature": record.boundary_signature, "boundary_occurrence": record.boundary_occurrence, + "boundary_event_id": record.boundary_event_id, "summary": record.summary, "source": record.source, "request_chars_before": before_chars, @@ -581,6 +734,9 @@ async def _apply_scoped( 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: @@ -619,38 +775,46 @@ async def _apply_scoped( 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 and self._session_memory_extractor.uses_session_state): + 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, - ) - target_reached = (tracker.budget(request, ctx).estimate.tokens - <= token_budget_before.warning_threshold_tokens if token_mode else - estimate_request_chars(request) <= config.autocompact_target_chars) - if not target_reached: - request.contents = [content.model_copy(deep=True) for content in original_contents] - compact_record = None + if 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, + ) + target_reached = (tracker.budget( + request, ctx).estimate.tokens <= token_budget_before.warning_threshold_tokens if token_mode + else estimate_request_chars(request) <= config.autocompact_target_chars) + if not target_reached: + request.contents = [content.model_copy(deep=True) for content in original_contents] + compact_record = None if compact_record is None: keep_count = min( @@ -673,22 +837,27 @@ async def _apply_scoped( ), 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: - 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") + 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") + await self._persist_session_compaction(ctx, compact_record) 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, + comparison_tokens_before if token_mode else None, + comparison_tokens_after if token_mode else None, + "estimated" if token_mode else None, ) state.latest_compaction = compact_record state.consecutive_failures = 0 @@ -700,9 +869,9 @@ async def _apply_scoped( 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 diff --git a/trpc_agent_sdk/sessions/compact/_base_config.py b/trpc_agent_sdk/sessions/compact/_base_config.py new file mode 100644 index 000000000..71a90ca08 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/_base_config.py @@ -0,0 +1,29 @@ +# 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 configuration contract for Session Compact strategies.""" + +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 ._base_manager import BaseSessionCompactManager + + +class BaseSessionCompactConfig(ABC): + """Create and attach one concrete Session Compact strategy.""" + + @abstractmethod + def setup( + self, + agent: Any, + session_service: Any, + ) -> "BaseSessionCompactManager": + """Create the strategy manager and attach it to the SessionService.""" 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..f7a38bb9a --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/_base_manager.py @@ -0,0 +1,57 @@ +# 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 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 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 delete_session( + self, + *, + app_name: str, + user_id: str, + session_id: str, + ) -> None: + """Delete side data owned by this manager for one session.""" + + @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 100% rename from trpc_agent_sdk/advanced_memory/_callbacks.py rename to trpc_agent_sdk/sessions/compact/_callbacks.py diff --git a/trpc_agent_sdk/advanced_memory/_config.py b/trpc_agent_sdk/sessions/compact/_config.py similarity index 94% rename from trpc_agent_sdk/advanced_memory/_config.py rename to trpc_agent_sdk/sessions/compact/_config.py index 35881cb64..456ff3dcd 100644 --- a/trpc_agent_sdk/advanced_memory/_config.py +++ b/trpc_agent_sdk/sessions/compact/_config.py @@ -14,6 +14,8 @@ from typing import Any from typing import Literal +from ._base_config import BaseSessionCompactConfig + DEFAULT_COMPACTABLE_TOOL_NAMES = ( "Read", "Bash", @@ -97,8 +99,8 @@ def _validate_path_components(values: tuple[str, ...]) -> None: @dataclass(frozen=True) -class AdvancedMemoryConfig: - """Configure the independent memory directory and storage limits.""" +class AdvancedCompactConfig(BaseSessionCompactConfig): + """Configure Advanced Session Compact and its shared memory runtime.""" enabled: bool = True root_dir: Path = field(default_factory=Path.cwd) @@ -177,8 +179,22 @@ class AdvancedMemoryConfig: preload_memory_candidate_limit: int = 200 session_ttl_delete_transcripts: bool = False + def setup(self, agent: Any, session_service: Any) -> Any: + """Create and attach the Advanced Session Compact manager.""" + from ._integration import setup_advanced_session_compact + + return setup_advanced_session_compact( + agent, + session_service, + self, + ) + 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: 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..ece6c8f28 100644 --- a/trpc_agent_sdk/advanced_memory/_formats.py +++ b/trpc_agent_sdk/sessions/compact/_formats.py @@ -8,7 +8,9 @@ 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 +139,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 +193,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 100% rename from trpc_agent_sdk/advanced_memory/_history_snip.py rename to trpc_agent_sdk/sessions/compact/_history_snip.py diff --git a/trpc_agent_sdk/sessions/compact/_integration.py b/trpc_agent_sdk/sessions/compact/_integration.py new file mode 100644 index 000000000..f83da03bb --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/_integration.py @@ -0,0 +1,155 @@ +# 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 the context-compression pipeline.""" + +from __future__ import annotations + +from dataclasses import replace +from typing import Any +from typing import TYPE_CHECKING + +from ._autocompact import LegacySummaryGenerator +from ._autocompact import setup_autocompact +from ._history_snip import setup_history_snip +from ._microcompact import setup_microcompact +from ._runtime import AdvancedMemoryRuntime +from ._config import AdvancedCompactConfig +from ._manager import AdvancedSessionCompactManager +from ._session_memory import SessionMemoryExtractor +from ._session_memory import SessionMemoryGenerator +from ._tool_result_budget import setup_tool_result_budget + +if TYPE_CHECKING: + from trpc_agent_sdk.agents import LlmAgent + from trpc_agent_sdk.sessions import SessionServiceABC + + +def setup_context_compression( + agent: "LlmAgent", + session_service: "SessionServiceABC", + memory_runtime: AdvancedMemoryRuntime, + summary_generator: LegacySummaryGenerator | None = None, + *, + compact_model: Any | None = None, + session_memory_generator: SessionMemoryGenerator | None = None, + session_memory_model: Any | None = None, +) -> "SessionServiceABC": + """Install native Session compression on an existing SessionService. + + The original service remains responsible for persistence. Session Compact + is attached through the BaseSessionService manager lifecycle. + """ + session_config = getattr(session_service, "session_config", None) + if session_config is None or not getattr(session_config, "store_historical_events", False): + raise ValueError( + "Context compression requires " + "SessionServiceConfig(store_historical_events=True)" + ) + if getattr(session_service, "summarizer_manager", None) is not None: + raise ValueError( + "Context compression and SummarizerSessionManager are mutually exclusive" + ) + + manager = getattr(session_service, "session_compact_manager", None) + if manager is not None: + if not isinstance(manager, AdvancedSessionCompactManager): + raise ValueError( + "Advanced context compression requires an " + "AdvancedSessionCompactManager" + ) + if manager.runtime is not memory_runtime: + raise ValueError("Context compression session service uses another runtime") + extractor = manager.session_memory_extractor + 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: + attach_manager = getattr(session_service, "set_session_compact_manager", None) + if not callable(attach_manager): + raise TypeError( + "Context compression requires a BaseSessionService with " + "set_session_compact_manager()" + ) + extractor = SessionMemoryExtractor( + memory_runtime, + session_memory_generator, + model=session_memory_model, + ) + manager = AdvancedSessionCompactManager( + memory_runtime, + extractor, + ) + attach_manager(manager) + setup_tool_result_budget(agent, memory_runtime) + setup_history_snip(agent, memory_runtime) + setup_microcompact(agent, memory_runtime) + autocompact = setup_autocompact( + agent, + memory_runtime, + summary_generator, + model=compact_model, + ) + autocompact.attach_session_memory_extractor(extractor) + return session_service + + +def setup_advanced_session_compact( + agent: Any, + session_service: "SessionServiceABC", + compact_config: AdvancedCompactConfig, + *, + summary_generator: LegacySummaryGenerator | None = None, + compact_model: Any | None = None, + session_memory_generator: SessionMemoryGenerator | None = None, + session_memory_model: Any | None = None, +) -> AdvancedSessionCompactManager: + """Configure Advanced Compact from a standard SessionService backend.""" + from trpc_agent_sdk.sessions import InMemorySessionService + from trpc_agent_sdk.sessions import RedisSessionService + from trpc_agent_sdk.sessions import SqlSessionService + + if isinstance(session_service, RedisSessionService): + resolved_config = replace( + compact_config, + storage_backend="redis", + redis_url=session_service.db_url, + redis_is_async=session_service.is_async, + ) + elif isinstance(session_service, SqlSessionService): + resolved_config = replace( + compact_config, + storage_backend="sql", + sql_url=session_service.db_url, + sql_is_async=session_service.is_async, + ) + elif isinstance(session_service, InMemorySessionService): + resolved_config = replace(compact_config, storage_backend="local") + else: + raise TypeError( + "Advanced Compact supports InMemorySessionService, " + "RedisSessionService, and SqlSessionService" + ) + runtime = AdvancedMemoryRuntime.create(resolved_config) + extractor = SessionMemoryExtractor( + runtime, + session_memory_generator, + model=session_memory_model, + ) + manager = AdvancedSessionCompactManager(runtime, extractor) + setup_tool_result_budget(agent, runtime) + setup_history_snip(agent, runtime) + setup_microcompact(agent, runtime) + autocompact = setup_autocompact( + agent, + runtime, + summary_generator, + model=compact_model, + ) + autocompact.attach_session_memory_extractor(extractor) + session_service.set_session_compact_manager(manager) + return manager diff --git a/trpc_agent_sdk/sessions/compact/_manager.py b/trpc_agent_sdk/sessions/compact/_manager.py new file mode 100644 index 000000000..ad0b0e6ef --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/_manager.py @@ -0,0 +1,102 @@ +# 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 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 ._runtime import AdvancedMemoryRuntime + from ._session_memory import SessionMemoryExtractor + + +class AdvancedSessionCompactManager(BaseSessionCompactManager): + """Coordinate Advanced Compact state without wrapping a SessionService.""" + + def __init__( + self, + runtime: "AdvancedMemoryRuntime", + session_memory_extractor: "SessionMemoryExtractor", + ) -> None: + """Store the compact runtime and post-turn memory extractor.""" + self._runtime = runtime + self._session_memory_extractor = session_memory_extractor + self._session_service: SessionServiceABC | None = None + + @property + def runtime(self) -> "AdvancedMemoryRuntime": + """Return the runtime shared by all compact stages.""" + return self._runtime + + @property + def session_memory_extractor(self) -> "SessionMemoryExtractor": + """Return the post-turn Session Memory extractor.""" + 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 + 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: + 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() + runtime = self._runtime.for_session(session) + if runtime.session_memory is None: + return None + return await runtime.session_memory.read(session.id) + + async def delete_session( + self, + *, + app_name: str, + user_id: str, + session_id: str, + ) -> None: + """Delete compact side data after the framework Session is deleted.""" + await self._runtime.for_scope(app_name, user_id).delete_session(session_id) + + async def close(self) -> None: + """Release Compact backend resources owned by this manager.""" + await self._runtime.close() diff --git a/trpc_agent_sdk/advanced_memory/_microcompact.py b/trpc_agent_sdk/sessions/compact/_microcompact.py similarity index 100% rename from trpc_agent_sdk/advanced_memory/_microcompact.py rename to trpc_agent_sdk/sessions/compact/_microcompact.py diff --git a/trpc_agent_sdk/advanced_memory/_paths.py b/trpc_agent_sdk/sessions/compact/_paths.py similarity index 96% rename from trpc_agent_sdk/advanced_memory/_paths.py rename to trpc_agent_sdk/sessions/compact/_paths.py index f68646127..87a473c17 100644 --- a/trpc_agent_sdk/advanced_memory/_paths.py +++ b/trpc_agent_sdk/sessions/compact/_paths.py @@ -12,7 +12,7 @@ from dataclasses import dataclass from pathlib import Path -from ._config import AdvancedMemoryConfig +from ._config import AdvancedCompactConfig _SAFE_COMPONENT_PATTERN = re.compile(r"[^A-Za-z0-9._-]+") @@ -60,7 +60,7 @@ def storage_key(self) -> str: class AdvancedMemoryPaths: """Build all disk paths for long-term and session memory.""" - config: AdvancedMemoryConfig + config: AdvancedCompactConfig scope: MemoryScope | None = None def for_scope(self, app_name: str, user_id: str) -> "AdvancedMemoryPaths": @@ -162,6 +162,10 @@ def storage_reference( return str(local_path) if self.scope is None: raise ValueError("A scoped path is required for non-local memory storage") + if resource == "session_memory": + return ("session-state://" + f"{self.scope.app_name}/{self.scope.user_id}/{session_id}/" + "_trpc_agent:summary") app_component = self.tenant_root_dir.parent.name user_component = self.tenant_root_dir.name @@ -176,8 +180,6 @@ def storage_reference( session_base = f"{self.config.redis_key_prefix}:{{{app_component}:{user_component}:{safe_session_id}}}" if resource == "transcript": key = f"{session_base}:transcript" - elif resource == "session_memory": - key = f"{session_base}:summary" else: key = f"{session_base}:tool:{result_id}" return f"advanced-memory://redis/{key}" @@ -190,8 +192,6 @@ def storage_reference( suffix = f"memory/topic/{local_path.name}" elif resource == "transcript": suffix = f"{session_id}/transcript" - elif resource == "session_memory": - suffix = f"{session_id}/summary" else: suffix = f"{session_id}/tool/{self.tool_result_path(session_id or '', result_id or '').stem}" return f"advanced-memory://sql/{app_name}/{user_id}/{suffix}" diff --git a/trpc_agent_sdk/sessions/compact/_redis_stores.py b/trpc_agent_sdk/sessions/compact/_redis_stores.py new file mode 100644 index 000000000..b60363d8c --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/_redis_stores.py @@ -0,0 +1,297 @@ +"""Redis implementations of the Advanced Memory storage contracts.""" + +from __future__ import annotations + +import asyncio +import json +from collections.abc import Mapping +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 AdvancedCompactConfig +from ._formats import MemoryDocument, MemoryIndexEntry +from ._paths import AdvancedMemoryPaths + +_APPEND_UNIQUE_SCRIPT = """ +if redis.call('SADD', KEYS[2], ARGV[1]) == 0 then return 0 end +redis.call('XADD', KEYS[1], '*', 'data', ARGV[2]) +return 1 +""" + +_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: AdvancedCompactConfig, 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}}}" + self._app_base = f"{config.redis_key_prefix}:{{{app_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 _session_base(self, session_id: str) -> str: + safe_session_id = self._paths.session_dir(session_id).name + tenant = f"{self._paths.tenant_root_dir.parent.name}:{self._paths.tenant_root_dir.name}" + return f"{self._config.redis_key_prefix}:{{{tenant}:{safe_session_id}}}" + + def _session_registry(self, session_id: str) -> str: + return f"{self._session_base(session_id)}:keys" + + 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_session_ttl(self, session_id: str, *keys: str) -> None: + skip_prefixes: tuple[str, ...] = () + if not self._config.session_ttl_delete_transcripts: + skip_prefixes = (f"{self._session_base(session_id)}:transcript", ) + await self._refresh_ttl_group( + self._session_registry(session_id), + list(keys), + self._config.session_ttl_seconds, + skip_prefixes=skip_prefixes, + ) + + async def _refresh_memory_ttl(self, *keys: str) -> None: + await self._refresh_ttl_group( + self._memory_registry(), + list(keys), + self._config.memory_ttl_seconds, + ) + + async def delete_session(self, session_id: str) -> None: + """Delete all Advanced Memory keys for one session.""" + session_base = self._session_base(session_id) + registry = self._session_registry(session_id) + keys: set[str] = {registry} + tracked = await self._command("smembers", registry) or [] + keys.update(value for value in (self._text(item) for item in tracked) if value) + + cursor: Any = 0 + pattern = f"{session_base}:*" + while True: + cursor, scanned = await self._command( + "scan", + cursor, + match=pattern, + count=100, + ) + keys.update(value for value in (self._text(item) for item in scanned) if value) + if int(cursor) == 0: + break + if keys: + await self._command("delete", *keys) + + @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() + lines, used_bytes = [], 0 + for line in value.splitlines(keepends=True)[:self._config.memory_index_max_lines]: + size = len(line.encode(self._config.encoding)) + if used_bytes + size > self._config.memory_index_max_bytes: + break + lines.append(line) + used_bytes += size + return "".join(lines) + + 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] + + +class RedisToolResultStore(_RedisStore): + + async def write(self, session_id: str, result_id: str, serialized_result: str) -> Path: + key = f"{self._session_base(session_id)}:tool:{result_id}" + await self._command("set", key, serialized_result) + await self._refresh_session_ttl(session_id, key) + return Path(f"advanced-memory://{key}") + + async def read(self, session_id: str, result_id: str) -> str | None: + key = f"{self._session_base(session_id)}:tool:{result_id}" + value = await self._command("get", key) + await self._refresh_session_ttl(session_id, key) + return self._text(value) + + +class RedisTranscriptStore(_RedisStore): + + @staticmethod + def _validate_record(record: Mapping[str, Any]) -> None: + """Reject Event and Session Memory duplication in Redis.""" + if record.get("kind") in {"event", "session-memory-checkpoint"}: + raise ValueError("Redis transcripts only store context-compression records") + + async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: + self._validate_record(record) + payload = dict(record) + payload.setdefault("recorded_at", datetime.now(timezone.utc).isoformat()) + stream = f"{self._session_base(session_id)}:transcript" + await self._command("xadd", stream, {"data": json.dumps(payload)}) + await self._refresh_session_ttl(session_id, stream) + return Path(f"advanced-memory://{stream}") + + async def append_unique(self, session_id: str, record: Mapping[str, Any], *, unique_key: str) -> tuple[Path, bool]: + self._validate_record(record) + payload = dict(record) + value = payload.get(unique_key) + if not isinstance(value, str) or not 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()) + stream = f"{self._session_base(session_id)}:transcript" + seen = f"{stream}:seen:{unique_key}" + async with self._storage.create_db_session() as connection: + added = await self._storage.execute_command( + connection, + RedisCommand( + method="eval", + args=(_APPEND_UNIQUE_SCRIPT, 2, stream, seen, value, json.dumps(payload, ensure_ascii=False)), + )) + await self._refresh_session_ttl(session_id, stream, seen) + return Path(f"advanced-memory://{stream}"), bool(added) + + async def read_all(self, session_id: str) -> list[dict[str, Any]]: + stream = f"{self._session_base(session_id)}:transcript" + entries = await self._command("xrange", stream, "-", "+") + await self._refresh_session_ttl(session_id, stream) + records: list[dict[str, Any]] = [] + for _, fields in entries: + value = fields.get(b"data") if isinstance(fields, dict) else None + value = value or fields.get("data") + text = self._text(value) + if text: + records.append(json.loads(text)) + return records diff --git a/trpc_agent_sdk/advanced_memory/_runtime.py b/trpc_agent_sdk/sessions/compact/_runtime.py similarity index 87% rename from trpc_agent_sdk/advanced_memory/_runtime.py rename to trpc_agent_sdk/sessions/compact/_runtime.py index 046357fcd..e0bd6d24f 100644 --- a/trpc_agent_sdk/advanced_memory/_runtime.py +++ b/trpc_agent_sdk/sessions/compact/_runtime.py @@ -14,7 +14,8 @@ import threading from typing import Any -from ._config import AdvancedMemoryConfig +from ._config import AdvancedCompactConfig +from ._coordination import CrossLoopLock from ._coordination import SessionOperationCoordinator from ._paths import AdvancedMemoryPaths from ._paths import MemoryScope @@ -29,11 +30,11 @@ class AdvancedMemoryRuntime: """Aggregate configuration, paths, and the three storage objects.""" - config: AdvancedMemoryConfig + config: AdvancedCompactConfig paths: AdvancedMemoryPaths coordination: SessionOperationCoordinator long_term_memory: LongTermMemoryStore - session_memory: SessionMemoryStore + session_memory: SessionMemoryStore | None tool_results: ToolResultStore transcripts: TranscriptStore _scoped_runtimes: dict[MemoryScope, "ScopedAdvancedMemoryRuntime"] = field( @@ -50,11 +51,17 @@ class AdvancedMemoryRuntime: _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: AdvancedMemoryConfig | None = None) -> "AdvancedMemoryRuntime": + def create(cls, config: AdvancedCompactConfig | None = None) -> "AdvancedMemoryRuntime": """Create a runtime isolated from the legacy mechanism.""" - resolved_config = config or AdvancedMemoryConfig() + resolved_config = config or AdvancedCompactConfig() paths = AdvancedMemoryPaths(resolved_config) redis_storage = None sql_storage = None @@ -81,7 +88,8 @@ def create(cls, config: AdvancedMemoryConfig | None = None) -> "AdvancedMemoryRu paths=paths, coordination=SessionOperationCoordinator(), long_term_memory=LongTermMemoryStore(resolved_config, paths), - session_memory=SessionMemoryStore(resolved_config, paths), + session_memory=(SessionMemoryStore(resolved_config, paths) + if resolved_config.storage_backend == "local" else None), tool_results=ToolResultStore(resolved_config, paths), transcripts=TranscriptStore(resolved_config, paths), _redis_storage=redis_storage, @@ -100,7 +108,6 @@ def for_scope(self, app_name: str, user_id: str) -> "ScopedAdvancedMemoryRuntime if self.config.storage_backend == "redis": from trpc_agent_sdk.storage import RedisStorage from ._redis_stores import RedisLongTermMemoryStore - from ._redis_stores import RedisSessionMemoryStore from ._redis_stores import RedisToolResultStore from ._redis_stores import RedisTranscriptStore @@ -109,19 +116,18 @@ def for_scope(self, app_name: str, user_id: str) -> "ScopedAdvancedMemoryRuntime is_async=self.config.redis_is_async, ) long_term_memory = RedisLongTermMemoryStore(self.config, paths, storage) - session_memory = RedisSessionMemoryStore(self.config, paths, storage) + session_memory = None tool_results = RedisToolResultStore(self.config, paths, storage) transcripts = RedisTranscriptStore(self.config, paths, storage) elif self.config.storage_backend == "sql": from ._sql_stores import SqlLongTermMemoryStore - from ._sql_stores import SqlSessionMemoryStore from ._sql_stores import SqlToolResultStore from ._sql_stores import SqlTranscriptStore 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) - session_memory = SqlSessionMemoryStore(self.config, paths, storage) + session_memory = None tool_results = SqlToolResultStore(self.config, paths, storage) transcripts = SqlTranscriptStore(self.config, paths, storage) else: @@ -189,14 +195,18 @@ async def initialize(self) -> bool: async def close(self) -> None: """Release shared external backend resources.""" - 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_storage is not None: - if self._sql_cleanup is not None: - await self._sql_cleanup.close() - await self._sql_storage.close() + 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_storage is not None: + if self._sql_cleanup is not None: + await self._sql_cleanup.close() + await self._sql_storage.close() + object.__setattr__(self, "_closed", True) @dataclass(frozen=True) @@ -207,12 +217,12 @@ class ScopedAdvancedMemoryRuntime: scope: MemoryScope paths: AdvancedMemoryPaths long_term_memory: LongTermMemoryStore - session_memory: SessionMemoryStore + session_memory: SessionMemoryStore | None tool_results: ToolResultStore transcripts: TranscriptStore @property - def config(self) -> AdvancedMemoryConfig: + def config(self) -> AdvancedCompactConfig: """Return the root runtime configuration.""" return self.root.config @@ -242,7 +252,7 @@ async def delete_session(self, session_id: str) -> None: session_dir = self.paths.session_dir(session_id) await asyncio.to_thread(shutil.rmtree, session_dir, True) return - delete_session = getattr(self.session_memory, "delete_session", None) + delete_session = getattr(self.tool_results, "delete_session", None) if delete_session is None: raise RuntimeError("Configured Advanced Memory backend cannot delete sessions") await delete_session(session_id) 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 05623d48d..35ddc81f2 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,12 +28,16 @@ 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 ._formats import build_session_memory_state +from ._formats import parse_session_memory_state from ._runtime import AdvancedMemoryRuntime 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 @@ -340,6 +346,7 @@ def __init__( 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,12 +356,56 @@ def __init__( model, section_max_chars=memory_runtime.config.session_memory_section_max_chars, ) + self._session_service = session_service @property def runtime(self) -> AdvancedMemoryRuntime: """Return the runtime bound to this extractor.""" return self._runtime + @property + def uses_session_state(self) -> bool: + """Return whether this backend stores Session Memory in Session.state.""" + return self._runtime.config.storage_backend in {"redis", "sql"} + + 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]], @@ -609,9 +660,59 @@ def missing_context(end: int) -> list[str]: async def _read_current_memory(self, session: "SessionABC") -> str: """Read old session memory or return the complete empty template.""" - current = await self._runtime.for_session(session).session_memory.read(session.id) + if self.uses_session_state: + parsed = parse_session_memory_state(session.state.get(SESSION_MEMORY_STATE_KEY)) + if parsed is not None: + return parsed[0].to_markdown() + return SessionMemoryDocument().to_markdown() + store = self._runtime.for_session(session).session_memory + if store is None: + raise RuntimeError("Session Memory store is unavailable") + current = await store.read(session.id) return current if current 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: "SessionABC", @@ -634,6 +735,34 @@ async def _persist_checkpoint( document.key_results, document.worklog, ) + if self.uses_session_state: + if self._session_service is None: + raise RuntimeError("Redis/SQL 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}, + ) + return runtime = self._runtime.for_session(session) await runtime.transcripts.append_unique( session.id, @@ -668,8 +797,14 @@ async def extract_if_needed( async with self._runtime.coordination.guard(session_key) as acquired: if not acquired: return SessionMemoryExtractionResult(False, "coordination-timeout") - records = await runtime.transcripts.read_all(session.id) - checkpoint = self._last_checkpoint(records) + if self.uses_session_state: + records = self._session_event_records(session) + checkpoint, checkpoint_context_tokens = self._state_checkpoint(session) + else: + records = await runtime.transcripts.read_all(session.id) + checkpoint = self._last_checkpoint(records) + checkpoint_context_tokens = (checkpoint.get("context_tokens") if checkpoint is not None + and isinstance(checkpoint.get("context_tokens"), int) else None) 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( @@ -684,8 +819,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 @@ -718,7 +851,10 @@ async def extract_if_needed( max_chars=config.session_memory_section_max_chars, total_max_chars=config.session_memory_total_max_chars, ) - await runtime.session_memory.write(session.id, document) + if not self.uses_session_state: + if runtime.session_memory is None: + raise RuntimeError("Session Memory store is unavailable") + await runtime.session_memory.write(session.id, document) await self._persist_checkpoint( session, included, diff --git a/trpc_agent_sdk/advanced_memory/_session_service.py b/trpc_agent_sdk/sessions/compact/_session_service.py similarity index 91% rename from trpc_agent_sdk/advanced_memory/_session_service.py rename to trpc_agent_sdk/sessions/compact/_session_service.py index 3c6574cee..e1f1f2441 100644 --- a/trpc_agent_sdk/advanced_memory/_session_service.py +++ b/trpc_agent_sdk/sessions/compact/_session_service.py @@ -36,6 +36,8 @@ def __init__( session_memory_extractor: SessionMemoryExtractor | None = None, ) -> None: """Store the legacy service and optional Advanced Memory runtime.""" + if isinstance(delegate, TranscriptSessionService): + raise ValueError("Transcript session service is already wrapped") self._delegate = delegate self._memory_runtime = memory_runtime self._session_memory_extractor = session_memory_extractor @@ -55,6 +57,16 @@ def memory_runtime(self) -> AdvancedMemoryRuntime: """Return the Advanced Memory runtime used by the decorator.""" return self._memory_runtime + @property + def session_config(self) -> Any: + """Expose the original service configuration.""" + return getattr(self._delegate, "session_config", None) + + @property + def summarizer_manager(self) -> Any: + """Expose the original service summarizer, when configured.""" + return getattr(self._delegate, "summarizer_manager", None) + @property def session_memory_extractor(self) -> SessionMemoryExtractor | None: """Return the session memory extractor used after each turn.""" @@ -197,6 +209,14 @@ async def update_session(self, session: SessionABC) -> None: """Delegate session updates to the underlying service.""" await self._delegate.update_session(session) + async def patch_session_state( + self, + session: SessionABC, + state_delta: dict[str, Any], + ) -> None: + """Delegate state-only updates without touching persisted Events.""" + await self._delegate.patch_session_state(session, state_delta) + async def create_session_summary( self, session: SessionABC, diff --git a/trpc_agent_sdk/sessions/compact/_sql_stores.py b/trpc_agent_sdk/sessions/compact/_sql_stores.py new file mode 100644 index 000000000..4c77eae64 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/_sql_stores.py @@ -0,0 +1,528 @@ +"""SQL implementations of the Advanced Memory storage contracts.""" + +from __future__ import annotations + +import json +import asyncio +import hashlib +import uuid +from datetime import datetime, timedelta, timezone +from dataclasses import replace +from pathlib import Path +from collections.abc import Mapping +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 AdvancedCompactConfig +from ._formats import MemoryDocument, MemoryIndexEntry +from ._paths import AdvancedMemoryPaths + + +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 SqlTranscript(AdvancedMemorySqlBase): + __tablename__ = "advanced_memory_transcripts" + + 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) + session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + record_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + payload: Mapped[str] = mapped_column(Text) + recorded_at: Mapped[datetime] = mapped_column(PreciseTimestamp, default=func.now()) + expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) + + +class SqlTranscriptSeen(AdvancedMemorySqlBase): + __tablename__ = "advanced_memory_transcript_seen" + + dedupe_id: Mapped[str] = mapped_column(String(64), primary_key=True) + app_name: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) + user_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) + session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) + unique_key: Mapped[str] = mapped_column(String(DEFAULT_MAX_VARCHAR_LENGTH)) + unique_value: Mapped[str] = mapped_column(String(DEFAULT_MAX_VARCHAR_LENGTH)) + expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) + + +class SqlToolResult(AdvancedMemorySqlBase): + __tablename__ = "advanced_memory_tool_results" + + 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) + session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + result_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_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: AdvancedCompactConfig, 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 + + async def _refresh_session_scope(self, db: Any, session_id: str) -> None: + expiry = self._expiry(self._config.session_ttl_seconds) + if expiry is None: + return + tables = ((SqlToolResult, (self._app_name, self._user_id, session_id)), ) + if self._config.session_ttl_delete_transcripts: + tables = ( + (SqlTranscript, (self._app_name, self._user_id, session_id)), + (SqlTranscriptSeen, (self._app_name, self._user_id, session_id)), + *tables, + ) + for model, key in tables: + rows = await self._storage.query( + db, + SqlKey(key=key, storage_cls=model), + SqlCondition(filters=[ + getattr(model, "app_name") == self._app_name, + getattr(model, "user_id") == self._user_id, + getattr(model, "session_id") == session_id, + getattr(model, "expires_at").is_(None) | (getattr(model, "expires_at") > self._now()), + ]), + ) + for row in rows: + row.expires_at = expiry + + async def delete_session(self, session_id: str) -> None: + """Delete all Advanced Memory rows for one session.""" + models = ( + SqlTranscript, + SqlTranscriptSeen, + SqlToolResult, + ) + filters = { + SqlTranscript: [ + SqlTranscript.app_name == self._app_name, + SqlTranscript.user_id == self._user_id, + SqlTranscript.session_id == session_id, + ], + SqlTranscriptSeen: [ + SqlTranscriptSeen.app_name == self._app_name, + SqlTranscriptSeen.user_id == self._user_id, + SqlTranscriptSeen.session_id == session_id, + ], + SqlToolResult: [ + SqlToolResult.app_name == self._app_name, + SqlToolResult.user_id == self._user_id, + SqlToolResult.session_id == session_id, + ], + } + async with self._storage.create_db_session() as db: + for model in models: + await self._storage.delete( + db, + SqlKey(key=tuple(), storage_cls=model), + SqlCondition(filters=filters[model]), + ) + await self._storage.commit(db) + + +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) + await self._storage.commit(db) + content = row.content + lines, used_bytes = [], 0 + for line in content.splitlines(keepends=True)[:self._config.memory_index_max_lines]: + size = len(line.encode(self._config.encoding)) + if used_bytes + size > self._config.memory_index_max_bytes: + break + lines.append(line) + used_bytes += size + return "".join(lines) + + 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 SqlToolResultStore(_SqlStore): + + async def write(self, session_id: str, result_id: str, serialized_result: str) -> Path: + async with self._storage.create_db_session() as db: + key = (self._app_name, self._user_id, session_id, result_id) + row = await self._storage.get(db, SqlKey(key=key, storage_cls=SqlToolResult)) + if row is None: + row = SqlToolResult( + app_name=key[0], + user_id=key[1], + session_id=key[2], + result_id=key[3], + ) + await self._storage.add(db, row) + row.content = serialized_result + row.updated_at = self._now() + row.expires_at = self._expiry(self._config.session_ttl_seconds) + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/tool/{result_id}") + + async def read(self, session_id: str, result_id: str) -> str | None: + async with self._storage.create_db_session() as db: + row = await self._storage.get( + db, + SqlKey(key=(self._app_name, self._user_id, session_id, result_id), storage_cls=SqlToolResult), + ) + if row is None or self._expired(row.expires_at): + return None + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return row.content + + +class SqlTranscriptStore(_SqlStore): + + @staticmethod + def _validate_record(record: Mapping[str, Any]) -> None: + """Reject Event and Session Memory duplication in SQL.""" + if record.get("kind") in {"event", "session-memory-checkpoint"}: + raise ValueError("SQL transcripts only store context-compression records") + + def _dedupe_id(self, session_id: str, unique_key: str, value: str) -> str: + raw = "\0".join((self._app_name, self._user_id, session_id, unique_key, value)) + return hashlib.sha256(raw.encode("utf-8")).hexdigest() + + async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: + self._validate_record(record) + payload = dict(record) + payload.setdefault("recorded_at", self._now().replace(tzinfo=timezone.utc).isoformat()) + async with self._storage.create_db_session() as db: + await self._storage.add( + db, + SqlTranscript( + app_name=self._app_name, + user_id=self._user_id, + session_id=session_id, + record_id=uuid.uuid4().hex, + payload=json.dumps(payload, ensure_ascii=False), + expires_at=(self._expiry(self._config.session_ttl_seconds) + if self._config.session_ttl_delete_transcripts else None), + )) + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript") + + async def append_unique( + self, + session_id: str, + record: Mapping[str, Any], + *, + unique_key: str, + ) -> tuple[Path, bool]: + self._validate_record(record) + payload = dict(record) + value = payload.get(unique_key) + if not isinstance(value, str) or not value: + raise ValueError(f"Transcript unique key {unique_key!r} must be a non-empty string") + async with self._storage.create_db_session() as db: + dedupe_id = self._dedupe_id(session_id, unique_key, value) + seen_key = (self._app_name, self._user_id, session_id, unique_key, value) + seen = await self._storage.get( + db, + SqlKey(key=(dedupe_id, ), storage_cls=SqlTranscriptSeen), + ) + if seen is not None and not self._expired(seen.expires_at): + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript"), False + if seen is not None: + await self._storage.delete( + db, + SqlKey(key=(dedupe_id, ), storage_cls=SqlTranscriptSeen), + SqlCondition(filters=[ + SqlTranscriptSeen.dedupe_id == dedupe_id, + ]), + ) + payload.setdefault("recorded_at", self._now().replace(tzinfo=timezone.utc).isoformat()) + await self._storage.add( + db, + SqlTranscriptSeen( + dedupe_id=dedupe_id, + app_name=seen_key[0], + user_id=seen_key[1], + session_id=seen_key[2], + unique_key=seen_key[3], + unique_value=seen_key[4], + expires_at=(self._expiry(self._config.session_ttl_seconds) + if self._config.session_ttl_delete_transcripts else None), + )) + await self._storage.add( + db, + SqlTranscript( + app_name=self._app_name, + user_id=self._user_id, + session_id=session_id, + record_id=uuid.uuid4().hex, + payload=json.dumps(payload, ensure_ascii=False), + expires_at=(self._expiry(self._config.session_ttl_seconds) + if self._config.session_ttl_delete_transcripts else None), + )) + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript"), True + + async def read_all(self, session_id: str) -> list[dict[str, Any]]: + async with self._storage.create_db_session() as db: + rows = await self._storage.query( + db, + SqlKey(key=(self._app_name, self._user_id, session_id), storage_cls=SqlTranscript), + SqlCondition( + filters=[ + SqlTranscript.app_name == self._app_name, + SqlTranscript.user_id == self._user_id, + SqlTranscript.session_id == session_id, + SqlTranscript.expires_at.is_(None) | (SqlTranscript.expires_at > self._now()), + ], + order_func=SqlTranscript.recorded_at.asc, + ), + ) + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return [json.loads(row.payload) for row in rows] + + +class SqlAdvancedMemoryCleanup: + """Periodically remove expired Advanced Memory SQL rows.""" + + _models = ( + SqlMemoryIndex, + SqlMemoryTopic, + SqlTranscript, + SqlTranscriptSeen, + SqlToolResult, + ) + + def __init__(self, config: AdvancedCompactConfig, 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 + and self._config.session_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: + models = self._models if self._config.session_ttl_delete_transcripts else tuple( + model for model in self._models if model is not SqlTranscript) + for model in models: + await self._storage.delete( + db, + SqlKey(key=tuple(), storage_cls=model), + SqlCondition(filters=[model.expires_at.is_not(None), model.expires_at <= now]), + ) + 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", + "SqlToolResultStore", + "SqlTranscriptStore", +] diff --git a/trpc_agent_sdk/sessions/compact/_storage.py b/trpc_agent_sdk/sessions/compact/_storage.py new file mode 100644 index 000000000..98b174872 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/_storage.py @@ -0,0 +1,499 @@ +# 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 shutil +import tempfile +import threading +import time +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 AdvancedCompactConfig +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 + + +def _is_expired(path: Path, ttl: int | None) -> bool: + if ttl is None or not path.exists(): + return False + return time.time() - path.stat().st_mtime >= ttl + + +def _touch(path: Path) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.touch() + + +def _expire_memory_dir(memory_dir: Path, config: AdvancedCompactConfig) -> bool: + """Expire the whole long-term memory group using index activity time.""" + index_path = memory_dir / config.memory_index_name + if not _is_expired(index_path, config.memory_ttl_seconds): + return False + for path in memory_dir.glob("*.md"): + path.unlink(missing_ok=True) + return True + + +def _refresh_memory_dir(memory_dir: Path) -> None: + """Refresh activity for every file in the long-term memory group.""" + for path in memory_dir.glob("*.md"): + _touch(path) + + +def _session_activity_path(session_dir: Path) -> Path: + return session_dir / ".advanced-memory-activity" + + +def _expire_session_dir(session_dir: Path, config: AdvancedCompactConfig) -> bool: + """Expire all Advanced Memory data belonging to one local session.""" + if not session_dir.exists() or config.session_ttl_seconds is None: + return False + activity_path = _session_activity_path(session_dir) + if activity_path.exists(): + expired = _is_expired(activity_path, config.session_ttl_seconds) + else: + files = [path for path in session_dir.rglob("*") if path.is_file()] + expired = bool(files) and time.time() - max(path.stat().st_mtime + for path in files) >= config.session_ttl_seconds + if expired: + if config.session_ttl_delete_transcripts: + shutil.rmtree(session_dir, ignore_errors=True) + else: + transcript_path = session_dir / config.transcript_name + for child in session_dir.iterdir(): + if child == transcript_path: + continue + if child.is_dir(): + shutil.rmtree(child, ignore_errors=True) + else: + child.unlink(missing_ok=True) + return expired + + +def _refresh_session_dir(session_dir: Path) -> None: + _touch(_session_activity_path(session_dir)) + + +class LongTermMemoryStore: + """Manage MEMORY.md and its detail files in the same directory.""" + + def __init__(self, config: AdvancedCompactConfig, 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 _expire_memory_dir(self._paths.memory_dir, self._config) or not self.index_path.exists(): + return "" + _refresh_memory_dir(self._paths.memory_dir) + 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) + _refresh_memory_dir(self._paths.memory_dir) + + 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_topic_sync, path) + + def _read_topic_sync(self, path: Path) -> str | None: + if _expire_memory_dir(self._paths.memory_dir, self._config) or not path.exists(): + return None + _refresh_memory_dir(self._paths.memory_dir) + return path.read_text(encoding=self._config.encoding) + + 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_sync, path) + + def _read_frontmatter_sync(self, path: Path) -> str | None: + """Synchronously read a topic's bounded frontmatter block.""" + if _expire_memory_dir(self._paths.memory_dir, self._config) or not path.exists(): + return None + _refresh_memory_dir(self._paths.memory_dir) + 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(self._write_topic_sync, path, document.to_markdown()) + return path + + def _write_topic_sync(self, path: Path, content: str) -> None: + _expire_memory_dir(self._paths.memory_dir, self._config) + _atomic_write_text(path, content, encoding=self._config.encoding) + _refresh_memory_dir(self._paths.memory_dir) + + 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 _expire_memory_dir(self._paths.memory_dir, self._config): + return [] + if not self._paths.memory_dir.exists(): + return [] + _refresh_memory_dir(self._paths.memory_dir) + 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: AdvancedCompactConfig, 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, session_id, path) + + def _read_sync(self, session_id: str, path: Path) -> str | None: + """Synchronously read session memory.""" + session_dir = self._paths.session_dir(session_id) + if _expire_session_dir(session_dir, self._config) or not path.exists(): + return None + _refresh_session_dir(session_dir) + 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( + self._write_sync, + session_id, + path, + document.to_markdown(), + ) + return path + + def _write_sync(self, session_id: str, path: Path, content: str) -> None: + session_dir = self._paths.session_dir(session_id) + _expire_session_dir(session_dir, self._config) + _atomic_write_text(path, content, encoding=self._config.encoding) + _refresh_session_dir(session_dir) + + +class ToolResultStore: + """Persist complete tool results that exceed the context budget.""" + + def __init__(self, config: AdvancedCompactConfig, 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( + self._write_sync, + session_id, + path, + serialized_result, + ) + 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, session_id, path) + + def _read_sync(self, session_id: str, path: Path) -> str | None: + """Synchronously read an optional complete tool-result file.""" + session_dir = self._paths.session_dir(session_id) + if _expire_session_dir(session_dir, self._config) or not path.exists(): + return None + _refresh_session_dir(session_dir) + return path.read_text(encoding=self._config.encoding) + + def _write_sync(self, session_id: str, path: Path, content: str) -> None: + session_dir = self._paths.session_dir(session_id) + _expire_session_dir(session_dir, self._config) + _atomic_write_text(path, content, encoding=self._config.encoding) + _refresh_session_dir(session_dir) + + +class TranscriptStore: + """Store complete per-session records as append-only JSONL.""" + + def __init__(self, config: AdvancedCompactConfig, 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.""" + _expire_session_dir(path.parent, self._config) + path.parent.mkdir(parents=True, exist_ok=True) + with self._write_lock: + self._append_serialized_unlocked(path, serialized) + _refresh_session_dir(path.parent) + + 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.""" + with self._write_lock: + if _expire_session_dir(path.parent, self._config): + for cache_key in list(self._seen_unique_values): + if cache_key[0] == path: + self._seen_unique_values.pop(cache_key, None) + path.parent.mkdir(parents=True, exist_ok=True) + cache_key = (path, unique_key) + 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) + _refresh_session_dir(path.parent) + 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: + expired = _expire_session_dir(path.parent, self._config) + if expired and self._config.session_ttl_delete_transcripts: + return [] + if not path.exists(): + return [] + _refresh_session_dir(path.parent) + 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 + + +class LocalAdvancedMemoryCleanup: + """Periodically remove expired local Advanced Memory data.""" + + def __init__(self, config: AdvancedCompactConfig) -> 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: + return + if self._config.memory_ttl_seconds is None and self._config.session_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] + session_roots = [root / self._config.session_dir_name] + tenants_root = root / "tenants" + if tenants_root.exists(): + for app_dir in tenants_root.iterdir(): + if app_dir.is_dir(): + for user_dir in app_dir.iterdir(): + if user_dir.is_dir(): + memory_dirs.append(user_dir / self._config.memory_dir_name) + session_roots.append(user_dir / self._config.session_dir_name) + for memory_dir in memory_dirs: + _expire_memory_dir(memory_dir, self._config) + for session_root in session_roots: + if session_root.exists(): + for session_dir in session_root.iterdir(): + if session_dir.is_dir(): + _expire_session_dir(session_dir, self._config) + + async def _run(self) -> None: + if self._stop_event is None: + return + ttls = [ + ttl for ttl in ( + self._config.memory_ttl_seconds, + self._config.session_ttl_seconds, + ) if ttl is not None + ] + interval = min(ttls) if ttls else 60 + try: + while not self._stop_event.is_set(): + try: + await asyncio.wait_for(self._stop_event.wait(), timeout=interval) + break + except asyncio.TimeoutError: + await self.cleanup_once() + except asyncio.CancelledError: + raise + + async def close(self) -> None: + if self._task is not None: + await self.cleanup_once() + 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/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 100% rename from trpc_agent_sdk/advanced_memory/_tool_result_budget.py rename to trpc_agent_sdk/sessions/compact/_tool_result_budget.py diff --git a/trpc_agent_sdk/advanced_memory/_transcript.py b/trpc_agent_sdk/sessions/compact/_transcript.py similarity index 100% rename from trpc_agent_sdk/advanced_memory/_transcript.py rename to trpc_agent_sdk/sessions/compact/_transcript.py diff --git a/trpc_agent_sdk/tools/_advanced_memory_tool.py b/trpc_agent_sdk/tools/_advanced_memory_tool.py index 8d6fdbdbe..1601b41eb 100644 --- a/trpc_agent_sdk/tools/_advanced_memory_tool.py +++ b/trpc_agent_sdk/tools/_advanced_memory_tool.py @@ -11,12 +11,12 @@ 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.sessions.compact._formats import MemoryDocument +from trpc_agent_sdk.sessions.compact._formats import MemoryIndexEntry +from trpc_agent_sdk.sessions.compact._formats import MemoryType +from trpc_agent_sdk.sessions.compact._formats import memory_freshness +from trpc_agent_sdk.sessions.compact._formats import parse_memory_updated_at +from trpc_agent_sdk.sessions.compact._runtime import AdvancedMemoryRuntime from ._function_tool import FunctionTool From eb29c64c17a24eba5b085eba415bc7b9628aca56 Mon Sep 17 00:00:00 2001 From: congkechen Date: Thu, 10 Sep 2026 15:29:04 +0800 Subject: [PATCH 4/6] =?UTF-8?q?feature:=20=E4=BC=98=E5=8C=96=20Session=20C?= =?UTF-8?q?ompact=20=E5=AE=9E=E7=8E=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../README.md | 19 +- .../run_agent.py | 17 +- .../README.md | 2 +- .../run_agent.py | 4 +- .../README.md | 2 +- .../run_agent.py | 4 +- .../README.md | 20 +- .../run_agent.py | 14 +- .../.env | 6 +- .../README.md | 20 +- .../run_agent.py | 19 +- .../test_advanced_memory_tools.py | 6 +- tests/advanced_memory/test_memory_context.py | 86 +-- tests/advanced_memory/test_preload_memory.py | 8 +- tests/advanced_memory/test_redis_stores.py | 142 ----- tests/advanced_memory/test_sql_stores.py | 113 ---- tests/advanced_memory/test_storage.py | 466 --------------- tests/sessions/compact/test_autocompact.py | 552 ------------------ .../test_context_compression_integration.py | 452 -------------- tests/sessions/compact/test_history_snip.py | 225 ------- tests/sessions/compact/test_microcompact.py | 178 ------ .../sessions/compact/test_session_compact.py | 128 ++++ .../compact/test_session_memory_extractor.py | 522 ----------------- .../compact/test_session_memory_state.py | 160 ----- tests/sessions/compact/test_token_budget.py | 9 +- .../compact/test_tool_result_budget.py | 332 ----------- .../test_transcript_session_service.py | 138 ----- trpc_agent_sdk/advanced_memory/__init__.py | 14 +- trpc_agent_sdk/advanced_memory/_config.py | 83 +++ .../advanced_memory/_integration.py | 64 +- .../advanced_memory/_memory_context.py | 2 +- trpc_agent_sdk/advanced_memory/_paths.py | 110 ++++ .../advanced_memory/_preload_memory.py | 2 +- .../advanced_memory/_redis_stores.py | 303 +++++++++- trpc_agent_sdk/advanced_memory/_runtime.py | 206 +++++++ trpc_agent_sdk/advanced_memory/_sql_stores.py | 534 ++++++++++++++++- trpc_agent_sdk/advanced_memory/_storage.py | 190 +++++- .../advanced_memory/_storage_backend.py | 4 +- trpc_agent_sdk/memory/__init__.py | 8 +- .../memory/_advanced_memory_service.py | 12 +- trpc_agent_sdk/runners.py | 7 +- trpc_agent_sdk/sessions/__init__.py | 14 +- .../sessions/_base_session_service.py | 37 +- .../sessions/_in_memory_session_service.py | 8 - .../sessions/_redis_session_service.py | 8 - .../sessions/_sql_session_service.py | 8 - trpc_agent_sdk/sessions/compact/__init__.py | 28 +- .../sessions/compact/_autocompact.py | 228 ++------ .../sessions/compact/_base_config.py | 29 - .../sessions/compact/_base_manager.py | 15 +- trpc_agent_sdk/sessions/compact/_callbacks.py | 4 +- trpc_agent_sdk/sessions/compact/_config.py | 235 +------- trpc_agent_sdk/sessions/compact/_formats.py | 2 + .../sessions/compact/_history_snip.py | 51 +- .../sessions/compact/_integration.py | 155 ----- trpc_agent_sdk/sessions/compact/_manager.py | 93 ++- .../sessions/compact/_microcompact.py | 51 +- trpc_agent_sdk/sessions/compact/_paths.py | 208 ------- .../sessions/compact/_redis_stores.py | 297 ---------- trpc_agent_sdk/sessions/compact/_runtime.py | 243 +------- .../sessions/compact/_session_memory.py | 131 ++--- .../sessions/compact/_session_service.py | 236 -------- .../sessions/compact/_sql_stores.py | 528 ----------------- trpc_agent_sdk/sessions/compact/_storage.py | 499 ---------------- .../sessions/compact/_tool_result_budget.py | 170 ++---- .../sessions/compact/_transcript.py | 49 -- trpc_agent_sdk/tools/_advanced_memory_tool.py | 2 +- 67 files changed, 1930 insertions(+), 6582 deletions(-) delete mode 100644 tests/advanced_memory/test_redis_stores.py delete mode 100644 tests/advanced_memory/test_sql_stores.py delete mode 100644 tests/advanced_memory/test_storage.py delete mode 100644 tests/sessions/compact/test_autocompact.py delete mode 100644 tests/sessions/compact/test_context_compression_integration.py delete mode 100644 tests/sessions/compact/test_history_snip.py delete mode 100644 tests/sessions/compact/test_microcompact.py create mode 100644 tests/sessions/compact/test_session_compact.py delete mode 100644 tests/sessions/compact/test_session_memory_extractor.py delete mode 100644 tests/sessions/compact/test_session_memory_state.py delete mode 100644 tests/sessions/compact/test_tool_result_budget.py delete mode 100644 tests/sessions/compact/test_transcript_session_service.py create mode 100644 trpc_agent_sdk/advanced_memory/_config.py create mode 100644 trpc_agent_sdk/advanced_memory/_paths.py create mode 100644 trpc_agent_sdk/advanced_memory/_runtime.py delete mode 100644 trpc_agent_sdk/sessions/compact/_base_config.py delete mode 100644 trpc_agent_sdk/sessions/compact/_integration.py delete mode 100644 trpc_agent_sdk/sessions/compact/_paths.py delete mode 100644 trpc_agent_sdk/sessions/compact/_redis_stores.py delete mode 100644 trpc_agent_sdk/sessions/compact/_session_service.py delete mode 100644 trpc_agent_sdk/sessions/compact/_sql_stores.py delete mode 100644 trpc_agent_sdk/sessions/compact/_storage.py delete mode 100644 trpc_agent_sdk/sessions/compact/_transcript.py diff --git a/examples/memory_service_with_advanced_memory/README.md b/examples/memory_service_with_advanced_memory/README.md index dd58e491a..1e23757f7 100644 --- a/examples/memory_service_with_advanced_memory/README.md +++ b/examples/memory_service_with_advanced_memory/README.md @@ -25,7 +25,7 @@ AdvancedMemoryService ## 核心组装 ```python -config = AdvancedCompactConfig( +config = AdvancedMemoryServiceConfig( root_dir=Path(__file__).resolve().parent, ) @@ -33,14 +33,12 @@ session_service = InMemorySessionService( session_config=SessionServiceConfig( store_historical_events=True, ), -) -compact_manager = setup_advanced_session_compact( - agent, - session_service, - config, + session_compact_manager=AdvancedSessionCompactManager( + config=AdvancedCompactConfig(), + ), ) -memory_service = AdvancedMemoryService(runtime=compact_manager.runtime) +memory_service = AdvancedMemoryService(config=config) runner = Runner( app_name="advanced_memory_demo", agent=agent, @@ -49,11 +47,8 @@ runner = Runner( ) ``` -Session Compact 与 Advanced Memory 可以共享一个 Runtime;Runtime 的 `close()` -支持幂等调用,因此两个 Service 的正常关闭流程不会造成重复释放错误。 - -也可以直接构造实现了 `BaseSessionCompactManager` 的自定义 Manager,并通过 -`session_compact_manager=` 注入标准 SessionService。 +Session Compact 与 Advanced Memory 使用独立配置和 Runtime。Compact 只使用 +SessionService 的 events、historical_events 和 state。 ## 运行 diff --git a/examples/memory_service_with_advanced_memory/run_agent.py b/examples/memory_service_with_advanced_memory/run_agent.py index 17e8298cb..efd8014e0 100644 --- a/examples/memory_service_with_advanced_memory/run_agent.py +++ b/examples/memory_service_with_advanced_memory/run_agent.py @@ -12,11 +12,12 @@ from pathlib import Path from dotenv import load_dotenv -from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig +from trpc_agent_sdk.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 setup_advanced_session_compact +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 @@ -30,13 +31,15 @@ def create_services(agent) -> tuple[InMemorySessionService, AdvancedMemoryServic memory_ttl = os.getenv("M_TTL") session_ttl = os.getenv("SESSION_TTL") session_ttl_seconds = int(session_ttl) if session_ttl else 0 - config = AdvancedCompactConfig( + config = AdvancedMemoryServiceConfig( root_dir=Path(__file__).resolve().parent, memory_ttl_seconds=int(memory_ttl) if memory_ttl else None, session_ttl_seconds=session_ttl_seconds or None, memory_focus_instruction=("特别关注并主动记住用户长期稳定的兴趣爱好、" "编程语言偏好、开发习惯和测试习惯。"), ) + compact_config = AdvancedCompactConfig() + compact_manager = AdvancedSessionCompactManager(config=compact_config) session_service = InMemorySessionService( session_config=SessionServiceConfig( ttl=SessionServiceConfig.create_ttl_config( @@ -46,13 +49,9 @@ def create_services(agent) -> tuple[InMemorySessionService, AdvancedMemoryServic ), store_historical_events=True, ), + session_compact_manager=compact_manager, ) - compact_manager = setup_advanced_session_compact( - agent, - session_service, - config, - ) - return session_service, AdvancedMemoryService(runtime=compact_manager.runtime) + return session_service, AdvancedMemoryService(config=config) async def run_turn(runner, *, user_id: str, session_id: str, prompt: str) -> None: diff --git a/examples/memory_service_with_advanced_memory_redis/README.md b/examples/memory_service_with_advanced_memory_redis/README.md index c97328e13..d9a8273d0 100644 --- a/examples/memory_service_with_advanced_memory_redis/README.md +++ b/examples/memory_service_with_advanced_memory_redis/README.md @@ -177,7 +177,7 @@ Redis 版本最核心的构建过程可以简化为三步: redis_url = "redis://:password@localhost:6379/0" memory_service = AdvancedMemoryService( - AdvancedCompactConfig( + AdvancedMemoryServiceConfig( storage_backend="redis", redis_url=redis_url, memory_ttl_seconds=120, # from M_TTL; omit to disable expiration diff --git a/examples/memory_service_with_advanced_memory_redis/run_agent.py b/examples/memory_service_with_advanced_memory_redis/run_agent.py index 93dce8789..e8b175ce0 100644 --- a/examples/memory_service_with_advanced_memory_redis/run_agent.py +++ b/examples/memory_service_with_advanced_memory_redis/run_agent.py @@ -14,7 +14,7 @@ from dotenv import load_dotenv from agent.agent import create_agent -from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig +from trpc_agent_sdk.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 @@ -63,7 +63,7 @@ def build_redis_url_from_environment() -> str: def create_advanced_memory_service(redis_url: str) -> AdvancedMemoryService: """Create the long-term Advanced Memory service backed by Redis.""" memory_ttl = os.getenv("M_TTL") - config = AdvancedCompactConfig( + config = AdvancedMemoryServiceConfig( storage_backend="redis", redis_url=redis_url, redis_key_prefix="advanced-memory-redis-demo:v1", diff --git a/examples/memory_service_with_advanced_memory_sql/README.md b/examples/memory_service_with_advanced_memory_sql/README.md index 87dbac3f4..bc7c7dcf0 100644 --- a/examples/memory_service_with_advanced_memory_sql/README.md +++ b/examples/memory_service_with_advanced_memory_sql/README.md @@ -87,7 +87,7 @@ SQL 版本最核心的构建过程可以简化为三步: sql_url = "mysql+aiomysql://user:password@localhost:3306/trpc_agent_advanced_memory" memory_service = AdvancedMemoryService( - AdvancedCompactConfig( + AdvancedMemoryServiceConfig( storage_backend="sql", sql_url=sql_url, sql_is_async=True, diff --git a/examples/memory_service_with_advanced_memory_sql/run_agent.py b/examples/memory_service_with_advanced_memory_sql/run_agent.py index 6fdd3d2f1..0fbde7f74 100644 --- a/examples/memory_service_with_advanced_memory_sql/run_agent.py +++ b/examples/memory_service_with_advanced_memory_sql/run_agent.py @@ -14,7 +14,7 @@ from dotenv import load_dotenv from agent.agent import create_agent -from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig +from trpc_agent_sdk.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 @@ -60,7 +60,7 @@ def sql_is_async() -> bool: def create_advanced_memory_service(sql_url: str) -> AdvancedMemoryService: """Create the long-term Advanced Memory service backed by SQL.""" memory_ttl = os.getenv("M_TTL") - config = AdvancedCompactConfig( + config = AdvancedMemoryServiceConfig( storage_backend="sql", sql_url=sql_url, sql_is_async=sql_is_async(), diff --git a/examples/session_service_with_advanced_memory_redis/README.md b/examples/session_service_with_advanced_memory_redis/README.md index dff135499..91886f100 100644 --- a/examples/session_service_with_advanced_memory_redis/README.md +++ b/examples/session_service_with_advanced_memory_redis/README.md @@ -15,16 +15,13 @@ ```text AdvancedCompactConfig - ↓ Runner 自动创建 + ↓ AdvancedSessionCompactManager RedisSessionService ├── AdvancedSessionCompactManager ├── events: summary + recent Events ├── historical_events: 被压缩的原始 Events └── state["_trpc_agent:summary"] -AdvancedMemoryRuntime -├── 精简 compression transcript -└── 完整 Tool Result 旁路存储 ``` 核心调用: @@ -34,7 +31,6 @@ session_config = SessionServiceConfig( store_historical_events=True, ) compact_config = AdvancedCompactConfig( - redis_key_prefix="session-compression-demo:v1", model_context_window_tokens=4096, token_autocompact_ratio=0.30, ) @@ -42,7 +38,7 @@ session_service = RedisSessionService( db_url=redis_url, is_async=True, session_config=session_config, - session_compact_config=compact_config, + session_compact_manager=AdvancedSessionCompactManager(config=compact_config), ) runner = Runner( @@ -52,9 +48,8 @@ runner = Runner( ) ``` -`Runner` 会读取 `session_compact_config`,自动从 `RedisSessionService` 获取 URL 和 -异步模式,创建 `AdvancedSessionCompactManager` 并通过基类接口注入。 -用户不需要手动调用 `setup_advanced_session_compact`,也不需要直接创建 Manager。 +`RedisSessionService` 会接收 `session_compact_manager`。Compact 只使用 SessionService 的 +`events`、`historical_events` 和 `state`,不创建额外的 Redis 存储。 ## 兼容已有 Session @@ -97,13 +92,12 @@ python run_agent.py ``` 脚本默认使用 `simple-demo`,可通过 `SESSION_ID` 修改。重复运行可以验证 -活跃窗口、历史原始 Events、Session Memory 和完整 Tool Result 都能跨进程恢复。 +活跃窗口、历史原始 Events 和 Session Memory 都能跨进程恢复。 运行结束会输出 `Active Events`、`Historical Events`、活跃窗口是否以 summary 开头,以及 Session Memory state 是否存在,方便直接确认压缩是否触发。 ## 存储职责 -- `RedisSessionService`:Session、活跃 Events、historical Events、state 和 Session Memory。 -- Advanced Memory Redis stores:压缩重放记录和完整 Tool Result。 -- Redis transcript 不保存 `kind=event`,也不保存 `session-memory-checkpoint`。 +- `RedisSessionService`:Session、活跃 Events、historical Events 和 state。 +- Compact 不创建独立的 Redis transcript、Tool Result 或 session-memory 存储。 diff --git a/examples/session_service_with_advanced_memory_redis/run_agent.py b/examples/session_service_with_advanced_memory_redis/run_agent.py index 4ae9f332d..77efbaca2 100644 --- a/examples/session_service_with_advanced_memory_redis/run_agent.py +++ b/examples/session_service_with_advanced_memory_redis/run_agent.py @@ -5,7 +5,6 @@ # 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 @@ -16,6 +15,7 @@ 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 @@ -43,7 +43,6 @@ def redis_url() -> str: def create_compact_config() -> AdvancedCompactConfig: """Configure only the settings needed to demonstrate one compaction.""" return AdvancedCompactConfig( - redis_key_prefix="session-compression-demo:v1", model_context_window_tokens=4096, max_output_tokens=256, token_warning_ratio=0.25, @@ -64,12 +63,13 @@ async def main() -> None: 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_config=compact_config, + session_compact_manager=compact_manager, ) runner = Runner( app_name=app_name, @@ -78,10 +78,10 @@ async def main() -> None: ) 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.", + "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( diff --git a/examples/session_service_with_advanced_memory_sql/.env b/examples/session_service_with_advanced_memory_sql/.env index 0809508e9..693f8ecb1 100644 --- a/examples/session_service_with_advanced_memory_sql/.env +++ b/examples/session_service_with_advanced_memory_sql/.env @@ -2,10 +2,10 @@ TRPC_AGENT_API_KEY= TRPC_AGENT_BASE_URL= TRPC_AGENT_MODEL_NAME= -MYSQL_USER=root +MYSQL_USER= MYSQL_PASSWORD= -MYSQL_HOST=127.0.0.1 -MYSQL_PORT=3306 +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 index b50276df7..522ce8555 100644 --- a/examples/session_service_with_advanced_memory_sql/README.md +++ b/examples/session_service_with_advanced_memory_sql/README.md @@ -16,17 +16,13 @@ ```text AdvancedCompactConfig - ↓ Runner 自动创建 + ↓ AdvancedSessionCompactManager SqlSessionService ├── AdvancedSessionCompactManager ├── events: summary + recent Events ├── sessions.historical_events: 被压缩的原始 Events └── sessions.state["_trpc_agent:summary"] -AdvancedMemoryRuntime -├── advanced_memory_transcripts -├── advanced_memory_transcript_seen -└── advanced_memory_tool_results ``` 核心调用: @@ -43,7 +39,7 @@ session_service = SqlSessionService( db_url=sql_url, is_async=False, session_config=session_config, - session_compact_config=compact_config, + session_compact_manager=AdvancedSessionCompactManager(config=compact_config), ) runner = Runner( @@ -53,9 +49,8 @@ runner = Runner( ) ``` -`Runner` 会读取 `session_compact_config`,自动从 `SqlSessionService` 获取 URL 和异步 -模式,创建 `AdvancedSessionCompactManager` 并通过基类接口注入。 -用户不需要手动调用 `setup_advanced_session_compact`,也不需要直接创建 Manager。 +`SqlSessionService` 会接收 `session_compact_manager`。Compact 只使用 SessionService 的 +`events`、`historical_events` 和 `state`,不创建额外的 SQL 表。 ## 兼容已有 Session @@ -101,8 +96,5 @@ python run_agent.py ## 存储职责 -- `SqlSessionService`:Session、活跃 Events、historical Events、state 和 Session Memory。 -- Advanced Memory SQL stores:压缩重放记录和完整 Tool Result。 -- 不再创建 `advanced_memory_session_memory` 表。 -- Advanced Memory transcript 不保存 `kind=event` 或 - `session-memory-checkpoint`。 +- `SqlSessionService`:Session、活跃 Events、historical Events 和 state。 +- Compact 不创建独立的 SQL transcript、Tool Result 或 session-memory 表。 diff --git a/examples/session_service_with_advanced_memory_sql/run_agent.py b/examples/session_service_with_advanced_memory_sql/run_agent.py index 7c4274bb2..ffee4b3df 100644 --- a/examples/session_service_with_advanced_memory_sql/run_agent.py +++ b/examples/session_service_with_advanced_memory_sql/run_agent.py @@ -5,7 +5,6 @@ # 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 @@ -16,6 +15,7 @@ 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 @@ -32,10 +32,8 @@ def sql_url() -> str: 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" - ) + return (f"mysql+pymysql://{db_user}:{db_password}@" + f"{db_host}:{db_port}/{db_name}?charset=utf8mb4") def create_compact_config() -> AdvancedCompactConfig: @@ -61,12 +59,13 @@ async def main() -> None: 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_config=compact_config, + session_compact_manager=compact_manager, ) runner = Runner( app_name=app_name, @@ -75,10 +74,10 @@ async def main() -> None: ) 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.", + "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( diff --git a/tests/advanced_memory/test_advanced_memory_tools.py b/tests/advanced_memory/test_advanced_memory_tools.py index 750777d1b..59124377e 100644 --- a/tests/advanced_memory/test_advanced_memory_tools.py +++ b/tests/advanced_memory/test_advanced_memory_tools.py @@ -8,7 +8,7 @@ import pytest -from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig +from trpc_agent_sdk.advanced_memory import AdvancedMemoryServiceConfig from trpc_agent_sdk.advanced_memory import AdvancedMemoryPaths from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime from trpc_agent_sdk.tools import AdvancedMemoryTools @@ -17,7 +17,7 @@ def _runtime(tmp_path: Path) -> AdvancedMemoryRuntime: """Create a test runtime with long-term memory enabled.""" - return AdvancedMemoryRuntime.create(AdvancedCompactConfig( + return AdvancedMemoryRuntime.create(AdvancedMemoryServiceConfig( enabled=True, root_dir=tmp_path, )).for_scope("demo-app", "demo-user") @@ -80,7 +80,7 @@ async def test_list_memory_index_reports_backend_storage_reference( expected_prefix: str, ) -> None: """Avoid exposing a local filesystem path for external memory stores.""" - config = AdvancedCompactConfig( + 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, diff --git a/tests/advanced_memory/test_memory_context.py b/tests/advanced_memory/test_memory_context.py index 703ab3a96..a2385d9d3 100644 --- a/tests/advanced_memory/test_memory_context.py +++ b/tests/advanced_memory/test_memory_context.py @@ -7,36 +7,20 @@ import pytest -from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig +from trpc_agent_sdk.advanced_memory import AdvancedMemoryServiceConfig from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime 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 setup_long_term_memory from trpc_agent_sdk.memory import AdvancedMemoryService from trpc_agent_sdk.models import LlmRequest -from trpc_agent_sdk.sessions.compact import AutoCompactCallback -from trpc_agent_sdk.sessions.compact import HistorySnipCallback -from trpc_agent_sdk.sessions.compact import MicrocompactCallback -from trpc_agent_sdk.sessions.compact import setup_context_compression -from trpc_agent_sdk.sessions.compact import ToolResultBudgetCallback from trpc_agent_sdk.sessions.compact._callbacks import install_staged_callback from trpc_agent_sdk.sessions import InMemorySessionService -from trpc_agent_sdk.sessions import SessionServiceConfig - - -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(AdvancedCompactConfig( + return AdvancedMemoryRuntime.create(AdvancedMemoryServiceConfig( enabled=True, root_dir=tmp_path, )) @@ -108,7 +92,7 @@ async def test_long_term_memory_index_is_injected_once(tmp_path: Path) -> 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( - AdvancedCompactConfig( + AdvancedMemoryServiceConfig( enabled=True, root_dir=tmp_path, memory_focus_instruction="重点记住用户长期稳定的兴趣爱好。", @@ -123,56 +107,6 @@ async def test_custom_memory_focus_is_injected_into_system_instruction(tmp_path: assert "重点记住用户长期稳定的兴趣爱好。" in instruction -async def test_context_setup_installs_four_compaction_stages(tmp_path: Path) -> None: - """Ensure Session compact setup installs only the four compact stages.""" - runtime = _runtime(tmp_path) - agent = SimpleNamespace(before_model_callback=None) - session_service = InMemorySessionService( - session_config=SessionServiceConfig(store_historical_events=True), - ) - - setup_context_compression( - agent, - session_service, - runtime, - FakeSummaryGenerator(), - ) - - assert session_service.session_compact_manager.runtime is 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) - await session_service.close() - - -async def test_explicit_memory_and_compact_setup_compose(tmp_path: Path, ) -> None: - """Ensure long-term memory and Session compact are composed explicitly.""" - runtime = _runtime(tmp_path) - agent = SimpleNamespace(before_model_callback=None, tools=[]) - session_service = InMemorySessionService( - session_config=SessionServiceConfig(store_historical_events=True), - ) - long_term = setup_long_term_memory(agent, runtime) - compact = setup_context_compression( - agent, - session_service, - runtime, - FakeSummaryGenerator(), - ) - - assert compact is session_service - assert session_service.session_compact_manager is not None - assert long_term.tools is not None - assert len(agent.before_model_callback) == 5 - tool_names = {tool.name for tool in agent.tools} - assert tool_names == { - "save_memory", - "read_memory", - "list_memory_index", - } - - 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) @@ -188,19 +122,19 @@ async def test_memory_service_does_not_install_session_compression(tmp_path: Pat agent.before_model_callback[0], LongTermMemoryContextCallback, ) - assert {tool.name - for tool in agent.tools} == { - "save_memory", - "read_memory", - "list_memory_index", - } + 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(AdvancedCompactConfig(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_preload_memory.py b/tests/advanced_memory/test_preload_memory.py index 2f4a3f054..a2c61deb8 100644 --- a/tests/advanced_memory/test_preload_memory.py +++ b/tests/advanced_memory/test_preload_memory.py @@ -5,7 +5,7 @@ from pathlib import Path from types import SimpleNamespace -from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig +from trpc_agent_sdk.advanced_memory import AdvancedMemoryServiceConfig from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime from trpc_agent_sdk.advanced_memory import MemoryDocument from trpc_agent_sdk.advanced_memory import MemoryPreloader @@ -32,7 +32,7 @@ async def select(self, query, candidates, ctx, *, limit): async def test_preloader_injects_selected_topic_with_budget(tmp_path: Path) -> None: """Ensure selected topic content is rendered and bounded.""" runtime = AdvancedMemoryRuntime.create( - AdvancedCompactConfig( + AdvancedMemoryServiceConfig( enabled=True, root_dir=tmp_path, preload_memory_enabled=True, @@ -68,7 +68,7 @@ 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( - AdvancedCompactConfig( + AdvancedMemoryServiceConfig( enabled=True, root_dir=tmp_path, preload_memory_enabled=True, @@ -103,7 +103,7 @@ 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( - AdvancedCompactConfig( + AdvancedMemoryServiceConfig( enabled=True, root_dir=tmp_path, preload_memory_enabled=True, diff --git a/tests/advanced_memory/test_redis_stores.py b/tests/advanced_memory/test_redis_stores.py deleted file mode 100644 index b681bbc9f..000000000 --- a/tests/advanced_memory/test_redis_stores.py +++ /dev/null @@ -1,142 +0,0 @@ -"""Tests for Redis Advanced Memory storage and TTL grouping.""" - -from __future__ import annotations - -from unittest.mock import AsyncMock, MagicMock -from pathlib import Path - -import pytest - -from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig -from trpc_agent_sdk.advanced_memory import AdvancedMemoryPaths -from trpc_agent_sdk.advanced_memory import MemoryIndexEntry -from trpc_agent_sdk.advanced_memory._redis_stores import RedisLongTermMemoryStore -from trpc_agent_sdk.sessions.compact._redis_stores import RedisToolResultStore -from trpc_agent_sdk.sessions.compact._redis_stores import RedisTranscriptStore - - -def _store(store_type: type, **overrides: object): - config = AdvancedCompactConfig( - storage_backend="redis", - redis_url="redis://localhost:6379/0", - root_dir=Path("/tmp/advanced-memory-redis-tests"), - memory_ttl_seconds=120, - session_ttl_seconds=60, - **overrides, - ) - paths = AdvancedMemoryPaths(config).for_scope("app", "user") - store = store_type(config, paths, MagicMock()) - - async def command(method: str, *args: object, **kwargs: object): - if method == "set" and args and str(args[0]).endswith(":memory:lock"): - return True - return [] - - store._command = AsyncMock(side_effect=command) - return store - - -@pytest.mark.asyncio -async def test_memory_writes_refresh_all_memory_keys() -> None: - store = _store(RedisLongTermMemoryStore) - - await store.write_index([ - MemoryIndexEntry(name="Profile", filename="profile.md", summary="User profile"), - ]) - - commands = [call.args for call in store._command.await_args_list] - assert ("set", f"{store._user_base}:memory:index", "- [Profile](profile.md):User profile\n") in commands - assert ("sadd", f"{store._user_base}:memory:keys", f"{store._user_base}:memory:index") in commands - assert ("expire", f"{store._user_base}:memory:index", 120) in commands - assert ("expire", f"{store._user_base}:memory:keys", 120) in commands - - -@pytest.mark.asyncio -async def test_session_writes_refresh_all_session_keys() -> None: - store = _store(RedisToolResultStore) - - await store.write("session-1", "result-1", "complete result") - - session_base = store._session_base("session-1") - commands = [call.args for call in store._command.await_args_list] - tool_key = f"{session_base}:tool:result-1" - assert any(command[0] == "set" and command[1] == tool_key for command in commands) - assert ("sadd", f"{session_base}:keys", tool_key) in commands - assert ("expire", tool_key, 60) in commands - assert ("expire", f"{session_base}:keys", 60) in commands - - -@pytest.mark.asyncio -async def test_ttl_refresh_includes_previously_tracked_keys() -> None: - store = _store(RedisToolResultStore, session_ttl_delete_transcripts=True) - session_base = store._session_base("session-1") - old_key = f"{session_base}:transcript" - store._command = AsyncMock(side_effect=[ - None, # SADD - [old_key.encode()], # SMEMBERS - None, # EXPIRE old key - None, # EXPIRE current key - None, # EXPIRE registry - ]) - - current_key = f"{session_base}:tool:result-1" - await store._refresh_session_ttl("session-1", current_key) - - commands = [call.args for call in store._command.await_args_list] - assert ("expire", old_key, 60) in commands - assert ("expire", current_key, 60) in commands - - -@pytest.mark.asyncio -async def test_ttl_refresh_preserves_transcript_by_default() -> None: - store = _store(RedisToolResultStore) - session_base = store._session_base("session-1") - old_key = f"{session_base}:transcript" - old_seen_key = f"{old_key}:seen:event_id" - store._command = AsyncMock(side_effect=[ - None, # SADD - [old_key.encode(), old_seen_key.encode()], # SMEMBERS - None, # EXPIRE current key - None, # EXPIRE registry - ]) - - current_key = f"{session_base}:tool:result-1" - await store._refresh_session_ttl("session-1", current_key) - - commands = [call.args for call in store._command.await_args_list] - assert ("expire", old_key, 60) not in commands - assert ("expire", old_seen_key, 60) not in commands - assert ("expire", current_key, 60) in commands - - -@pytest.mark.asyncio -async def test_transcript_rejects_event_copies() -> None: - store = _store(RedisTranscriptStore) - - with pytest.raises(ValueError, match="context-compression"): - await store.append( - "session-1", - { - "kind": "event", - "event_id": "event-1" - }, - ) - - -@pytest.mark.asyncio -async def test_memory_write_lock_releases_with_token_check() -> None: - store = _store(RedisLongTermMemoryStore) - - async with store._memory_write_lock(): - pass - - lock_key = f"{store._user_base}:memory:lock" - lock_sets = [call for call in store._command.await_args_list if call.args[:2] == ("set", lock_key)] - releases = [call for call in store._command.await_args_list if call.args and call.args[0] == "eval"] - assert lock_sets - assert lock_sets[0].kwargs["nx"] is True - assert lock_sets[0].kwargs["ex"] == 30 - assert releases - assert releases[0].args[2] == 1 - assert releases[0].args[3] == lock_key - assert releases[0].args[4] == lock_sets[0].args[2] diff --git a/tests/advanced_memory/test_sql_stores.py b/tests/advanced_memory/test_sql_stores.py deleted file mode 100644 index b29745c55..000000000 --- a/tests/advanced_memory/test_sql_stores.py +++ /dev/null @@ -1,113 +0,0 @@ -"""SQLite tests for the Advanced Memory SQL backend.""" - -from __future__ import annotations - -from pathlib import Path - -import pytest - -from trpc_agent_sdk.advanced_memory import ( - AdvancedCompactConfig, - AdvancedMemoryRuntime, - MemoryDocument, - MemoryIndexEntry, - MemoryType, -) - - -def _runtime(tmp_path: Path) -> AdvancedMemoryRuntime: - return AdvancedMemoryRuntime.create( - AdvancedCompactConfig( - storage_backend="sql", - sql_url=f"sqlite:///{tmp_path / 'advanced-memory.db'}", - sql_is_async=False, - memory_ttl_seconds=120, - session_ttl_seconds=60, - )) - - -async def test_sql_stores_round_trip_and_deduplicate(tmp_path: Path) -> None: - root = _runtime(tmp_path) - scoped = root.for_scope("app", "user") - await scoped.initialize() - - await scoped.long_term_memory.write_index([ - MemoryIndexEntry(name="Profile", filename="profile.md", summary="Profile"), - ]) - await scoped.long_term_memory.write_topic( - "profile", - MemoryDocument( - name="Profile", - description="Profile", - memory_type=MemoryType.USER, - content="A user profile", - ), - ) - await scoped.tool_results.write("session", "result", '{"ok": true}') - await scoped.transcripts.append( - "session", - { - "kind": "autocompact-failure", - "attempt_id": "one" - }, - ) - _, first = await scoped.transcripts.append_unique( - "session", - { - "kind": "history-snip", - "snip_id": "two" - }, - unique_key="snip_id", - ) - _, second = await scoped.transcripts.append_unique( - "session", - { - "kind": "history-snip", - "snip_id": "two" - }, - unique_key="snip_id", - ) - - assert first is True - assert second is False - assert "profile.md" in await scoped.long_term_memory.read_index() - assert await scoped.long_term_memory.read_topic("profile") - assert scoped.session_memory is None - assert await scoped.tool_results.read("session", "result") == '{"ok": true}' - assert len(await scoped.transcripts.read_all("session")) == 2 - - await root.close() - - -async def test_sql_transcript_rejects_event_copies(tmp_path: Path) -> None: - root = _runtime(tmp_path) - scoped = root.for_scope("app", "user") - await scoped.initialize() - - with pytest.raises(ValueError, match="context-compression"): - await scoped.transcripts.append( - "session", - { - "kind": "event", - "event_id": "event-1" - }, - ) - - await root.close() - - -async def test_sql_stores_isolate_users(tmp_path: Path) -> None: - root = _runtime(tmp_path) - first = root.for_scope("app", "first") - second = root.for_scope("app", "second") - await first.initialize() - await second.initialize() - - await first.long_term_memory.write_index([ - MemoryIndexEntry(name="First", filename="first.md", summary="First"), - ]) - - assert "first.md" in await first.long_term_memory.read_index() - assert "first.md" not in await second.long_term_memory.read_index() - - await root.close() diff --git a/tests/advanced_memory/test_storage.py b/tests/advanced_memory/test_storage.py deleted file mode 100644 index e4f025cd5..000000000 --- a/tests/advanced_memory/test_storage.py +++ /dev/null @@ -1,466 +0,0 @@ -"""Unit tests for the independent Advanced Memory stores.""" - -from __future__ import annotations - -import asyncio -import json -import os -import threading -from datetime import datetime -from datetime import timezone -from pathlib import Path - -import pytest - -from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig -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 memory_freshness -from trpc_agent_sdk.advanced_memory import parse_memory_updated_at -from trpc_agent_sdk.sessions.compact import SESSION_MEMORY_SECTIONS -from trpc_agent_sdk.sessions.compact import SessionMemoryDocument - - -def _enabled_config(tmp_path: Path, **overrides: object) -> AdvancedCompactConfig: - """Create an enabled configuration rooted at the test directory.""" - return AdvancedCompactConfig(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 = AdvancedCompactConfig() - - 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"): - AdvancedCompactConfig() - - -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"): - AdvancedCompactConfig() - - -def test_config_rejects_unknown_storage_backend(tmp_path: Path) -> None: - """Prevent misspelled external backends from silently using local files.""" - with pytest.raises(ValueError, match="storage_backend must be one of"): - AdvancedCompactConfig( - root_dir=tmp_path, - storage_backend="redisx", # type: ignore[arg-type] - ) - - -async def test_disabled_runtime_does_not_create_directories(tmp_path: Path) -> None: - """Ensure disabled runtime initialization creates no directories.""" - runtime = AdvancedMemoryRuntime.create(AdvancedCompactConfig(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_runtime_close_is_idempotent(tmp_path: Path) -> None: - """Allow a shared Runtime to be closed by more than one service owner.""" - runtime = AdvancedMemoryRuntime.create(_enabled_config(tmp_path)) - await runtime.initialize() - - await runtime.close() - await runtime.close() - - -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_scoped_storage_isolates_users_and_allows_same_session_id(tmp_path: Path) -> None: - """Keep all Advanced Memory records inside the app and user namespace.""" - runtime = AdvancedMemoryRuntime.create(_enabled_config(tmp_path)) - first = runtime.for_scope("demo-app", "user-a") - second = runtime.for_scope("demo-app", "user-b") - await first.initialize() - await second.initialize() - - await first.long_term_memory.write_index([MemoryIndexEntry(name="A", filename="a.md", summary="A")]) - await second.long_term_memory.write_index([MemoryIndexEntry(name="B", filename="b.md", summary="B")]) - await first.session_memory.write("shared", SessionMemoryDocument(session_title="A")) - await second.session_memory.write("shared", SessionMemoryDocument(session_title="B")) - await first.transcripts.append("shared", {"kind": "event", "event_id": "a"}) - await second.transcripts.append("shared", {"kind": "event", "event_id": "b"}) - - assert "a.md" in await first.long_term_memory.read_index() - assert "b.md" not in await first.long_term_memory.read_index() - assert "b.md" in await second.long_term_memory.read_index() - assert (await first.session_memory.read("shared")) != await second.session_memory.read("shared") - assert [record["event_id"] for record in await first.transcripts.read_all("shared")] == ["a"] - assert [record["event_id"] for record in await second.transcripts.read_all("shared")] == ["b"] - assert first.paths.session_dir("shared") != second.paths.session_dir("shared") - - -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_unique_cache_is_reset_after_session_ttl(tmp_path: Path, ) -> None: - """Allow a reused session ID to append after transcript deletion.""" - runtime = AdvancedMemoryRuntime.create( - _enabled_config( - tmp_path, - session_ttl_seconds=1, - session_ttl_delete_transcripts=True, - )) - transcript = runtime.transcripts - await transcript.append_unique( - "session-a", - { - "kind": "event", - "event_id": "event-1" - }, - unique_key="event_id", - ) - activity_path = runtime.paths.session_dir("session-a") / ".advanced-memory-activity" - os.utime(activity_path, (1.0, 1.0)) - - _, appended = await transcript.append_unique( - "session-a", - { - "kind": "event", - "event_id": "event-1" - }, - unique_key="event_id", - ) - - assert appended is True - assert len(await transcript.read_all("session-a")) == 1 - await runtime.close() - - -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 = AdvancedCompactConfig( - 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() == "" - - -async def test_local_ttl_expires_memory_and_session_groups(tmp_path: Path) -> None: - """Expire local memory groups after their last activity.""" - runtime = AdvancedMemoryRuntime.create( - _enabled_config( - tmp_path, - memory_ttl_seconds=1, - session_ttl_seconds=1, - session_ttl_delete_transcripts=True, - )) - scoped = runtime.for_scope("app", "user") - await scoped.initialize() - await scoped.long_term_memory.write_index([ - MemoryIndexEntry(name="Profile", filename="profile.md", summary="Profile"), - ]) - await scoped.long_term_memory.write_topic( - "profile", - MemoryDocument(name="Profile", description="Profile", memory_type=MemoryType.USER, content="data"), - ) - await scoped.session_memory.write("session", SessionMemoryDocument(session_title="Session")) - await scoped.tool_results.write("session", "result", "data") - await scoped.transcripts.append("session", {"event_id": "event"}) - - old = 1.0 - os.utime(scoped.paths.memory_index_path, (old, old)) - os.utime(scoped.paths.session_dir("session") / ".advanced-memory-activity", (old, old)) - - assert await scoped.long_term_memory.read_index() == "" - assert await scoped.long_term_memory.read_topic("profile") is None - assert await scoped.session_memory.read("session") is None - assert not scoped.paths.session_dir("session").exists() - await runtime.close() - - -async def test_local_session_ttl_preserves_transcripts_by_default(tmp_path: Path) -> None: - """Keep local transcripts when session TTL cleanup uses its default.""" - runtime = AdvancedMemoryRuntime.create(_enabled_config( - tmp_path, - session_ttl_seconds=1, - )) - scoped = runtime.for_scope("app", "user") - await scoped.initialize() - await scoped.session_memory.write("session", SessionMemoryDocument(session_title="Session")) - await scoped.transcripts.append("session", {"event_id": "event"}) - - activity_path = scoped.paths.session_dir("session") / ".advanced-memory-activity" - os.utime(activity_path, (1.0, 1.0)) - - assert await scoped.session_memory.read("session") is None - assert scoped.paths.transcript_path("session").exists() - records = await scoped.transcripts.read_all("session") - assert len(records) == 1 - assert records[0]["event_id"] == "event" - await runtime.close() - - -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"): - AdvancedCompactConfig(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/sessions/compact/test_autocompact.py b/tests/sessions/compact/test_autocompact.py deleted file mode 100644 index 03eaba20b..000000000 --- a/tests/sessions/compact/test_autocompact.py +++ /dev/null @@ -1,552 +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.sessions.compact import AutoCompact -from trpc_agent_sdk.sessions.compact import AutoCompactCallback -from trpc_agent_sdk.sessions.compact import AdvancedCompactConfig -from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime -from trpc_agent_sdk.sessions.compact import HistorySnipCallback -from trpc_agent_sdk.sessions.compact import MicrocompactCallback -from trpc_agent_sdk.sessions.compact import SessionMemoryDocument -from trpc_agent_sdk.sessions.compact import setup_autocompact -from trpc_agent_sdk.sessions.compact import setup_history_snip -from trpc_agent_sdk.sessions.compact import setup_microcompact -from trpc_agent_sdk.sessions.compact import setup_tool_result_budget -from trpc_agent_sdk.sessions.compact 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.sessions import SessionServiceConfig -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( - AdvancedCompactConfig( - 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, - )).for_scope("demo-app", "demo-user") - - -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", - session=SimpleNamespace( - app_name="demo-app", - user_id="demo-user", - id=session_id, - ), - 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_compact_persists_summary_and_archives_replaced_events(tmp_path: Path) -> None: - """Ensure AutoCompact writes the compressed window through SessionService.""" - runtime = _runtime(tmp_path) - service = InMemorySessionService( - session_config=SessionServiceConfig(store_historical_events=True), - ) - session = await service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="session-a", - ) - request = _request(5) - for index, content in enumerate(request.contents): - await service.append_event( - session, - Event( - id=f"event-{index}", - invocation_id="invocation-1", - author="user" if index % 2 == 0 else "agent", - content=content.model_copy(deep=True), - ), - ) - ctx = SimpleNamespace( - session_id=session.id, - app_name=session.app_name, - session=session, - session_service=service, - agent=SimpleNamespace(model="fake-model"), - ) - - result = await AutoCompact(runtime, FakeSummaryGenerator()).apply( - request, - session_id=session.id, - ctx=ctx, - force=True, - ) - - assert result.compacted - restored = await service.get_session( - app_name=session.app_name, - user_id=session.user_id, - session_id=session.id, - ) - assert restored is not None - assert restored.events[0].is_summary_event() - assert [event.id for event in restored.events[1:]] == ["event-3", "event-4"] - assert [event.id for event in restored.historical_events] == [ - "event-0", - "event-1", - "event-2", - ] - assert not restored.compact_events( - Event(author="system", content=Content(parts=[Part.from_text(text="duplicate")])), - "event-2", - compaction_id=restored.events[0].custom_metadata["session_compaction_id"], - ) - - -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( - AdvancedCompactConfig( - 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, - )).for_scope("demo-app", "demo-user") - 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_token_reduction_uses_consistent_full_request_estimates(tmp_path: Path) -> None: - """Do not compare a usage-based before value with an estimated after value.""" - runtime = AdvancedMemoryRuntime.create( - AdvancedCompactConfig( - enabled=True, - root_dir=tmp_path, - model_context_window_tokens=20_000, - max_output_tokens=100, - token_warning_ratio=0.4, - token_autocompact_ratio=0.5, - autocompact_keep_recent_contents=2, - )).for_scope("demo-app", "demo-user") - request = _request(5) - ctx = _ctx() - ctx.session.events = [ - SimpleNamespace( - content=request.contents[0].model_copy(deep=True), - usage_metadata=SimpleNamespace(total_token_count=12_000), - custom_metadata={}, - ), - ] - - result = await AutoCompact(runtime, FakeSummaryGenerator()).apply( - request, - session_id="session-a", - ctx=ctx, - ) - - assert result.compacted - assert result.request_tokens_after < result.request_tokens_before - assert result.request_tokens_before < 12_000 - assert result.token_source == "estimated" - - -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 / "tenants" / "demo-app" / "demo-user" / "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/sessions/compact/test_context_compression_integration.py b/tests/sessions/compact/test_context_compression_integration.py deleted file mode 100644 index 6ac51305d..000000000 --- a/tests/sessions/compact/test_context_compression_integration.py +++ /dev/null @@ -1,452 +0,0 @@ -"""Tests for request compression over an unchanged SessionService.""" - -from pathlib import Path -from types import SimpleNamespace - -import pytest - -from trpc_agent_sdk.evaluation._eval_session_service import EvalSessionService -from trpc_agent_sdk.sessions.compact import AdvancedCompactConfig -from trpc_agent_sdk.sessions.compact import BaseSessionCompactConfig -from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime -from trpc_agent_sdk.sessions.compact import BaseSessionCompactManager -from trpc_agent_sdk.sessions.compact import AutoCompactCallback -from trpc_agent_sdk.sessions.compact import HistorySnipCallback -from trpc_agent_sdk.sessions.compact import MicrocompactCallback -from trpc_agent_sdk.sessions.compact import SESSION_MEMORY_STATE_KEY -from trpc_agent_sdk.sessions.compact import SessionMemoryDocument -from trpc_agent_sdk.sessions.compact import ToolResultBudget -from trpc_agent_sdk.sessions.compact import ToolResultBudgetCallback -from trpc_agent_sdk.sessions.compact import setup_advanced_session_compact -from trpc_agent_sdk.sessions.compact import setup_context_compression -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 import SqlSessionService -from trpc_agent_sdk.types import Content -from trpc_agent_sdk.types import FunctionResponse -from trpc_agent_sdk.types import Part - - -class FakeSummaryGenerator: - """Return a deterministic autocompact summary.""" - - async def generate(self, history: str, ctx) -> str: - del history, ctx - return "summary" - - -class FakeSessionMemoryGenerator: - """Return deterministic structured Session Memory.""" - - async def generate(self, extraction_input, ctx) -> SessionMemoryDocument: - del ctx - return SessionMemoryDocument( - session_title="Post-turn memory", - current_state=f"Processed {extraction_input.last_event_id}", - ) - - -class DummySummarizerManager: - """Provide the BaseSessionService attachment protocol.""" - - def set_session_service(self, service) -> None: - self.service = service - - -def _runtime(tmp_path: Path) -> AdvancedMemoryRuntime: - return AdvancedMemoryRuntime.create( - AdvancedCompactConfig( - root_dir=tmp_path, - tool_result_max_chars=200, - tool_results_per_message_max_chars=5_000, - tool_result_preview_chars=40, - )) - - -def _session_service() -> InMemorySessionService: - return InMemorySessionService( - session_config=SessionServiceConfig(store_historical_events=True), - ) - - -async def test_session_service_accepts_base_compact_manager(tmp_path: Path) -> None: - """Inject the Advanced manager through the common manager contract.""" - agent = SimpleNamespace(before_model_callback=None) - service = InMemorySessionService( - session_config=SessionServiceConfig(store_historical_events=True), - ) - manager = setup_advanced_session_compact( - agent, - service, - AdvancedCompactConfig(root_dir=tmp_path), - session_memory_generator=FakeSessionMemoryGenerator(), - ) - - assert isinstance(service.session_compact_manager, BaseSessionCompactManager) - assert service.session_compact_manager is manager - await service.close() - - -def test_advanced_config_implements_compact_config_contract() -> None: - """Concrete strategies must be selectable through the config base class.""" - assert issubclass(AdvancedCompactConfig, BaseSessionCompactConfig) - - -async def test_advanced_setup_infers_sql_backend_from_session_service( - tmp_path: Path, -) -> None: - """Use the SessionService as the single source of backend settings.""" - database_url = f"sqlite:///{tmp_path / 'compact.db'}" - service = SqlSessionService( - db_url=database_url, - is_async=False, - session_config=SessionServiceConfig(store_historical_events=True), - ) - manager = setup_advanced_session_compact( - SimpleNamespace(before_model_callback=None), - service, - AdvancedCompactConfig(root_dir=tmp_path), - session_memory_generator=FakeSessionMemoryGenerator(), - ) - - assert manager.runtime.config.storage_backend == "sql" - assert manager.runtime.config.sql_url == database_url - assert manager.runtime.config.sql_is_async is False - await service.close() - - -@pytest.mark.asyncio -async def test_runner_auto_installs_compact_from_session_config(tmp_path: Path) -> None: - """Let Runner create the manager from the declarative SessionService config.""" - from trpc_agent_sdk.runners import Runner - - agent = SimpleNamespace( - name="compact-agent", - tools=[], - before_model_callback=None, - get_subagents=lambda: [], - ) - service = InMemorySessionService( - session_config=SessionServiceConfig(store_historical_events=True), - session_compact_config=AdvancedCompactConfig(root_dir=tmp_path), - ) - - runner = Runner( - app_name="compact-test", - agent=agent, - session_service=service, - enable_post_turn_processing=False, - ) - - assert service.session_compact_manager is not None - assert service.session_compact_manager.runtime.config.root_dir == tmp_path.resolve() - await runner.close() - - -def _tool_event(output: str) -> Event: - return Event( - id="event-1", - invocation_id="invocation-1", - author="user", - content=Content(parts=[ - Part(function_response=FunctionResponse( - id="result-1", - name="demo_tool", - response={"output": output}, - )) - ]), - ) - - -async def test_setup_attaches_manager_to_original_service(tmp_path: Path) -> None: - """Install only the four request callbacks over the original service.""" - runtime = _runtime(tmp_path) - delegate = _session_service() - agent = SimpleNamespace(before_model_callback=None) - - service = setup_context_compression( - agent, - delegate, - runtime, - FakeSummaryGenerator(), - ) - - assert service is delegate - assert service.session_compact_manager is not None - assert service.session_compact_manager.runtime is runtime - assert [type(callback) for callback in agent.before_model_callback] == [ - ToolResultBudgetCallback, - HistorySnipCallback, - MicrocompactCallback, - AutoCompactCallback, - ] - - -async def test_setup_rejects_original_session_summarizer(tmp_path: Path) -> None: - """Prevent two independent mechanisms from writing summary Events.""" - delegate = InMemorySessionService( - summarizer_manager=DummySummarizerManager(), - session_config=SessionServiceConfig(store_historical_events=True), - ) - agent = SimpleNamespace(before_model_callback=None) - - with pytest.raises(ValueError, match="mutually exclusive"): - setup_context_compression( - agent, - delegate, - _runtime(tmp_path), - FakeSummaryGenerator(), - ) - await delegate.close() - - -async def test_manager_keeps_events_in_original_service_only(tmp_path: Path) -> None: - """Read and append Events without a second Event transcript.""" - runtime = _runtime(tmp_path) - delegate = _session_service() - session = await delegate.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="legacy-session", - ) - old_event = Event( - id="old-event", - invocation_id="invocation-1", - author="user", - content=Content(parts=[Part.from_text(text="old event")]), - ) - await delegate.append_event(session, old_event) - - agent = SimpleNamespace(before_model_callback=None) - service = setup_context_compression(agent, delegate, runtime, FakeSummaryGenerator()) - loaded = await service.get_session( - app_name="demo-app", - user_id="demo-user", - session_id=session.id, - ) - assert loaded is not None - await service.append_event(loaded, _tool_event("x" * 500)) - - stored = await delegate.get_session( - app_name="demo-app", - user_id="demo-user", - session_id=session.id, - ) - assert stored is not None - assert [event.id for event in stored.events] == ["old-event", "event-1"] - assert await runtime.for_session(stored).transcripts.read_all(stored.id) == [] - - -async def test_request_replacement_does_not_rewrite_stored_event(tmp_path: Path) -> None: - """Replace a request copy while retaining the complete persisted result.""" - runtime = _runtime(tmp_path) - delegate = _session_service() - agent = SimpleNamespace(before_model_callback=None) - service = setup_context_compression(agent, delegate, runtime, FakeSummaryGenerator()) - session = await service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="budget-session", - ) - await service.append_event(session, _tool_event("x" * 500)) - request = LlmRequest( - model="test-model", - contents=[session.events[0].content.model_copy(deep=True)], - ) - - result = await ToolResultBudget(runtime.for_session(session)).apply( - request, - session_id=session.id, - ) - - stored = await delegate.get_session( - app_name="demo-app", - user_id="demo-user", - session_id=session.id, - ) - assert result.replaced_count == 1 - assert "persisted_output" in request.contents[0].parts[0].function_response.response - assert stored is not None - assert stored.events[0].content.parts[0].function_response.response == { - "output": "x" * 500 - } - records = await runtime.for_session(stored).transcripts.read_all(stored.id) - assert all(record.get("kind") != "event" for record in records) - - -async def test_setup_is_idempotent_and_validates_runtime_first(tmp_path: Path) -> None: - """Reuse one manager and reject a different runtime without changing callbacks.""" - runtime = _runtime(tmp_path / "one") - agent = SimpleNamespace(before_model_callback=None) - service = setup_context_compression( - agent, - _session_service(), - runtime, - FakeSummaryGenerator(), - ) - repeated = setup_context_compression(agent, service, runtime, FakeSummaryGenerator()) - assert repeated is service - assert len(agent.before_model_callback) == 4 - - clean_agent = SimpleNamespace(before_model_callback=None) - with pytest.raises(ValueError, match="another runtime"): - setup_context_compression( - clean_agent, - service, - _runtime(tmp_path / "two"), - FakeSummaryGenerator(), - ) - assert clean_agent.before_model_callback is None - - -async def test_compact_manager_is_mutually_exclusive_with_native_summarizer(tmp_path: Path) -> None: - """Prevent adding the native summarizer after compact setup.""" - service = _session_service() - setup_context_compression( - SimpleNamespace(before_model_callback=None), - service, - _runtime(tmp_path), - FakeSummaryGenerator(), - ) - - with pytest.raises(ValueError, match="mutually exclusive"): - service.set_summarizer_manager(DummySummarizerManager()) - - -async def test_original_service_delete_cleans_compact_side_data(tmp_path: Path) -> None: - """Run compact cleanup through the original SessionService lifecycle.""" - runtime = _runtime(tmp_path) - service = _session_service() - setup_context_compression( - SimpleNamespace(before_model_callback=None), - service, - runtime, - FakeSummaryGenerator(), - ) - session = await service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="delete-me", - ) - scoped = runtime.for_session(session) - await scoped.transcripts.append(session.id, {"kind": "test-record"}) - - await service.delete_session( - app_name=session.app_name, - user_id=session.user_id, - session_id=session.id, - ) - - assert await scoped.transcripts.read_all(session.id) == [] - - -async def test_eval_session_service_forwards_compact_manager(tmp_path: Path) -> None: - """Keep evaluation wrappers on the inner service's compact lifecycle.""" - inner = _session_service() - service = EvalSessionService(inner) - runtime = _runtime(tmp_path) - - configured = setup_context_compression( - SimpleNamespace(before_model_callback=None), - service, - runtime, - FakeSummaryGenerator(), - ) - - assert configured is service - assert service.session_compact_manager is inner.session_compact_manager - assert service.session_compact_manager.runtime is runtime - - -async def test_sql_delegate_keeps_its_existing_event_storage(tmp_path: Path) -> None: - """Ensure manager composition works with the SQL SessionService.""" - runtime = _runtime(tmp_path / "advanced") - delegate = SqlSessionService( - db_url=f"sqlite:///{tmp_path / 'sessions.db'}", - is_async=False, - ) - session = await delegate.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="sql-session", - ) - event = Event( - id="sql-event", - invocation_id="invocation-1", - author="user", - content=Content(parts=[Part.from_text(text="stored by SQL")]), - ) - await delegate.append_event(session, event) - - service = setup_context_compression( - SimpleNamespace(before_model_callback=None), - delegate, - runtime, - FakeSummaryGenerator(), - ) - loaded = await service.get_session( - app_name="demo-app", - user_id="demo-user", - session_id=session.id, - ) - - assert loaded is not None - assert [item.id for item in loaded.events] == ["sql-event"] - assert await runtime.for_session(loaded).transcripts.read_all(loaded.id) == [] - await service.close() - await runtime.close() - - -async def test_post_turn_hook_updates_session_memory_state(tmp_path: Path) -> None: - """Ensure the existing Runner summary hook updates Session Memory.""" - database = tmp_path / "post-turn.db" - runtime = AdvancedMemoryRuntime.create( - AdvancedCompactConfig( - storage_backend="sql", - sql_url=f"sqlite:///{database}", - sql_is_async=False, - session_memory_initial_chars=1, - session_memory_update_chars=1, - ), - ) - delegate = SqlSessionService( - db_url=f"sqlite:///{database}", - is_async=False, - ) - agent = SimpleNamespace(before_model_callback=None) - service = setup_context_compression( - agent, - delegate, - runtime, - FakeSummaryGenerator(), - session_memory_generator=FakeSessionMemoryGenerator(), - ) - session = await service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="post-turn", - ) - await service.append_event(session, _tool_event("post-turn content")) - ctx = SimpleNamespace( - session=session, - session_service=service, - agent=SimpleNamespace(model="fake-model"), - ) - - await service.create_session_summary(session, ctx=ctx) - - assert SESSION_MEMORY_STATE_KEY in session.state - loaded = await service.get_session( - app_name=session.app_name, - user_id=session.user_id, - session_id=session.id, - ) - assert loaded is not None - assert SESSION_MEMORY_STATE_KEY in loaded.state - summary = await service.get_session_summary(loaded) - assert summary is not None - assert "Post-turn memory" in summary - await service.close() - await runtime.close() diff --git a/tests/sessions/compact/test_history_snip.py b/tests/sessions/compact/test_history_snip.py deleted file mode 100644 index 7c39c88d8..000000000 --- a/tests/sessions/compact/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.sessions.compact import AdvancedCompactConfig -from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime -from trpc_agent_sdk.sessions.compact import HistorySnip -from trpc_agent_sdk.sessions.compact import HistorySnipCallback -from trpc_agent_sdk.sessions.compact import Microcompact -from trpc_agent_sdk.sessions.compact import MicrocompactCallback -from trpc_agent_sdk.sessions.compact import setup_history_snip -from trpc_agent_sdk.sessions.compact import setup_microcompact -from trpc_agent_sdk.sessions.compact import setup_tool_result_budget -from trpc_agent_sdk.sessions.compact import ToolResultBudgetCallback -from trpc_agent_sdk.sessions.compact 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( - AdvancedCompactConfig( - 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( - AdvancedCompactConfig( - 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( - AdvancedCompactConfig( - 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/sessions/compact/test_microcompact.py b/tests/sessions/compact/test_microcompact.py deleted file mode 100644 index 4b76d961d..000000000 --- a/tests/sessions/compact/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.sessions.compact import AdvancedCompactConfig -from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime -from trpc_agent_sdk.sessions.compact import Microcompact -from trpc_agent_sdk.sessions.compact import MicrocompactCallback -from trpc_agent_sdk.sessions.compact import setup_microcompact -from trpc_agent_sdk.sessions.compact import setup_tool_result_budget -from trpc_agent_sdk.sessions.compact 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( - AdvancedCompactConfig( - 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/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/sessions/compact/test_session_memory_extractor.py b/tests/sessions/compact/test_session_memory_extractor.py deleted file mode 100644 index 5ebdf4dc0..000000000 --- a/tests/sessions/compact/test_session_memory_extractor.py +++ /dev/null @@ -1,522 +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.sessions.compact import AdvancedCompactConfig -from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime -from trpc_agent_sdk.sessions.compact import ForkedSessionMemoryGenerator -from trpc_agent_sdk.sessions.compact import SessionMemoryExtractionInput -from trpc_agent_sdk.sessions.compact import SessionMemoryDocument -from trpc_agent_sdk.sessions.compact import SessionMemoryExtractor -from trpc_agent_sdk.sessions.compact 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( - AdvancedCompactConfig( - 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")) - - -def _scoped(runtime: AdvancedMemoryRuntime): - """Return the tenant runtime used by the test sessions.""" - return runtime.for_scope("demo-app", "demo-user") - - -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), - ) - - scoped = _scoped(runtime) - memory = await scoped.session_memory.read(session.id) - records = await scoped.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( - AdvancedCompactConfig( - 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 _scoped(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 _scoped(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 _scoped(runtime).transcripts.read_all(session.id) - assert result.reason == "extraction-failed" - assert await _scoped(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 _scoped(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/sessions/compact/test_session_memory_state.py b/tests/sessions/compact/test_session_memory_state.py deleted file mode 100644 index ee0b60492..000000000 --- a/tests/sessions/compact/test_session_memory_state.py +++ /dev/null @@ -1,160 +0,0 @@ -"""Session-state persistence tests for Redis/SQL Advanced Memory.""" - -from pathlib import Path -from types import SimpleNamespace - -from trpc_agent_sdk.sessions.compact import AdvancedCompactConfig -from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime -from trpc_agent_sdk.sessions.compact import AutoCompact -from trpc_agent_sdk.sessions.compact import SessionMemoryDocument -from trpc_agent_sdk.sessions.compact import SessionMemoryExtractor -from trpc_agent_sdk.sessions.compact._formats import SESSION_MEMORY_STATE_KEY -from trpc_agent_sdk.sessions.compact._formats import parse_session_memory_state -from trpc_agent_sdk.events import Event -from trpc_agent_sdk.models import LlmRequest -from trpc_agent_sdk.sessions import SqlSessionService -from trpc_agent_sdk.types import Content -from trpc_agent_sdk.types import Part - - -class _Generator: - - def __init__(self) -> None: - self.inputs = [] - - async def generate(self, extraction_input, ctx) -> SessionMemoryDocument: - del ctx - self.inputs.append(extraction_input) - return SessionMemoryDocument( - session_title="State-backed session", - current_state=f"Processed {extraction_input.last_event_id}", - ) - - -class _LegacyGenerator: - - async def generate(self, history, ctx) -> str: - del history, ctx - return "legacy" - - -def _event(event_id: str, text: str) -> Event: - return Event( - id=event_id, - invocation_id="invocation", - author="agent", - content=Content(role="model", parts=[Part.from_text(text=text)]), - ) - - -async def test_sql_session_memory_is_persisted_in_session_state(tmp_path: Path, ) -> None: - database = tmp_path / "state-memory.db" - runtime = AdvancedMemoryRuntime.create( - AdvancedCompactConfig( - storage_backend="sql", - sql_url=f"sqlite:///{database}", - sql_is_async=False, - session_memory_initial_chars=1, - session_memory_update_chars=1, - )) - service = SqlSessionService(db_url=f"sqlite:///{database}", is_async=False) - session = await service.create_session( - app_name="app", - user_id="user", - session_id="session", - ) - await service.append_event(session, _event("event-1", "x" * 2_000)) - generator = _Generator() - extractor = SessionMemoryExtractor( - runtime, - generator, - session_service=service, - ) - ctx = SimpleNamespace( - session=session, - agent=SimpleNamespace(model="test-model"), - ) - - result = await extractor.extract_if_needed(session, ctx, force=True) - - loaded = await service.get_session( - app_name="app", - user_id="user", - session_id="session", - ) - assert result.extracted is True - assert loaded is not None - parsed = parse_session_memory_state(loaded.state[SESSION_MEMORY_STATE_KEY]) - assert parsed is not None - document, checkpoint, _ = parsed - assert document.current_state == "Processed event-1" - assert checkpoint["last_event_id"] == "event-1" - assert len(loaded.events) == 1 - assert runtime.for_session(loaded).session_memory is None - assert await runtime.for_session(loaded).transcripts.read_all(loaded.id) == [] - await service.close() - await runtime.close() - - -async def test_autocompact_generates_state_memory_only_when_invoked(tmp_path: Path, ) -> None: - database = tmp_path / "autocompact-state.db" - runtime = AdvancedMemoryRuntime.create( - AdvancedCompactConfig( - storage_backend="sql", - sql_url=f"sqlite:///{database}", - sql_is_async=False, - autocompact_target_chars=20_000, - session_memory_initial_chars=1, - session_memory_update_chars=1, - )) - service = SqlSessionService(db_url=f"sqlite:///{database}", is_async=False) - session = await service.create_session( - app_name="app", - user_id="user", - session_id="session", - ) - for index in range(3): - await service.append_event( - session, - _event(f"event-{index}", f"message-{index}-" + "x" * 3_000), - ) - generator = _Generator() - extractor = SessionMemoryExtractor( - runtime, - generator, - session_service=service, - ) - compressor = AutoCompact(runtime, _LegacyGenerator()) - compressor.attach_session_memory_extractor(extractor) - ctx = SimpleNamespace( - session=session, - session_service=service, - agent=SimpleNamespace(model="test-model"), - ) - request = LlmRequest( - model="test-model", - contents=[event.content.model_copy(deep=True) for event in session.events], - ) - - result = await compressor.apply( - request, - session_id=session.id, - ctx=ctx, - force=True, - ) - - assert result.compacted is True - assert result.source == "session-memory" - assert generator.inputs - assert SESSION_MEMORY_STATE_KEY in session.state - assert session.events[0].is_summary_event() - assert [event.id for event in session.historical_events] == [ - "event-0", - "event-1", - "event-2", - ] - records = await runtime.for_session(session).transcripts.read_all(session.id) - assert [record["kind"] for record in records] == ["autocompact-success"] - assert all(record["kind"] != "event" for record in records) - await service.close() - await runtime.close() diff --git a/tests/sessions/compact/test_token_budget.py b/tests/sessions/compact/test_token_budget.py index 4e2a29fbb..9e2c918a6 100644 --- a/tests/sessions/compact/test_token_budget.py +++ b/tests/sessions/compact/test_token_budget.py @@ -39,7 +39,6 @@ def test_usage_baseline_adds_only_contents_after_matching_event(tmp_path) -> Non tracker = TokenContextTracker( 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(AdvancedCompactConfig(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(AdvancedCompactConfig(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 @@ -96,7 +95,6 @@ def test_budget_reserves_max_output_and_calculates_three_thresholds(tmp_path) -> tracker = TokenContextTracker( 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(AdvancedCompactConfig(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/compact/test_tool_result_budget.py b/tests/sessions/compact/test_tool_result_budget.py deleted file mode 100644 index 2a3e806b4..000000000 --- a/tests/sessions/compact/test_tool_result_budget.py +++ /dev/null @@ -1,332 +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.sessions.compact import AdvancedCompactConfig -from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime -from trpc_agent_sdk.sessions.compact import setup_tool_result_budget -from trpc_agent_sdk.sessions.compact import ToolResultBudget -from trpc_agent_sdk.sessions.compact 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( - AdvancedCompactConfig( - 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_sql_replacement_reports_sql_storage_path(tmp_path: Path) -> None: - """Expose the path returned by the SQL tool-result store.""" - root = AdvancedMemoryRuntime.create(AdvancedCompactConfig( - enabled=True, - storage_backend="sql", - sql_url=f"sqlite:///{tmp_path / 'memory.db'}", - sql_is_async=False, - tool_result_max_chars=200, - tool_results_per_message_max_chars=5_000, - tool_result_preview_chars=40, - )) - runtime = root.for_scope("demo-app", "demo-user") - budget = ToolResultBudget(runtime) - request, _ = _request(("result-1", "x" * 500)) - - await budget.apply(request, session_id="session-a") - - replacement = request.contents[0].parts[0].function_response.response - assert replacement["persisted_output"]["path"].startswith("advanced-memory://sql/") - assert await runtime.tool_results.read("session-a", "result-1") is not None - - -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/sessions/compact/test_transcript_session_service.py b/tests/sessions/compact/test_transcript_session_service.py deleted file mode 100644 index f4a20fe7f..000000000 --- a/tests/sessions/compact/test_transcript_session_service.py +++ /dev/null @@ -1,138 +0,0 @@ -"""Unit tests for TranscriptSessionService automatic recording.""" - -from __future__ import annotations - -from pathlib import Path - -import pytest - -from trpc_agent_sdk.sessions.compact import AdvancedCompactConfig -from trpc_agent_sdk.sessions.compact import AdvancedMemoryRuntime -from trpc_agent_sdk.sessions.compact 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(AdvancedCompactConfig(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.for_session(session).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(AdvancedCompactConfig(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.for_session(session).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(AdvancedCompactConfig(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.for_session(session).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(AdvancedCompactConfig(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(AdvancedCompactConfig(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.for_session(session).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(AdvancedCompactConfig(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_nested_transcript_wrapper_is_rejected(tmp_path: Path) -> None: - """Ensure a transcript decorator cannot wrap another decorator.""" - runtime = AdvancedMemoryRuntime.create(AdvancedCompactConfig(enabled=True, root_dir=tmp_path)) - inner = TranscriptSessionService(InMemorySessionService(), runtime) - - with pytest.raises(ValueError, match="already wrapped"): - TranscriptSessionService(inner, runtime) - - -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(AdvancedCompactConfig(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.for_session(session).transcripts.read_all(session.id) == [] diff --git a/trpc_agent_sdk/advanced_memory/__init__.py b/trpc_agent_sdk/advanced_memory/__init__.py index 658252342..0f7fe7ca5 100644 --- a/trpc_agent_sdk/advanced_memory/__init__.py +++ b/trpc_agent_sdk/advanced_memory/__init__.py @@ -5,17 +5,17 @@ # tRPC-Agent-Python is licensed under Apache-2.0. """Optional long-term memory APIs.""" -from trpc_agent_sdk.sessions.compact._config import AdvancedCompactConfig +from ._config import AdvancedMemoryServiceConfig from trpc_agent_sdk.sessions.compact._formats import MemoryDocument from trpc_agent_sdk.sessions.compact._formats import MemoryIndexEntry from trpc_agent_sdk.sessions.compact._formats import MemoryType from trpc_agent_sdk.sessions.compact._formats import memory_freshness from trpc_agent_sdk.sessions.compact._formats import parse_memory_updated_at -from trpc_agent_sdk.sessions.compact._paths import AdvancedMemoryPaths -from trpc_agent_sdk.sessions.compact._paths import MemoryScope -from trpc_agent_sdk.sessions.compact._runtime import AdvancedMemoryRuntime -from trpc_agent_sdk.sessions.compact._runtime import ScopedAdvancedMemoryRuntime -from trpc_agent_sdk.sessions.compact._storage import LongTermMemoryStore +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 @@ -32,7 +32,7 @@ __all__ = [ "AdvancedMemoryStorageBackend", - "AdvancedCompactConfig", + "AdvancedMemoryServiceConfig", "LongTermMemoryIntegration", "AdvancedMemoryPaths", "AdvancedMemoryRuntime", diff --git a/trpc_agent_sdk/advanced_memory/_config.py b/trpc_agent_sdk/advanced_memory/_config.py new file mode 100644 index 000000000..99588568c --- /dev/null +++ b/trpc_agent_sdk/advanced_memory/_config.py @@ -0,0 +1,83 @@ +# 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/advanced_memory/_integration.py b/trpc_agent_sdk/advanced_memory/_integration.py index b68d8577b..3c01e02f6 100644 --- a/trpc_agent_sdk/advanced_memory/_integration.py +++ b/trpc_agent_sdk/advanced_memory/_integration.py @@ -11,7 +11,7 @@ from typing import Any from typing import TYPE_CHECKING -from trpc_agent_sdk.sessions.compact._runtime import AdvancedMemoryRuntime +from ._runtime import AdvancedMemoryRuntime from ._memory_context import LongTermMemoryContext from ._memory_context import setup_long_term_memory_context @@ -35,38 +35,22 @@ def _setup_long_term_memory_tools( ) -> "AdvancedMemoryTools": """Install the three official memory tools idempotently.""" from trpc_agent_sdk.tools._advanced_memory_tool import ( - ADVANCED_MEMORY_TOOL_NAMES, - ) + 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 - ] + 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 - } + 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" - ) + 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" - ) + 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 - } + 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" - ) + raise ValueError("Advanced Memory tools are only partially installed") return owner tools = AdvancedMemoryTools(memory_runtime) agent.tools.extend(tools.as_tools()) @@ -79,39 +63,28 @@ def _setup_preload_memory_tool( 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 - ): + 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" - ] + 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" - ) + 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, - ) - ) + agent.tools.append(PreloadMemoryTool( + memory_preloader=preloader.preload, + use_legacy_memory=use_legacy_memory, + )) def setup_long_term_memory( @@ -123,11 +96,8 @@ def setup_long_term_memory( ) -> 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 - ) + 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, diff --git a/trpc_agent_sdk/advanced_memory/_memory_context.py b/trpc_agent_sdk/advanced_memory/_memory_context.py index 73bbfd336..b3e62a9df 100644 --- a/trpc_agent_sdk/advanced_memory/_memory_context.py +++ b/trpc_agent_sdk/advanced_memory/_memory_context.py @@ -10,7 +10,7 @@ from typing import TYPE_CHECKING from trpc_agent_sdk.sessions.compact._callbacks import install_staged_callback -from trpc_agent_sdk.sessions.compact._runtime import AdvancedMemoryRuntime +from ._runtime import AdvancedMemoryRuntime if TYPE_CHECKING: from trpc_agent_sdk.agents import LlmAgent diff --git a/trpc_agent_sdk/advanced_memory/_paths.py b/trpc_agent_sdk/advanced_memory/_paths.py new file mode 100644 index 000000000..768c86ff4 --- /dev/null +++ b/trpc_agent_sdk/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/advanced_memory/_preload_memory.py index 7866bb603..d032a8c42 100644 --- a/trpc_agent_sdk/advanced_memory/_preload_memory.py +++ b/trpc_agent_sdk/advanced_memory/_preload_memory.py @@ -23,7 +23,7 @@ from trpc_agent_sdk.sessions import InMemorySessionService from trpc_agent_sdk.sessions.compact._formats import memory_freshness from trpc_agent_sdk.sessions.compact._formats import parse_memory_updated_at -from trpc_agent_sdk.sessions.compact._runtime import AdvancedMemoryRuntime +from ._runtime import AdvancedMemoryRuntime from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import Part diff --git a/trpc_agent_sdk/advanced_memory/_redis_stores.py b/trpc_agent_sdk/advanced_memory/_redis_stores.py index f49479063..07f4f73ae 100644 --- a/trpc_agent_sdk/advanced_memory/_redis_stores.py +++ b/trpc_agent_sdk/advanced_memory/_redis_stores.py @@ -1,5 +1,302 @@ -"""Redis stores owned by long-term Advanced Memory.""" +"""Redis implementations of the Advanced Memory storage contracts.""" -from trpc_agent_sdk.sessions.compact._redis_stores import RedisLongTermMemoryStore +from __future__ import annotations -__all__ = ["RedisLongTermMemoryStore"] +import asyncio +import json +from collections.abc import Mapping +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 trpc_agent_sdk.sessions.compact._formats import MemoryDocument, MemoryIndexEntry +from ._paths import AdvancedMemoryPaths + +_APPEND_UNIQUE_SCRIPT = """ +if redis.call('SADD', KEYS[2], ARGV[1]) == 0 then return 0 end +redis.call('XADD', KEYS[1], '*', 'data', ARGV[2]) +return 1 +""" + +_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}}}" + self._app_base = f"{config.redis_key_prefix}:{{{app_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 _session_base(self, session_id: str) -> str: + safe_session_id = self._paths.session_dir(session_id).name + tenant = f"{self._paths.tenant_root_dir.parent.name}:{self._paths.tenant_root_dir.name}" + return f"{self._config.redis_key_prefix}:{{{tenant}:{safe_session_id}}}" + + def _session_registry(self, session_id: str) -> str: + return f"{self._session_base(session_id)}:keys" + + 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_session_ttl(self, session_id: str, *keys: str) -> None: + skip_prefixes: tuple[str, ...] = () + if not self._config.session_ttl_delete_transcripts: + skip_prefixes = (f"{self._session_base(session_id)}:transcript", ) + await self._refresh_ttl_group( + self._session_registry(session_id), + list(keys), + self._config.session_ttl_seconds, + skip_prefixes=skip_prefixes, + ) + + async def _refresh_memory_ttl(self, *keys: str) -> None: + await self._refresh_ttl_group( + self._memory_registry(), + list(keys), + self._config.memory_ttl_seconds, + ) + + async def delete_session(self, session_id: str) -> None: + """Delete all Advanced Memory keys for one session.""" + session_base = self._session_base(session_id) + registry = self._session_registry(session_id) + keys: set[str] = {registry} + tracked = await self._command("smembers", registry) or [] + keys.update(value for value in (self._text(item) for item in tracked) if value) + + cursor: Any = 0 + pattern = f"{session_base}:*" + while True: + cursor, scanned = await self._command( + "scan", + cursor, + match=pattern, + count=100, + ) + keys.update(value for value in (self._text(item) for item in scanned) if value) + if int(cursor) == 0: + break + if keys: + await self._command("delete", *keys) + + @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() + lines, used_bytes = [], 0 + for line in value.splitlines(keepends=True)[:self._config.memory_index_max_lines]: + size = len(line.encode(self._config.encoding)) + if used_bytes + size > self._config.memory_index_max_bytes: + break + lines.append(line) + used_bytes += size + return "".join(lines) + + 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] + + +class RedisToolResultStore(_RedisStore): + + async def write(self, session_id: str, result_id: str, serialized_result: str) -> Path: + key = f"{self._session_base(session_id)}:tool:{result_id}" + await self._command("set", key, serialized_result) + await self._refresh_session_ttl(session_id, key) + return Path(f"advanced-memory://{key}") + + async def read(self, session_id: str, result_id: str) -> str | None: + key = f"{self._session_base(session_id)}:tool:{result_id}" + value = await self._command("get", key) + await self._refresh_session_ttl(session_id, key) + return self._text(value) + + +class RedisTranscriptStore(_RedisStore): + + @staticmethod + def _validate_record(record: Mapping[str, Any]) -> None: + """Reject Event and Session Memory duplication in Redis.""" + if record.get("kind") in {"event", "session-memory-checkpoint"}: + raise ValueError("Redis transcripts only store context-compression records") + + async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: + self._validate_record(record) + payload = dict(record) + payload.setdefault("recorded_at", datetime.now(timezone.utc).isoformat()) + stream = f"{self._session_base(session_id)}:transcript" + await self._command("xadd", stream, {"data": json.dumps(payload)}) + await self._refresh_session_ttl(session_id, stream) + return Path(f"advanced-memory://{stream}") + + async def append_unique(self, session_id: str, record: Mapping[str, Any], *, unique_key: str) -> tuple[Path, bool]: + self._validate_record(record) + payload = dict(record) + value = payload.get(unique_key) + if not isinstance(value, str) or not 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()) + stream = f"{self._session_base(session_id)}:transcript" + seen = f"{stream}:seen:{unique_key}" + async with self._storage.create_db_session() as connection: + added = await self._storage.execute_command( + connection, + RedisCommand( + method="eval", + args=(_APPEND_UNIQUE_SCRIPT, 2, stream, seen, value, json.dumps(payload, ensure_ascii=False)), + )) + await self._refresh_session_ttl(session_id, stream, seen) + return Path(f"advanced-memory://{stream}"), bool(added) + + async def read_all(self, session_id: str) -> list[dict[str, Any]]: + stream = f"{self._session_base(session_id)}:transcript" + entries = await self._command("xrange", stream, "-", "+") + await self._refresh_session_ttl(session_id, stream) + records: list[dict[str, Any]] = [] + for _, fields in entries: + value = fields.get(b"data") if isinstance(fields, dict) else None + value = value or fields.get("data") + text = self._text(value) + if text: + records.append(json.loads(text)) + return records diff --git a/trpc_agent_sdk/advanced_memory/_runtime.py b/trpc_agent_sdk/advanced_memory/_runtime.py new file mode 100644 index 000000000..31bd28134 --- /dev/null +++ b/trpc_agent_sdk/advanced_memory/_runtime.py @@ -0,0 +1,206 @@ +# 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 trpc_agent_sdk.sessions.compact._coordination import SessionOperationCoordinator +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 + coordination: SessionOperationCoordinator + 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) + _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 + 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 + sql_storage = SqlStorage( + is_async=resolved_config.sql_is_async, + db_url=resolved_config.sql_url, + metadata=AdvancedMemorySqlBase.metadata, + expire_on_commit=False, + ) + else: + local_cleanup = LocalAdvancedMemoryCleanup(resolved_config) + return cls( + config=resolved_config, + paths=paths, + coordination=SessionOperationCoordinator(), + long_term_memory=LongTermMemoryStore(resolved_config, paths), + _redis_storage=redis_storage, + _sql_storage=sql_storage, + _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") + 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_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 + + @property + def coordination(self) -> SessionOperationCoordinator: + """Return the shared coordinator.""" + return self.root.coordination + + def session_key(self, session_id: str) -> str: + """Return a lock/cache key unique across all tenants.""" + return f"{self.scope.storage_key}\0{session_id}" + + async def initialize(self) -> bool: + """Initialize only this tenant's local directories.""" + if not self.config.enabled: + return False + if self.config.storage_backend == "sql" and self.root._sql_cleanup is not None: + await self.root._sql_cleanup.start() + if self.config.storage_backend == "local" and self.root._local_cleanup is not None: + await self.root._local_cleanup.start() + await self.long_term_memory.initialize() + return True diff --git a/trpc_agent_sdk/advanced_memory/_sql_stores.py b/trpc_agent_sdk/advanced_memory/_sql_stores.py index d9173b3ce..11862c982 100644 --- a/trpc_agent_sdk/advanced_memory/_sql_stores.py +++ b/trpc_agent_sdk/advanced_memory/_sql_stores.py @@ -1,5 +1,533 @@ -"""SQL stores owned by long-term Advanced Memory.""" +"""SQL implementations of the Advanced Memory storage contracts.""" -from trpc_agent_sdk.sessions.compact._sql_stores import SqlLongTermMemoryStore +from __future__ import annotations -__all__ = ["SqlLongTermMemoryStore"] +import json +import asyncio +import hashlib +import uuid +from datetime import datetime, timedelta, timezone +from dataclasses import replace +from pathlib import Path +from collections.abc import Mapping +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 trpc_agent_sdk.sessions.compact._formats import MemoryDocument, MemoryIndexEntry +from ._paths import AdvancedMemoryPaths + + +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 SqlTranscript(AdvancedMemorySqlBase): + __tablename__ = "advanced_memory_transcripts" + + 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) + session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + record_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + payload: Mapped[str] = mapped_column(Text) + recorded_at: Mapped[datetime] = mapped_column(PreciseTimestamp, default=func.now()) + expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) + + +class SqlTranscriptSeen(AdvancedMemorySqlBase): + __tablename__ = "advanced_memory_transcript_seen" + + dedupe_id: Mapped[str] = mapped_column(String(64), primary_key=True) + app_name: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) + user_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) + session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) + unique_key: Mapped[str] = mapped_column(String(DEFAULT_MAX_VARCHAR_LENGTH)) + unique_value: Mapped[str] = mapped_column(String(DEFAULT_MAX_VARCHAR_LENGTH)) + expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) + + +class SqlToolResult(AdvancedMemorySqlBase): + __tablename__ = "advanced_memory_tool_results" + + 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) + session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + result_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_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 + + async def _refresh_session_scope(self, db: Any, session_id: str) -> None: + expiry = self._expiry(self._config.session_ttl_seconds) + if expiry is None: + return + tables = ((SqlToolResult, (self._app_name, self._user_id, session_id)), ) + if self._config.session_ttl_delete_transcripts: + tables = ( + (SqlTranscript, (self._app_name, self._user_id, session_id)), + (SqlTranscriptSeen, (self._app_name, self._user_id, session_id)), + *tables, + ) + for model, key in tables: + rows = await self._storage.query( + db, + SqlKey(key=key, storage_cls=model), + SqlCondition(filters=[ + getattr(model, "app_name") == self._app_name, + getattr(model, "user_id") == self._user_id, + getattr(model, "session_id") == session_id, + getattr(model, "expires_at").is_(None) | (getattr(model, "expires_at") > self._now()), + ]), + ) + for row in rows: + row.expires_at = expiry + + async def delete_session(self, session_id: str) -> None: + """Delete all Advanced Memory rows for one session.""" + models = ( + SqlTranscript, + SqlTranscriptSeen, + SqlToolResult, + ) + filters = { + SqlTranscript: [ + SqlTranscript.app_name == self._app_name, + SqlTranscript.user_id == self._user_id, + SqlTranscript.session_id == session_id, + ], + SqlTranscriptSeen: [ + SqlTranscriptSeen.app_name == self._app_name, + SqlTranscriptSeen.user_id == self._user_id, + SqlTranscriptSeen.session_id == session_id, + ], + SqlToolResult: [ + SqlToolResult.app_name == self._app_name, + SqlToolResult.user_id == self._user_id, + SqlToolResult.session_id == session_id, + ], + } + async with self._storage.create_db_session() as db: + for model in models: + await self._storage.delete( + db, + SqlKey(key=tuple(), storage_cls=model), + SqlCondition(filters=filters[model]), + ) + await self._storage.commit(db) + + +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) + await self._storage.commit(db) + content = row.content + lines, used_bytes = [], 0 + for line in content.splitlines(keepends=True)[:self._config.memory_index_max_lines]: + size = len(line.encode(self._config.encoding)) + if used_bytes + size > self._config.memory_index_max_bytes: + break + lines.append(line) + used_bytes += size + return "".join(lines) + + 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 SqlToolResultStore(_SqlStore): + + async def write(self, session_id: str, result_id: str, serialized_result: str) -> Path: + async with self._storage.create_db_session() as db: + key = (self._app_name, self._user_id, session_id, result_id) + row = await self._storage.get(db, SqlKey(key=key, storage_cls=SqlToolResult)) + if row is None: + row = SqlToolResult( + app_name=key[0], + user_id=key[1], + session_id=key[2], + result_id=key[3], + ) + await self._storage.add(db, row) + row.content = serialized_result + row.updated_at = self._now() + row.expires_at = self._expiry(self._config.session_ttl_seconds) + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/tool/{result_id}") + + async def read(self, session_id: str, result_id: str) -> str | None: + async with self._storage.create_db_session() as db: + row = await self._storage.get( + db, + SqlKey(key=(self._app_name, self._user_id, session_id, result_id), storage_cls=SqlToolResult), + ) + if row is None or self._expired(row.expires_at): + return None + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return row.content + + +class SqlTranscriptStore(_SqlStore): + + @staticmethod + def _validate_record(record: Mapping[str, Any]) -> None: + """Reject Event and Session Memory duplication in SQL.""" + if record.get("kind") in {"event", "session-memory-checkpoint"}: + raise ValueError("SQL transcripts only store context-compression records") + + def _dedupe_id(self, session_id: str, unique_key: str, value: str) -> str: + raw = "\0".join((self._app_name, self._user_id, session_id, unique_key, value)) + return hashlib.sha256(raw.encode("utf-8")).hexdigest() + + async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: + self._validate_record(record) + payload = dict(record) + payload.setdefault("recorded_at", self._now().replace(tzinfo=timezone.utc).isoformat()) + async with self._storage.create_db_session() as db: + await self._storage.add( + db, + SqlTranscript( + app_name=self._app_name, + user_id=self._user_id, + session_id=session_id, + record_id=uuid.uuid4().hex, + payload=json.dumps(payload, ensure_ascii=False), + expires_at=(self._expiry(self._config.session_ttl_seconds) + if self._config.session_ttl_delete_transcripts else None), + )) + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript") + + async def append_unique( + self, + session_id: str, + record: Mapping[str, Any], + *, + unique_key: str, + ) -> tuple[Path, bool]: + self._validate_record(record) + payload = dict(record) + value = payload.get(unique_key) + if not isinstance(value, str) or not value: + raise ValueError(f"Transcript unique key {unique_key!r} must be a non-empty string") + async with self._storage.create_db_session() as db: + dedupe_id = self._dedupe_id(session_id, unique_key, value) + seen_key = (self._app_name, self._user_id, session_id, unique_key, value) + seen = await self._storage.get( + db, + SqlKey(key=(dedupe_id, ), storage_cls=SqlTranscriptSeen), + ) + if seen is not None and not self._expired(seen.expires_at): + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript"), False + if seen is not None: + await self._storage.delete( + db, + SqlKey(key=(dedupe_id, ), storage_cls=SqlTranscriptSeen), + SqlCondition(filters=[ + SqlTranscriptSeen.dedupe_id == dedupe_id, + ]), + ) + payload.setdefault("recorded_at", self._now().replace(tzinfo=timezone.utc).isoformat()) + await self._storage.add( + db, + SqlTranscriptSeen( + dedupe_id=dedupe_id, + app_name=seen_key[0], + user_id=seen_key[1], + session_id=seen_key[2], + unique_key=seen_key[3], + unique_value=seen_key[4], + expires_at=(self._expiry(self._config.session_ttl_seconds) + if self._config.session_ttl_delete_transcripts else None), + )) + await self._storage.add( + db, + SqlTranscript( + app_name=self._app_name, + user_id=self._user_id, + session_id=session_id, + record_id=uuid.uuid4().hex, + payload=json.dumps(payload, ensure_ascii=False), + expires_at=(self._expiry(self._config.session_ttl_seconds) + if self._config.session_ttl_delete_transcripts else None), + )) + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript"), True + + async def read_all(self, session_id: str) -> list[dict[str, Any]]: + async with self._storage.create_db_session() as db: + rows = await self._storage.query( + db, + SqlKey(key=(self._app_name, self._user_id, session_id), storage_cls=SqlTranscript), + SqlCondition( + filters=[ + SqlTranscript.app_name == self._app_name, + SqlTranscript.user_id == self._user_id, + SqlTranscript.session_id == session_id, + SqlTranscript.expires_at.is_(None) | (SqlTranscript.expires_at > self._now()), + ], + order_func=SqlTranscript.recorded_at.asc, + ), + ) + await self._refresh_session_scope(db, session_id) + await self._storage.commit(db) + return [json.loads(row.payload) for row in rows] + + +class SqlAdvancedMemoryCleanup: + """Periodically remove expired Advanced Memory SQL rows.""" + + _models = ( + SqlMemoryIndex, + SqlMemoryTopic, + SqlTranscript, + SqlTranscriptSeen, + SqlToolResult, + ) + + 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 + and self._config.session_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: + models = self._models if self._config.session_ttl_delete_transcripts else tuple( + model for model in self._models if model is not SqlTranscript) + for model in models: + await self._storage.delete( + db, + SqlKey(key=tuple(), storage_cls=model), + SqlCondition(filters=[model.expires_at.is_not(None), model.expires_at <= now]), + ) + 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", + "SqlToolResultStore", + "SqlTranscriptStore", +] diff --git a/trpc_agent_sdk/advanced_memory/_storage.py b/trpc_agent_sdk/advanced_memory/_storage.py index 0e0a957ae..d4d17daea 100644 --- a/trpc_agent_sdk/advanced_memory/_storage.py +++ b/trpc_agent_sdk/advanced_memory/_storage.py @@ -1,5 +1,189 @@ -"""Local storage owned by long-term Advanced Memory.""" +# 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 trpc_agent_sdk.sessions.compact._storage import LongTermMemoryStore +from __future__ import annotations -__all__ = ["LongTermMemoryStore"] +import asyncio +import os +import tempfile +import time +from dataclasses import replace +from datetime import datetime +from datetime import timezone +from pathlib import Path + +from trpc_agent_sdk.sessions.compact._formats import MemoryDocument +from trpc_agent_sdk.sessions.compact._formats import MemoryIndexEntry + +from ._config import AdvancedMemoryServiceConfig +from ._paths import AdvancedMemoryPaths + + +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 "" + lines: list[str] = [] + used_bytes = 0 + with self.index_path.open(encoding=self._config.encoding) as source: + for _ in range(self._config.memory_index_max_lines): + line = source.readline() + if not line: + break + size = len(line.encode(self._config.encoding)) + if used_bytes + size > self._config.memory_index_max_bytes: + break + lines.append(line) + used_bytes += size + return "".join(lines) + + 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/advanced_memory/_storage_backend.py b/trpc_agent_sdk/advanced_memory/_storage_backend.py index 19a5b2729..4e10a2f0c 100644 --- a/trpc_agent_sdk/advanced_memory/_storage_backend.py +++ b/trpc_agent_sdk/advanced_memory/_storage_backend.py @@ -8,8 +8,8 @@ from typing import Protocol -from trpc_agent_sdk.sessions.compact._paths import MemoryScope -from trpc_agent_sdk.sessions.compact._runtime import ScopedAdvancedMemoryRuntime +from ._paths import MemoryScope +from ._runtime import ScopedAdvancedMemoryRuntime class AdvancedMemoryStorageBackend(Protocol): diff --git a/trpc_agent_sdk/memory/__init__.py b/trpc_agent_sdk/memory/__init__.py index d93e9dabc..a4b768e07 100644 --- a/trpc_agent_sdk/memory/__init__.py +++ b/trpc_agent_sdk/memory/__init__.py @@ -27,7 +27,7 @@ __all__ = [ "BaseMemoryService", "MemoryServiceConfig", - "AdvancedCompactConfig", + "AdvancedMemoryServiceConfig", "AdvancedMemoryService", "EventTtl", "InMemoryMemoryService", @@ -43,8 +43,8 @@ def __getattr__(name: str): """Lazily expose Advanced Memory configuration without import cycles.""" - if name == "AdvancedCompactConfig": - from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig + if name == "AdvancedMemoryServiceConfig": + from trpc_agent_sdk.advanced_memory import AdvancedMemoryServiceConfig - return AdvancedCompactConfig + return AdvancedMemoryServiceConfig 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 dc00d31c6..c626256cd 100644 --- a/trpc_agent_sdk/memory/_advanced_memory_service.py +++ b/trpc_agent_sdk/memory/_advanced_memory_service.py @@ -19,7 +19,7 @@ from trpc_agent_sdk.sessions import Session if TYPE_CHECKING: - from trpc_agent_sdk.advanced_memory import AdvancedCompactConfig + from trpc_agent_sdk.advanced_memory import AdvancedMemoryServiceConfig from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime from trpc_agent_sdk.advanced_memory import LongTermMemoryIntegration @@ -28,24 +28,24 @@ class AdvancedMemoryService(BaseMemoryService): """Expose user-scoped long-term Memory through the Runner memory API. ``Runner`` calls :meth:`bind` automatically. Session compression is - configured independently with ``setup_context_compression``. + configured independently through ``SessionService.session_compact_manager``. """ def __init__( self, - config: AdvancedCompactConfig | None = None, + config: AdvancedMemoryServiceConfig | None = None, *, runtime: AdvancedMemoryRuntime | 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 AdvancedCompactConfig + from trpc_agent_sdk.advanced_memory import AdvancedMemoryServiceConfig from trpc_agent_sdk.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 AdvancedCompactConfig()) + 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._preload_memory_model = preload_memory_model @@ -54,7 +54,7 @@ def __init__( self._bound_agent: Any | None = None @property - def config(self) -> AdvancedCompactConfig: + def config(self) -> AdvancedMemoryServiceConfig: """Return the Advanced Memory configuration.""" return self._runtime.config diff --git a/trpc_agent_sdk/runners.py b/trpc_agent_sdk/runners.py index 083c36803..418d09519 100644 --- a/trpc_agent_sdk/runners.py +++ b/trpc_agent_sdk/runners.py @@ -230,10 +230,9 @@ def __init__( if isinstance(memory_service, AdvancedMemoryService): session_service = memory_service.bind(agent, session_service) - compact_config = getattr(session_service, "session_compact_config", None) - from trpc_agent_sdk.sessions.compact import BaseSessionCompactConfig - if isinstance(compact_config, BaseSessionCompactConfig): - compact_config.setup(agent, session_service) + 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 5501aed23..9ce67dc43 100644 --- a/trpc_agent_sdk/sessions/__init__.py +++ b/trpc_agent_sdk/sessions/__init__.py @@ -54,12 +54,9 @@ "State", "BaseSessionService", "BaseSessionCompactManager", - "BaseSessionCompactConfig", "AdvancedCompactConfig", "AdvancedSessionCompactManager", "AutoCompact", - "setup_advanced_session_compact", - "setup_context_compression", "HistoryRecord", "InMemorySessionService", "SessionWithTTL", @@ -99,13 +96,10 @@ def __getattr__(name: str): """Lazily expose Advanced Memory without creating an import cycle.""" if name in { - "AdvancedCompactConfig", - "AdvancedSessionCompactManager", - "AutoCompact", - "BaseSessionCompactManager", - "BaseSessionCompactConfig", - "setup_advanced_session_compact", - "setup_context_compression", + "AdvancedCompactConfig", + "AdvancedSessionCompactManager", + "AutoCompact", + "BaseSessionCompactManager", }: from . import compact diff --git a/trpc_agent_sdk/sessions/_base_session_service.py b/trpc_agent_sdk/sessions/_base_session_service.py index 8827b2f2f..6cbd8fbed 100644 --- a/trpc_agent_sdk/sessions/_base_session_service.py +++ b/trpc_agent_sdk/sessions/_base_session_service.py @@ -39,7 +39,6 @@ if TYPE_CHECKING: from .compact import BaseSessionCompactManager - from .compact import BaseSessionCompactConfig class BaseSessionService(SessionServiceABC): @@ -51,23 +50,15 @@ class BaseSessionService(SessionServiceABC): def __init__(self, summarizer_manager: Optional[SummarizerSessionManager] = None, session_config: Optional[SessionServiceConfig] = None, - session_compact_config: Optional["BaseSessionCompactConfig"] = 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_config: Optional Advanced Compact configuration session_compact_manager: Optional pluggable Session Compact manager """ - if session_compact_config is not None and session_compact_manager is not None: - raise ValueError( - "Provide either session_compact_config or " - "session_compact_manager, not both" - ) self._summarizer_manager = summarizer_manager - self._session_compact_config = session_compact_config self._session_compact_manager: Optional[BaseSessionCompactManager] = None if session_config is None: session_config = SessionServiceConfig() @@ -89,11 +80,6 @@ def session_config(self) -> SessionServiceConfig: """Get the session service configuration.""" return self._session_config - @property - def session_compact_config(self) -> Optional["BaseSessionCompactConfig"]: - """Return deferred Session Compact configuration, if configured.""" - return self._session_compact_config - @property def session_compact_manager(self) -> Optional["BaseSessionCompactManager"]: """Get the Session Compact lifecycle manager.""" @@ -107,9 +93,7 @@ def set_summarizer_manager(self, summarizer_manager: SummarizerSessionManager, f 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" - ) + 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) @@ -121,9 +105,7 @@ def set_session_compact_manager( ) -> None: """Attach Session Compact through the native manager lifecycle.""" if self._summarizer_manager is not None: - raise ValueError( - "SummarizerSessionManager and BaseSessionCompactManager are mutually exclusive" - ) + 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 @@ -244,21 +226,6 @@ async def get_session_summary(self, session: Session) -> Optional[str]: return await self._session_compact_manager.get_session_summary(session) return None - async def _delete_session_compact_data( - self, - *, - app_name: str, - user_id: str, - session_id: str, - ) -> None: - """Delete side data owned by the configured compact manager.""" - if self._session_compact_manager: - await self._session_compact_manager.delete_session( - app_name=app_name, - user_id=user_id, - session_id=session_id, - ) - def filter_events(self, session: Session, need_copy: bool = False) -> Session: """Filter events based on the session config. diff --git a/trpc_agent_sdk/sessions/_in_memory_session_service.py b/trpc_agent_sdk/sessions/_in_memory_session_service.py index 642e06135..fdd1d1ce3 100644 --- a/trpc_agent_sdk/sessions/_in_memory_session_service.py +++ b/trpc_agent_sdk/sessions/_in_memory_session_service.py @@ -54,7 +54,6 @@ if TYPE_CHECKING: from .compact._base_manager import BaseSessionCompactManager - from .compact._base_config import BaseSessionCompactConfig class SessionWithTTL(BaseModel): @@ -114,12 +113,10 @@ class InMemorySessionService(BaseSessionService): def __init__(self, summarizer_manager: Optional[SummarizerSessionManager] = None, session_config: Optional[SessionServiceConfig] = None, - session_compact_config: "BaseSessionCompactConfig | None" = None, session_compact_manager: BaseSessionCompactManager | None = None): super().__init__( summarizer_manager=summarizer_manager, session_config=session_config, - session_compact_config=session_compact_config, session_compact_manager=session_compact_manager, ) # Storage with TTL support @@ -227,11 +224,6 @@ async def list_sessions(self, *, app_name: str, user_id: Optional[str] = None) - async def delete_session(self, *, app_name: str, user_id: str, session_id: str) -> None: 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] - await self._delete_session_compact_data( - app_name=app_name, - user_id=user_id, - session_id=session_id, - ) @override async def append_event(self, session: Session, event: Event) -> Event: diff --git a/trpc_agent_sdk/sessions/_redis_session_service.py b/trpc_agent_sdk/sessions/_redis_session_service.py index 650c7188c..2fce605a8 100644 --- a/trpc_agent_sdk/sessions/_redis_session_service.py +++ b/trpc_agent_sdk/sessions/_redis_session_service.py @@ -39,7 +39,6 @@ if TYPE_CHECKING: from .compact._base_manager import BaseSessionCompactManager - from .compact._base_config import BaseSessionCompactConfig def _session_key_prefix(app_name: str, user_id: Optional[str] = None) -> str: @@ -94,7 +93,6 @@ def __init__(self, summarizer_manager: Optional[SummarizerSessionManager] = None, session_config: Optional[SessionServiceConfig] = None, is_async: bool = False, - session_compact_config: "BaseSessionCompactConfig | None" = None, session_compact_manager: BaseSessionCompactManager | None = None, **kwargs: Any): self._db_url = db_url @@ -103,7 +101,6 @@ def __init__(self, super().__init__( summarizer_manager=summarizer_manager, session_config=session_config, - session_compact_config=session_compact_config, session_compact_manager=session_compact_manager, ) if is_default_config: @@ -220,11 +217,6 @@ async def delete_session(self, *, app_name: str, user_id: str, session_id: str) async with self._redis_storage.create_db_session() as redis_session: key = session_key(app_name, user_id, session_id) await self._redis_storage.delete(redis_session, key) - await self._delete_session_compact_data( - app_name=app_name, - user_id=user_id, - session_id=session_id, - ) @override async def append_event(self, session: Session, event: Event) -> Event: diff --git a/trpc_agent_sdk/sessions/_sql_session_service.py b/trpc_agent_sdk/sessions/_sql_session_service.py index 5cfb3e9f8..1f998ba9b 100644 --- a/trpc_agent_sdk/sessions/_sql_session_service.py +++ b/trpc_agent_sdk/sessions/_sql_session_service.py @@ -80,7 +80,6 @@ if TYPE_CHECKING: from .compact._base_manager import BaseSessionCompactManager - from .compact._base_config import BaseSessionCompactConfig def _event_field_or_default(field_name: str, value: Any) -> Any: @@ -396,7 +395,6 @@ def __init__(self, summarizer_manager: Optional[SummarizerSessionManager] = None, is_async: bool = False, session_config: Optional[SessionServiceConfig] = None, - session_compact_config: "BaseSessionCompactConfig | None" = None, session_compact_manager: BaseSessionCompactManager | None = None, **kwargs: Any): self._db_url = db_url @@ -405,7 +403,6 @@ def __init__(self, super().__init__( summarizer_manager=summarizer_manager, session_config=session_config, - session_compact_config=session_compact_config, session_compact_manager=session_compact_manager, ) if is_default_config: @@ -557,11 +554,6 @@ async def delete_session(self, app_name: str, user_id: str, session_id: str) -> session_key = SqlKey(key=(app_name, user_id, session_id), storage_cls=StorageSession) await self._sql_storage.delete(sql_session, session_key, conditions) await self._sql_storage.commit(sql_session) - await self._delete_session_compact_data( - app_name=app_name, - user_id=user_id, - session_id=session_id, - ) @override async def append_event(self, session: Session, event: Event) -> Event: diff --git a/trpc_agent_sdk/sessions/compact/__init__.py b/trpc_agent_sdk/sessions/compact/__init__.py index 4551f8745..1fe406010 100644 --- a/trpc_agent_sdk/sessions/compact/__init__.py +++ b/trpc_agent_sdk/sessions/compact/__init__.py @@ -12,7 +12,6 @@ from ._autocompact import ForkedLegacySummaryGenerator from ._autocompact import setup_autocompact from ._base_manager import BaseSessionCompactManager -from ._base_config import BaseSessionCompactConfig from ._config import AdvancedCompactConfig from ._formats import build_session_memory_state from ._formats import parse_session_memory_state @@ -25,17 +24,11 @@ from ._history_snip import HistorySnipCallback from ._history_snip import HistorySnipResult from ._history_snip import setup_history_snip -from ._integration import setup_advanced_session_compact -from ._integration import setup_context_compression 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 ._paths import MemoryScope -from ._runtime import AdvancedMemoryRuntime -from ._runtime import ScopedAdvancedMemoryRuntime from ._session_memory import build_session_memory_prompt from ._session_memory import ForkedSessionMemoryGenerator from ._session_memory import has_session_memory_content @@ -43,10 +36,6 @@ from ._session_memory import SessionMemoryExtractionInput from ._session_memory import SessionMemoryExtractionResult from ._session_memory import SessionMemoryExtractor -from ._session_service import TranscriptSessionService -from ._storage import SessionMemoryStore -from ._storage import ToolResultStore -from ._storage import TranscriptStore from ._token_budget import ContextBudget from ._token_budget import ContextTokenEstimate from ._token_budget import HeuristicTokenEstimator @@ -57,13 +46,11 @@ 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 ._runtime import ScopedSessionCompactRuntime +from ._runtime import SessionCompactRuntime __all__ = [ "AdvancedCompactConfig", - "BaseSessionCompactConfig", - "AdvancedMemoryPaths", - "AdvancedMemoryRuntime", "AutoCompact", "AutoCompactCallback", "AutoCompactResult", @@ -75,12 +62,10 @@ "HistorySnip", "HistorySnipCallback", "HistorySnipResult", - "MemoryScope", "Microcompact", "MicrocompactCallback", "MicrocompactResult", "ModelContextWindowResolver", - "ScopedAdvancedMemoryRuntime", "SESSION_MEMORY_SECTION_DESCRIPTIONS", "SESSION_MEMORY_SECTIONS", "SESSION_MEMORY_STATE_KEY", @@ -88,7 +73,6 @@ "SessionMemoryExtractionInput", "SessionMemoryExtractionResult", "SessionMemoryExtractor", - "SessionMemoryStore", "BaseSessionCompactManager", "AdvancedSessionCompactManager", "TokenContextTracker", @@ -96,10 +80,8 @@ "ToolResultBudget", "ToolResultBudgetCallback", "ToolResultBudgetResult", - "ToolResultStore", - "TRANSCRIPT_SCHEMA_VERSION", - "TranscriptSessionService", - "TranscriptStore", + "SessionCompactRuntime", + "ScopedSessionCompactRuntime", "build_session_memory_prompt", "build_session_memory_state", "content_signature", @@ -108,8 +90,6 @@ "limit_session_memory_document", "parse_session_memory_state", "setup_autocompact", - "setup_advanced_session_compact", - "setup_context_compression", "setup_history_snip", "setup_microcompact", "setup_tool_result_budget", diff --git a/trpc_agent_sdk/sessions/compact/_autocompact.py b/trpc_agent_sdk/sessions/compact/_autocompact.py index 045d3052d..af4630b5a 100644 --- a/trpc_agent_sdk/sessions/compact/_autocompact.py +++ b/trpc_agent_sdk/sessions/compact/_autocompact.py @@ -32,7 +32,7 @@ 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: @@ -41,14 +41,13 @@ 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) @@ -229,7 +228,7 @@ class AutoCompact: def __init__( self, - memory_runtime: AdvancedMemoryRuntime, + memory_runtime: SessionCompactRuntime, summary_generator: LegacySummaryGenerator | None = None, *, model: Any | None = None, @@ -246,7 +245,7 @@ def __init__( 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 @@ -269,36 +268,12 @@ def _session_lock(self, session_id: str) -> asyncio.Lock: return lock async def _load_state(self, session_id: str) -> AutoCompactState: - """Restore the latest compaction and failure count from the transcript.""" + """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, - record.get("boundary_event_id") - if isinstance(record.get("boundary_event_id"), str) else None, - record.get("compaction_id") - if isinstance(record.get("compaction_id"), str) else None, - ) - failures = 0 - elif record.get("kind") == "autocompact-failure": - failures += 1 - state = AutoCompactState(latest_compaction=latest, consecutive_failures=failures) + state = AutoCompactState(latest_compaction=None, consecutive_failures=0) self._states[state_key] = state return state @@ -310,17 +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.""" - if self._runtime.config.storage_backend in {"redis", "sql"}: - return (f"{summary.rstrip()}\n\n" - "For exact content from before compaction, read the original " - "SessionService Events. Current session memory is stored in " - f"session.state[{SESSION_MEMORY_STATE_KEY!r}].") + """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.storage_reference('transcript', session_id=session_id)}\n" - "Current session memory: " - f"{self._runtime.paths.storage_reference('session_memory', session_id=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, @@ -415,65 +384,22 @@ async def _latest_session_memory_record( session_id: str, ctx: "InvocationContext", ) -> tuple[str, str, int, str] | None: - """Read session memory and its checkpoint Event for model-free compaction.""" - if self._runtime.config.storage_backend in {"redis", "sql"}: - parsed = parse_session_memory_state(ctx.session.state.get(SESSION_MEMORY_STATE_KEY)) - if parsed is None: - return None - document, checkpoint, _ = parsed - signature = checkpoint.get("boundary_signature") - occurrence = checkpoint.get("boundary_occurrence") - event_id = checkpoint.get("last_event_id") - if (not isinstance(signature, str) or not isinstance(occurrence, int) or occurrence <= 0 - or not isinstance(event_id, str)): - return None - memory = document.to_markdown() - if memory.strip() == SessionMemoryDocument().to_markdown().strip(): - return None - return memory, signature, occurrence, event_id - async with self._runtime.coordination.guard( - session_id, - timeout=self._runtime.config.session_memory_wait_timeout_seconds, - ) as acquired: - if not acquired: - return None - if self._runtime.session_memory is None: - 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): - boundary = self._event_content_signature( - records, - record["last_event_id"], - ) - if boundary is not None: - return memory, boundary[0], boundary[1], 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 + """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, @@ -525,9 +451,7 @@ def _resolve_boundary_event_id( 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 + event for event in (getattr(ctx.session, "events", []) or []) if getattr(event, "content", None) is not None ] if len(content_events) <= 1: return None @@ -552,8 +476,8 @@ async def _persist_session_compaction( 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. setup_context_compression always - # supplies the framework Session and persists the compacted window. + # 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( @@ -626,67 +550,6 @@ 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.""" - compaction_id = record.compaction_id or f"autocompact:{uuid.uuid4().hex}" - await self._runtime.transcripts.append( - session_id, - { - "schema_version": AUTOCOMPACT_SCHEMA_VERSION, - "kind": "autocompact-success", - "compaction_id": compaction_id, - "boundary_signature": record.boundary_signature, - "boundary_occurrence": record.boundary_occurrence, - "boundary_event_id": record.boundary_event_id, - "summary": record.summary, - "source": record.source, - "request_chars_before": before_chars, - "request_chars_after": after_chars, - "request_tokens_before": before_tokens, - "request_tokens_after": after_tokens, - "token_source": token_source, - }, - ) - - async def _persist_failure( - self, - session_id: str, - error: Exception, - failures: int, - token_budget: Any | None = None, - ) -> None: - """Persist failures so the circuit breaker survives a restart.""" - await self._runtime.transcripts.append( - session_id, - { - "schema_version": - AUTOCOMPACT_SCHEMA_VERSION, - "kind": - "autocompact-failure", - "attempt_id": - f"autocompact:{uuid.uuid4().hex}", - "consecutive_failures": - failures, - "error": - str(error), - "request_tokens": (token_budget.estimate.tokens - if token_budget is not None and token_budget.token_mode_enabled else None), - "context_window_tokens": (token_budget.context_window_tokens - if token_budget is not None and token_budget.token_mode_enabled else None), - "token_source": (token_budget.estimate.source - if token_budget is not None and token_budget.token_mode_enabled else None), - }, - ) - async def apply( self, request: "LlmRequest", @@ -722,7 +585,6 @@ async def _apply_scoped( 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) @@ -734,9 +596,7 @@ async def _apply_scoped( 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 - ) + 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: @@ -775,7 +635,7 @@ async def _apply_scoped( original_contents = [content.model_copy(deep=True) for content in request.contents] try: compact_record: AutoCompactRecord | None = None - if (self._session_memory_extractor is not None and self._session_memory_extractor.uses_session_state): + if self._session_memory_extractor is not None: await self._session_memory_extractor.extract_if_needed( ctx.session, ctx, @@ -809,9 +669,11 @@ async def _apply_scoped( strict_boundary=True, boundary_event_id=boundary_event_id, ) - 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 token_mode: + target_reached = (tracker.budget(request, ctx).estimate.tokens + <= 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 @@ -841,7 +703,6 @@ async def _apply_scoped( ) request_chars_after = estimate_request_chars(request) - token_budget_after = tracker.budget(request, ctx) if token_mode: comparison_tokens_after = tracker.estimate_request_tokens(request) if (comparison_tokens_after >= comparison_tokens_before @@ -850,15 +711,6 @@ async def _apply_scoped( elif request_chars_after >= request_chars_before: raise ValueError("Autocompact did not reduce request size") await self._persist_session_compaction(ctx, compact_record) - await self._persist_success( - session_id, - compact_record, - request_chars_before, - request_chars_after, - comparison_tokens_before if token_mode else None, - comparison_tokens_after if token_mode else None, - "estimated" if token_mode else None, - ) state.latest_compaction = compact_record state.consecutive_failures = 0 return AutoCompactResult( @@ -876,12 +728,6 @@ async def _apply_scoped( 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, @@ -937,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_config.py b/trpc_agent_sdk/sessions/compact/_base_config.py deleted file mode 100644 index 71a90ca08..000000000 --- a/trpc_agent_sdk/sessions/compact/_base_config.py +++ /dev/null @@ -1,29 +0,0 @@ -# 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 configuration contract for Session Compact strategies.""" - -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 ._base_manager import BaseSessionCompactManager - - -class BaseSessionCompactConfig(ABC): - """Create and attach one concrete Session Compact strategy.""" - - @abstractmethod - def setup( - self, - agent: Any, - session_service: Any, - ) -> "BaseSessionCompactManager": - """Create the strategy manager and attach it to the SessionService.""" diff --git a/trpc_agent_sdk/sessions/compact/_base_manager.py b/trpc_agent_sdk/sessions/compact/_base_manager.py index f7a38bb9a..2a861ec8e 100644 --- a/trpc_agent_sdk/sessions/compact/_base_manager.py +++ b/trpc_agent_sdk/sessions/compact/_base_manager.py @@ -10,6 +10,7 @@ from abc import ABC from abc import abstractmethod +from typing import Any from typing import TYPE_CHECKING if TYPE_CHECKING: @@ -21,6 +22,10 @@ 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, @@ -42,16 +47,6 @@ async def create_session_summary( async def get_session_summary(self, session: "Session") -> str | None: """Return the compact representation exposed as a session summary.""" - @abstractmethod - async def delete_session( - self, - *, - app_name: str, - user_id: str, - session_id: str, - ) -> None: - """Delete side data owned by this manager for one session.""" - @abstractmethod async def close(self) -> None: """Release resources owned by this manager.""" diff --git a/trpc_agent_sdk/sessions/compact/_callbacks.py b/trpc_agent_sdk/sessions/compact/_callbacks.py index 99b422346..a678fffb1 100644 --- a/trpc_agent_sdk/sessions/compact/_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 index 456ff3dcd..1692bf2d6 100644 --- a/trpc_agent_sdk/sessions/compact/_config.py +++ b/trpc_agent_sdk/sessions/compact/_config.py @@ -1,129 +1,47 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# Tencent is pleased to support the open source ecosystem. # # Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Configuration for the independent Advanced Memory mechanism.""" +# Licensed under Apache-2.0. +"""Configuration for Session Compact.""" from __future__ import annotations -import os from dataclasses import dataclass from dataclasses import field -from pathlib import Path from typing import Any -from typing import Literal - -from ._base_config import BaseSessionCompactConfig DEFAULT_COMPACTABLE_TOOL_NAMES = ( "Read", "Bash", "Grep", "Glob", - "WebSearch", - "WebFetch", - "Edit", - "Write", + "Search", + "CodeSearch", ) -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}") + raise ValueError(f"{name} must be non-negative") 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 AdvancedCompactConfig(BaseSessionCompactConfig): - """Configure Advanced Session Compact and its shared memory runtime.""" +class AdvancedCompactConfig: + """Configure compression that is persisted by the SessionService.""" 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 - session_ttl_seconds: int | None = None - 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" - 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 - memory_focus_instruction: str | None = None tool_result_max_chars: int = 50_000 tool_results_per_message_max_chars: int = 200_000 tool_result_preview_chars: int = 2_000 @@ -132,16 +50,8 @@ class AdvancedCompactConfig(BaseSessionCompactConfig): 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, - )) + 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 @@ -171,133 +81,50 @@ class AdvancedCompactConfig(BaseSessionCompactConfig): 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 - session_ttl_delete_transcripts: bool = False - - def setup(self, agent: Any, session_service: Any) -> Any: - """Create and attach the Advanced Session Compact manager.""" - from ._integration import setup_advanced_session_compact - - return setup_advanced_session_compact( - agent, - session_service, - self, - ) 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.session_ttl_seconds is not None and self.session_ttl_seconds <= 0: - raise ValueError("session_ttl_seconds must be greater than zero when provided") - 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.session_dir_name, - self.memory_index_name, - self.transcript_name, - self.session_memory_name, - )) + """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, - ) - _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( + 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, - ) - _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_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, - ) - _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_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) - object.__setattr__(self, "root_dir", self.root_dir.expanduser().resolve()) diff --git a/trpc_agent_sdk/sessions/compact/_formats.py b/trpc_agent_sdk/sessions/compact/_formats.py index ece6c8f28..33c3841b6 100644 --- a/trpc_agent_sdk/sessions/compact/_formats.py +++ b/trpc_agent_sdk/sessions/compact/_formats.py @@ -5,6 +5,8 @@ # 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 diff --git a/trpc_agent_sdk/sessions/compact/_history_snip.py b/trpc_agent_sdk/sessions/compact/_history_snip.py index 72f959320..5f8867551 100644 --- a/trpc_agent_sdk/sessions/compact/_history_snip.py +++ b/trpc_agent_sdk/sessions/compact/_history_snip.py @@ -15,7 +15,7 @@ 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 @@ -27,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]" @@ -84,7 +83,7 @@ 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] = {} @@ -92,7 +91,7 @@ def __init__(self, memory_runtime: AdvancedMemoryRuntime) -> None: 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 @@ -106,25 +105,14 @@ def _session_lock(self, session_id: str) -> asyncio.Lock: return lock async def _load_state(self, session_id: str) -> HistorySnipState: - """Restore prior history-snip decisions from the transcript.""" + """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[state_key] = state return state @@ -162,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", @@ -199,7 +164,6 @@ async def apply( request_chars = estimate_request_chars(request) return HistorySnipResult(None, 0, 0, 0, request_chars, request_chars) if ctx is None or hasattr(self._runtime, "scope"): - await self._runtime.initialize() 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) @@ -271,7 +235,6 @@ async def _apply_scoped( 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 @@ -317,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/_integration.py b/trpc_agent_sdk/sessions/compact/_integration.py deleted file mode 100644 index f83da03bb..000000000 --- a/trpc_agent_sdk/sessions/compact/_integration.py +++ /dev/null @@ -1,155 +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 setup entry points for the context-compression pipeline.""" - -from __future__ import annotations - -from dataclasses import replace -from typing import Any -from typing import TYPE_CHECKING - -from ._autocompact import LegacySummaryGenerator -from ._autocompact import setup_autocompact -from ._history_snip import setup_history_snip -from ._microcompact import setup_microcompact -from ._runtime import AdvancedMemoryRuntime -from ._config import AdvancedCompactConfig -from ._manager import AdvancedSessionCompactManager -from ._session_memory import SessionMemoryExtractor -from ._session_memory import SessionMemoryGenerator -from ._tool_result_budget import setup_tool_result_budget - -if TYPE_CHECKING: - from trpc_agent_sdk.agents import LlmAgent - from trpc_agent_sdk.sessions import SessionServiceABC - - -def setup_context_compression( - agent: "LlmAgent", - session_service: "SessionServiceABC", - memory_runtime: AdvancedMemoryRuntime, - summary_generator: LegacySummaryGenerator | None = None, - *, - compact_model: Any | None = None, - session_memory_generator: SessionMemoryGenerator | None = None, - session_memory_model: Any | None = None, -) -> "SessionServiceABC": - """Install native Session compression on an existing SessionService. - - The original service remains responsible for persistence. Session Compact - is attached through the BaseSessionService manager lifecycle. - """ - session_config = getattr(session_service, "session_config", None) - if session_config is None or not getattr(session_config, "store_historical_events", False): - raise ValueError( - "Context compression requires " - "SessionServiceConfig(store_historical_events=True)" - ) - if getattr(session_service, "summarizer_manager", None) is not None: - raise ValueError( - "Context compression and SummarizerSessionManager are mutually exclusive" - ) - - manager = getattr(session_service, "session_compact_manager", None) - if manager is not None: - if not isinstance(manager, AdvancedSessionCompactManager): - raise ValueError( - "Advanced context compression requires an " - "AdvancedSessionCompactManager" - ) - if manager.runtime is not memory_runtime: - raise ValueError("Context compression session service uses another runtime") - extractor = manager.session_memory_extractor - 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: - attach_manager = getattr(session_service, "set_session_compact_manager", None) - if not callable(attach_manager): - raise TypeError( - "Context compression requires a BaseSessionService with " - "set_session_compact_manager()" - ) - extractor = SessionMemoryExtractor( - memory_runtime, - session_memory_generator, - model=session_memory_model, - ) - manager = AdvancedSessionCompactManager( - memory_runtime, - extractor, - ) - attach_manager(manager) - setup_tool_result_budget(agent, memory_runtime) - setup_history_snip(agent, memory_runtime) - setup_microcompact(agent, memory_runtime) - autocompact = setup_autocompact( - agent, - memory_runtime, - summary_generator, - model=compact_model, - ) - autocompact.attach_session_memory_extractor(extractor) - return session_service - - -def setup_advanced_session_compact( - agent: Any, - session_service: "SessionServiceABC", - compact_config: AdvancedCompactConfig, - *, - summary_generator: LegacySummaryGenerator | None = None, - compact_model: Any | None = None, - session_memory_generator: SessionMemoryGenerator | None = None, - session_memory_model: Any | None = None, -) -> AdvancedSessionCompactManager: - """Configure Advanced Compact from a standard SessionService backend.""" - from trpc_agent_sdk.sessions import InMemorySessionService - from trpc_agent_sdk.sessions import RedisSessionService - from trpc_agent_sdk.sessions import SqlSessionService - - if isinstance(session_service, RedisSessionService): - resolved_config = replace( - compact_config, - storage_backend="redis", - redis_url=session_service.db_url, - redis_is_async=session_service.is_async, - ) - elif isinstance(session_service, SqlSessionService): - resolved_config = replace( - compact_config, - storage_backend="sql", - sql_url=session_service.db_url, - sql_is_async=session_service.is_async, - ) - elif isinstance(session_service, InMemorySessionService): - resolved_config = replace(compact_config, storage_backend="local") - else: - raise TypeError( - "Advanced Compact supports InMemorySessionService, " - "RedisSessionService, and SqlSessionService" - ) - runtime = AdvancedMemoryRuntime.create(resolved_config) - extractor = SessionMemoryExtractor( - runtime, - session_memory_generator, - model=session_memory_model, - ) - manager = AdvancedSessionCompactManager(runtime, extractor) - setup_tool_result_budget(agent, runtime) - setup_history_snip(agent, runtime) - setup_microcompact(agent, runtime) - autocompact = setup_autocompact( - agent, - runtime, - summary_generator, - model=compact_model, - ) - autocompact.attach_session_memory_extractor(extractor) - session_service.set_session_compact_manager(manager) - return manager diff --git a/trpc_agent_sdk/sessions/compact/_manager.py b/trpc_agent_sdk/sessions/compact/_manager.py index ad0b0e6ef..f969d0cda 100644 --- a/trpc_agent_sdk/sessions/compact/_manager.py +++ b/trpc_agent_sdk/sessions/compact/_manager.py @@ -8,6 +8,7 @@ from __future__ import annotations +from typing import Any from typing import TYPE_CHECKING from ._base_manager import BaseSessionCompactManager @@ -19,8 +20,11 @@ from trpc_agent_sdk.context import InvocationContext from trpc_agent_sdk.sessions import Session - from ._runtime import AdvancedMemoryRuntime - from ._session_memory import SessionMemoryExtractor +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): @@ -28,22 +32,66 @@ class AdvancedSessionCompactManager(BaseSessionCompactManager): def __init__( self, - runtime: "AdvancedMemoryRuntime", - session_memory_extractor: "SessionMemoryExtractor", + 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 the compact runtime and post-turn memory extractor.""" - self._runtime = runtime - self._session_memory_extractor = session_memory_extractor + """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) -> "AdvancedMemoryRuntime": + 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( @@ -56,12 +104,11 @@ def set_session_service( 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)" - ) + raise ValueError("Advanced Session Compact requires " + "SessionServiceConfig(store_historical_events=True)") self._session_service = session_service - self._session_memory_extractor.attach_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, @@ -70,7 +117,7 @@ async def create_session_summary( ctx: "InvocationContext | None" = None, ) -> None: """Use the native post-turn hook to update persistent Session Memory.""" - if ctx is not None: + if ctx is not None and self._session_memory_extractor is not None: await self._session_memory_extractor.extract_if_needed( session, ctx, @@ -82,21 +129,7 @@ async def get_session_summary(self, session: "Session") -> str | None: parsed = parse_session_memory_state(session.state.get(SESSION_MEMORY_STATE_KEY)) if parsed is not None: return parsed[0].to_markdown() - runtime = self._runtime.for_session(session) - if runtime.session_memory is None: - return None - return await runtime.session_memory.read(session.id) - - async def delete_session( - self, - *, - app_name: str, - user_id: str, - session_id: str, - ) -> None: - """Delete compact side data after the framework Session is deleted.""" - await self._runtime.for_scope(app_name, user_id).delete_session(session_id) + return None async def close(self) -> None: - """Release Compact backend resources owned by this manager.""" - await self._runtime.close() + """Release Compact resources owned by the manager.""" diff --git a/trpc_agent_sdk/sessions/compact/_microcompact.py b/trpc_agent_sdk/sessions/compact/_microcompact.py index ec1ae94d5..896c45212 100644 --- a/trpc_agent_sdk/sessions/compact/_microcompact.py +++ b/trpc_agent_sdk/sessions/compact/_microcompact.py @@ -15,7 +15,7 @@ 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 +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]" @@ -75,7 +74,7 @@ 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] = {} @@ -83,7 +82,7 @@ def __init__(self, memory_runtime: AdvancedMemoryRuntime) -> None: 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 @@ -97,25 +96,14 @@ def _session_lock(self, session_id: str) -> asyncio.Lock: return lock async def _load_state(self, session_id: str) -> MicrocompactState: - """Restore cleaned tool-result identifiers from the transcript.""" + """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[state_key] = state return state @@ -152,29 +140,6 @@ 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", @@ -189,7 +154,6 @@ async def apply( if not config.enabled or not config.microcompact_enabled: return MicrocompactResult(None, 0, 0, 0) if ctx is None or hasattr(self._runtime, "scope"): - await self._runtime.initialize() return await self._apply_scoped( request, session_id=session_id, @@ -259,7 +223,6 @@ async def _apply_scoped( 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 @@ -300,7 +263,7 @@ async def __call__(self, ctx: "InvocationContext", request: "LlmRequest") -> Non 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/_paths.py b/trpc_agent_sdk/sessions/compact/_paths.py deleted file mode 100644 index 87a473c17..000000000 --- a/trpc_agent_sdk/sessions/compact/_paths.py +++ /dev/null @@ -1,208 +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 AdvancedCompactConfig - -_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.""" - if value != value.strip() or any(character.isspace() and character not in {" "} for character in value): - raise ValueError(f"{field_name} must not contain leading/trailing or control whitespace") - if any(ord(character) < 32 or ord(character) == 127 for character in value): - raise ValueError(f"{field_name} must not contain control characters") - 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 MemoryScope: - """Identify the application and user that own Advanced Memory data.""" - - 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 a stable process-local key for locks and caches.""" - return repr((self.app_name, self.user_id)) - - -@dataclass(frozen=True) -class AdvancedMemoryPaths: - """Build all disk paths for long-term and session memory.""" - - config: AdvancedCompactConfig - scope: MemoryScope | None = None - - def for_scope(self, app_name: str, user_id: str) -> "AdvancedMemoryPaths": - """Return paths rooted in the given application's user namespace.""" - return AdvancedMemoryPaths(self.config, MemoryScope(app_name, user_id)) - - @property - def tenant_root_dir(self) -> Path: - """Return this scope's root, or the legacy root when unscoped.""" - 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 a key suitable for lock and cache partitioning.""" - return self.scope.storage_key if self.scope is not None else "legacy\0global" - - @property - def memory_dir(self) -> Path: - """Return the long-term memory directory.""" - return self.tenant_root_dir / self.config.memory_dir_name - - @property - def session_root_dir(self) -> Path: - """Return the root directory for session memory.""" - return self.tenant_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 storage_reference( - self, - resource: str, - *, - session_id: str | None = None, - topic_name: str | None = None, - result_id: str | None = None, - ) -> str: - """Return a model-visible reference for a stored Advanced Memory resource.""" - if resource == "memory_index": - local_path = self.memory_index_path - elif resource == "memory_topic": - if topic_name is None: - raise ValueError("topic_name is required for a memory topic reference") - local_path = self.memory_topic_path(topic_name) - elif resource == "transcript": - if session_id is None: - raise ValueError("session_id is required for a transcript reference") - local_path = self.transcript_path(session_id) - elif resource == "session_memory": - if session_id is None: - raise ValueError("session_id is required for a session memory reference") - local_path = self.session_memory_path(session_id) - elif resource == "tool_result": - if session_id is None or result_id is None: - raise ValueError("session_id and result_id are required for a tool result reference") - local_path = self.tool_result_path(session_id, result_id) - else: - raise ValueError(f"Unknown Advanced Memory resource: {resource}") - if self.config.storage_backend == "local": - return str(local_path) - if self.scope is None: - raise ValueError("A scoped path is required for non-local memory storage") - if resource == "session_memory": - return ("session-state://" - f"{self.scope.app_name}/{self.scope.user_id}/{session_id}/" - "_trpc_agent:summary") - - app_component = self.tenant_root_dir.parent.name - user_component = self.tenant_root_dir.name - if self.config.storage_backend == "redis": - user_base = f"{self.config.redis_key_prefix}:{{{app_component}:{user_component}}}" - if resource == "memory_index": - key = f"{user_base}:memory:index" - elif resource == "memory_topic": - key = f"{user_base}:memory:topic:{local_path.name}" - else: - safe_session_id = self.session_dir(session_id or "").name - session_base = f"{self.config.redis_key_prefix}:{{{app_component}:{user_component}:{safe_session_id}}}" - if resource == "transcript": - key = f"{session_base}:transcript" - else: - key = f"{session_base}:tool:{result_id}" - return f"advanced-memory://redis/{key}" - - app_name = self.scope.app_name - user_id = self.scope.user_id - if resource == "memory_index": - suffix = "memory/index" - elif resource == "memory_topic": - suffix = f"memory/topic/{local_path.name}" - elif resource == "transcript": - suffix = f"{session_id}/transcript" - else: - suffix = f"{session_id}/tool/{self.tool_result_path(session_id or '', result_id or '').stem}" - return f"advanced-memory://sql/{app_name}/{user_id}/{suffix}" - - 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/sessions/compact/_redis_stores.py b/trpc_agent_sdk/sessions/compact/_redis_stores.py deleted file mode 100644 index b60363d8c..000000000 --- a/trpc_agent_sdk/sessions/compact/_redis_stores.py +++ /dev/null @@ -1,297 +0,0 @@ -"""Redis implementations of the Advanced Memory storage contracts.""" - -from __future__ import annotations - -import asyncio -import json -from collections.abc import Mapping -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 AdvancedCompactConfig -from ._formats import MemoryDocument, MemoryIndexEntry -from ._paths import AdvancedMemoryPaths - -_APPEND_UNIQUE_SCRIPT = """ -if redis.call('SADD', KEYS[2], ARGV[1]) == 0 then return 0 end -redis.call('XADD', KEYS[1], '*', 'data', ARGV[2]) -return 1 -""" - -_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: AdvancedCompactConfig, 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}}}" - self._app_base = f"{config.redis_key_prefix}:{{{app_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 _session_base(self, session_id: str) -> str: - safe_session_id = self._paths.session_dir(session_id).name - tenant = f"{self._paths.tenant_root_dir.parent.name}:{self._paths.tenant_root_dir.name}" - return f"{self._config.redis_key_prefix}:{{{tenant}:{safe_session_id}}}" - - def _session_registry(self, session_id: str) -> str: - return f"{self._session_base(session_id)}:keys" - - 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_session_ttl(self, session_id: str, *keys: str) -> None: - skip_prefixes: tuple[str, ...] = () - if not self._config.session_ttl_delete_transcripts: - skip_prefixes = (f"{self._session_base(session_id)}:transcript", ) - await self._refresh_ttl_group( - self._session_registry(session_id), - list(keys), - self._config.session_ttl_seconds, - skip_prefixes=skip_prefixes, - ) - - async def _refresh_memory_ttl(self, *keys: str) -> None: - await self._refresh_ttl_group( - self._memory_registry(), - list(keys), - self._config.memory_ttl_seconds, - ) - - async def delete_session(self, session_id: str) -> None: - """Delete all Advanced Memory keys for one session.""" - session_base = self._session_base(session_id) - registry = self._session_registry(session_id) - keys: set[str] = {registry} - tracked = await self._command("smembers", registry) or [] - keys.update(value for value in (self._text(item) for item in tracked) if value) - - cursor: Any = 0 - pattern = f"{session_base}:*" - while True: - cursor, scanned = await self._command( - "scan", - cursor, - match=pattern, - count=100, - ) - keys.update(value for value in (self._text(item) for item in scanned) if value) - if int(cursor) == 0: - break - if keys: - await self._command("delete", *keys) - - @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() - lines, used_bytes = [], 0 - for line in value.splitlines(keepends=True)[:self._config.memory_index_max_lines]: - size = len(line.encode(self._config.encoding)) - if used_bytes + size > self._config.memory_index_max_bytes: - break - lines.append(line) - used_bytes += size - return "".join(lines) - - 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] - - -class RedisToolResultStore(_RedisStore): - - async def write(self, session_id: str, result_id: str, serialized_result: str) -> Path: - key = f"{self._session_base(session_id)}:tool:{result_id}" - await self._command("set", key, serialized_result) - await self._refresh_session_ttl(session_id, key) - return Path(f"advanced-memory://{key}") - - async def read(self, session_id: str, result_id: str) -> str | None: - key = f"{self._session_base(session_id)}:tool:{result_id}" - value = await self._command("get", key) - await self._refresh_session_ttl(session_id, key) - return self._text(value) - - -class RedisTranscriptStore(_RedisStore): - - @staticmethod - def _validate_record(record: Mapping[str, Any]) -> None: - """Reject Event and Session Memory duplication in Redis.""" - if record.get("kind") in {"event", "session-memory-checkpoint"}: - raise ValueError("Redis transcripts only store context-compression records") - - async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: - self._validate_record(record) - payload = dict(record) - payload.setdefault("recorded_at", datetime.now(timezone.utc).isoformat()) - stream = f"{self._session_base(session_id)}:transcript" - await self._command("xadd", stream, {"data": json.dumps(payload)}) - await self._refresh_session_ttl(session_id, stream) - return Path(f"advanced-memory://{stream}") - - async def append_unique(self, session_id: str, record: Mapping[str, Any], *, unique_key: str) -> tuple[Path, bool]: - self._validate_record(record) - payload = dict(record) - value = payload.get(unique_key) - if not isinstance(value, str) or not 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()) - stream = f"{self._session_base(session_id)}:transcript" - seen = f"{stream}:seen:{unique_key}" - async with self._storage.create_db_session() as connection: - added = await self._storage.execute_command( - connection, - RedisCommand( - method="eval", - args=(_APPEND_UNIQUE_SCRIPT, 2, stream, seen, value, json.dumps(payload, ensure_ascii=False)), - )) - await self._refresh_session_ttl(session_id, stream, seen) - return Path(f"advanced-memory://{stream}"), bool(added) - - async def read_all(self, session_id: str) -> list[dict[str, Any]]: - stream = f"{self._session_base(session_id)}:transcript" - entries = await self._command("xrange", stream, "-", "+") - await self._refresh_session_ttl(session_id, stream) - records: list[dict[str, Any]] = [] - for _, fields in entries: - value = fields.get(b"data") if isinstance(fields, dict) else None - value = value or fields.get("data") - text = self._text(value) - if text: - records.append(json.loads(text)) - return records diff --git a/trpc_agent_sdk/sessions/compact/_runtime.py b/trpc_agent_sdk/sessions/compact/_runtime.py index e0bd6d24f..f5734e2c5 100644 --- a/trpc_agent_sdk/sessions/compact/_runtime.py +++ b/trpc_agent_sdk/sessions/compact/_runtime.py @@ -1,258 +1,53 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# Tencent is pleased to support the open source ecosystem. # # 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.""" +# Licensed under Apache-2.0. +"""Runtime coordination for Session Compact.""" from __future__ import annotations from dataclasses import dataclass -from dataclasses import field -import asyncio -import shutil -import threading -from typing import Any from ._config import AdvancedCompactConfig -from ._coordination import CrossLoopLock from ._coordination import SessionOperationCoordinator -from ._paths import AdvancedMemoryPaths -from ._paths import MemoryScope -from ._storage import LocalAdvancedMemoryCleanup -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.""" +@dataclass +class SessionCompactRuntime: + """Hold compression configuration and per-session coordination only.""" config: AdvancedCompactConfig - paths: AdvancedMemoryPaths coordination: SessionOperationCoordinator - long_term_memory: LongTermMemoryStore - session_memory: SessionMemoryStore | None - tool_results: ToolResultStore - transcripts: TranscriptStore - _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: AdvancedCompactConfig | None = None) -> "AdvancedMemoryRuntime": - """Create a runtime isolated from the legacy mechanism.""" - resolved_config = config or AdvancedCompactConfig() - 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 - sql_storage = SqlStorage( - is_async=resolved_config.sql_is_async, - db_url=resolved_config.sql_url, - metadata=AdvancedMemorySqlBase.metadata, - expire_on_commit=False, - ) - from ._sql_stores import SqlAdvancedMemoryCleanup - sql_cleanup = SqlAdvancedMemoryCleanup(resolved_config, sql_storage) - else: - local_cleanup = LocalAdvancedMemoryCleanup(resolved_config) - return cls( - config=resolved_config, - paths=paths, - coordination=SessionOperationCoordinator(), - long_term_memory=LongTermMemoryStore(resolved_config, paths), - session_memory=(SessionMemoryStore(resolved_config, paths) - if resolved_config.storage_backend == "local" else None), - tool_results=ToolResultStore(resolved_config, paths), - transcripts=TranscriptStore(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 - from ._redis_stores import RedisToolResultStore - from ._redis_stores import RedisTranscriptStore + def create(cls, config: AdvancedCompactConfig | None = None) -> "SessionCompactRuntime": + return cls(config or AdvancedCompactConfig(), SessionOperationCoordinator()) - 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) - session_memory = None - tool_results = RedisToolResultStore(self.config, paths, storage) - transcripts = RedisTranscriptStore(self.config, paths, storage) - elif self.config.storage_backend == "sql": - from ._sql_stores import SqlLongTermMemoryStore - from ._sql_stores import SqlToolResultStore - from ._sql_stores import SqlTranscriptStore - 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) - session_memory = None - tool_results = SqlToolResultStore(self.config, paths, storage) - transcripts = SqlTranscriptStore(self.config, paths, storage) - else: - long_term_memory = LongTermMemoryStore(self.config, paths) - session_memory = SessionMemoryStore(self.config, paths) - tool_results = ToolResultStore(self.config, paths) - transcripts = TranscriptStore(self.config, paths) - runtime = ScopedAdvancedMemoryRuntime( - root=self, - scope=scope, - paths=paths, - long_term_memory=long_term_memory, - session_memory=session_memory, - tool_results=tool_results, - transcripts=transcripts, - ) - self._scoped_runtimes[scope] = runtime - return runtime - - def for_session(self, session: object) -> "ScopedAdvancedMemoryRuntime": - """Return the scoped runtime for a SessionABC-compatible object.""" + 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("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. + raise ValueError("Session Compact requires session app_name and user_id") + return ScopedSessionCompactRuntime(self, f"{app_name}\0{user_id}") - 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() and not legacy_paths.session_root_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)) - if legacy_paths.session_root_dir.exists(): - shutil.move(str(legacy_paths.session_root_dir), str(scoped.paths.session_root_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") - async with self._sql_storage.create_db_session(): - pass - if self._sql_cleanup is not None: - await self._sql_cleanup.start() - 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 +@dataclass +class ScopedSessionCompactRuntime: + """Session-scoped view used by compression callbacks.""" - 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_storage is not None: - if self._sql_cleanup is not None: - await self._sql_cleanup.close() - 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 - session_memory: SessionMemoryStore | None - tool_results: ToolResultStore - transcripts: TranscriptStore + root: SessionCompactRuntime + scope: str @property def config(self) -> AdvancedCompactConfig: - """Return the root runtime configuration.""" return self.root.config @property def coordination(self) -> SessionOperationCoordinator: - """Return the shared coordinator.""" return self.root.coordination def session_key(self, session_id: str) -> str: - """Return a lock/cache key unique across all tenants.""" - return f"{self.scope.storage_key}\0{session_id}" - - async def initialize(self) -> bool: - """Initialize only this tenant's local directories.""" - if not self.config.enabled: - return False - if self.config.storage_backend == "sql" and self.root._sql_cleanup is not None: - await self.root._sql_cleanup.start() - if self.config.storage_backend == "local" and self.root._local_cleanup is not None: - await self.root._local_cleanup.start() - await self.long_term_memory.initialize() - return True + return f"{self.scope}\0{session_id}" - async def delete_session(self, session_id: str) -> None: - """Delete all Advanced Memory data belonging to one session.""" - if self.config.storage_backend == "local": - session_dir = self.paths.session_dir(session_id) - await asyncio.to_thread(shutil.rmtree, session_dir, True) - return - delete_session = getattr(self.tool_results, "delete_session", None) - if delete_session is None: - raise RuntimeError("Configured Advanced Memory backend cannot delete sessions") - await delete_session(session_id) + def for_session(self, session: object) -> "ScopedSessionCompactRuntime": + return self.root.for_session(session) diff --git a/trpc_agent_sdk/sessions/compact/_session_memory.py b/trpc_agent_sdk/sessions/compact/_session_memory.py index 35ddc81f2..2e9961289 100644 --- a/trpc_agent_sdk/sessions/compact/_session_memory.py +++ b/trpc_agent_sdk/sessions/compact/_session_memory.py @@ -32,7 +32,7 @@ from ._formats import SessionMemoryDocument from ._formats import build_session_memory_state from ._formats import parse_session_memory_state -from ._runtime import AdvancedMemoryRuntime +from ._runtime import SessionCompactRuntime from ._token_budget import TokenContextTracker if TYPE_CHECKING: @@ -40,7 +40,6 @@ 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( @@ -150,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. """ @@ -342,7 +341,7 @@ class SessionMemoryExtractor: def __init__( self, - memory_runtime: AdvancedMemoryRuntime, + memory_runtime: SessionCompactRuntime, generator: SessionMemoryGenerator | None = None, *, model: Any | None = None, @@ -359,15 +358,10 @@ def __init__( self._session_service = session_service @property - def runtime(self) -> AdvancedMemoryRuntime: + def runtime(self) -> SessionCompactRuntime: """Return the runtime bound to this extractor.""" return self._runtime - @property - def uses_session_state(self) -> bool: - """Return whether this backend stores Session Memory in Session.state.""" - return self._runtime.config.storage_backend in {"redis", "sql"} - 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: @@ -412,7 +406,7 @@ def _event_records_after_checkpoint( checkpoint_event_id: str | None, checkpoint_recorded_at: str | None = None, ) -> list[dict[str, Any]]: - """Return Event transcript records after the checkpoint in order.""" + """Return Session Event records after the checkpoint in order.""" event_records = [record for record in records if record.get("kind") == "event"] if checkpoint_event_id is None: return event_records @@ -432,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, @@ -538,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", []) @@ -553,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) @@ -659,17 +643,9 @@ def missing_context(end: int) -> list[str]: return [], None async def _read_current_memory(self, session: "SessionABC") -> str: - """Read old session memory or return the complete empty template.""" - if self.uses_session_state: - parsed = parse_session_memory_state(session.state.get(SESSION_MEMORY_STATE_KEY)) - if parsed is not None: - return parsed[0].to_markdown() - return SessionMemoryDocument().to_markdown() - store = self._runtime.for_session(session).session_memory - if store is None: - raise RuntimeError("Session Memory store is unavailable") - current = await store.read(session.id) - return current if current is not None else SessionMemoryDocument().to_markdown() + """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, @@ -735,49 +711,31 @@ async def _persist_checkpoint( document.key_results, document.worklog, ) - if self.uses_session_state: - if self._session_service is None: - raise RuntimeError("Redis/SQL 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}, - ) - return - runtime = self._runtime.for_session(session) - await 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( @@ -792,19 +750,12 @@ async def extract_if_needed( if not config.enabled or not config.session_memory_enabled: return SessionMemoryExtractionResult(False, "disabled") runtime = self._runtime.for_session(session) - await runtime.initialize() 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") - if self.uses_session_state: - records = self._session_event_records(session) - checkpoint, checkpoint_context_tokens = self._state_checkpoint(session) - else: - records = await runtime.transcripts.read_all(session.id) - checkpoint = self._last_checkpoint(records) - checkpoint_context_tokens = (checkpoint.get("context_tokens") if checkpoint is not None - and isinstance(checkpoint.get("context_tokens"), int) else None) + 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( @@ -851,10 +802,6 @@ async def extract_if_needed( max_chars=config.session_memory_section_max_chars, total_max_chars=config.session_memory_total_max_chars, ) - if not self.uses_session_state: - if runtime.session_memory is None: - raise RuntimeError("Session Memory store is unavailable") - await runtime.session_memory.write(session.id, document) await self._persist_checkpoint( session, included, diff --git a/trpc_agent_sdk/sessions/compact/_session_service.py b/trpc_agent_sdk/sessions/compact/_session_service.py deleted file mode 100644 index e1f1f2441..000000000 --- a/trpc_agent_sdk/sessions/compact/_session_service.py +++ /dev/null @@ -1,236 +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.""" - if isinstance(delegate, TranscriptSessionService): - raise ValueError("Transcript session service is already wrapped") - 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_config(self) -> Any: - """Expose the original service configuration.""" - return getattr(self._delegate, "session_config", None) - - @property - def summarizer_manager(self) -> Any: - """Expose the original service summarizer, when configured.""" - return getattr(self._delegate, "summarizer_manager", None) - - @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, session: SessionABC) -> 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.for_session(session).initialize() - - def _session_lock(self, session: SessionABC) -> CrossLoopLock: - """Return an independent asynchronous write lock per session.""" - key = self._memory_runtime.for_session(session).session_key(session.id) - lock = self._session_locks.get(key) - if lock is None: - lock = CrossLoopLock() - self._session_locks[key] = lock - return lock - - async def _load_parent_if_needed(self, session: SessionABC) -> None: - """Restore the parent-chain tail before the first session write.""" - runtime = self._memory_runtime.for_session(session) - key = runtime.session_key(session.id) - if key in self._loaded_parent_sessions: - return - records = await runtime.transcripts.read_all(session.id) - self._last_event_ids[key] = find_last_event_id(records) - self._loaded_parent_sessions.add(key) - - 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 the framework session and all Advanced Memory session data.""" - runtime = self._memory_runtime.for_scope(app_name, user_id) - scope_key = runtime.session_key(session_id) - lock = self._session_locks.setdefault(scope_key, CrossLoopLock()) - async with lock: - await self._delegate.delete_session( - app_name=app_name, - user_id=user_id, - session_id=session_id, - ) - await runtime.delete_session(session_id) - self._session_locks.pop(scope_key, None) - self._loaded_parent_sessions.discard(scope_key) - self._last_event_ids.pop(scope_key, 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(session) - runtime = self._memory_runtime.for_session(session) - key = runtime.session_key(session.id) - async with self._session_lock(session): - await self._load_parent_if_needed(session) - record = build_event_transcript_record( - session, - persisted_event, - parent_event_id=self._last_event_ids.get(key), - ) - _, appended = await runtime.transcripts.append_unique( - session.id, - record, - unique_key="event_id", - ) - if appended: - self._last_event_ids[key] = 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 patch_session_state( - self, - session: SessionABC, - state_delta: dict[str, Any], - ) -> None: - """Delegate state-only updates without touching persisted Events.""" - await self._delegate.patch_session_state(session, state_delta) - - 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/sessions/compact/_sql_stores.py b/trpc_agent_sdk/sessions/compact/_sql_stores.py deleted file mode 100644 index 4c77eae64..000000000 --- a/trpc_agent_sdk/sessions/compact/_sql_stores.py +++ /dev/null @@ -1,528 +0,0 @@ -"""SQL implementations of the Advanced Memory storage contracts.""" - -from __future__ import annotations - -import json -import asyncio -import hashlib -import uuid -from datetime import datetime, timedelta, timezone -from dataclasses import replace -from pathlib import Path -from collections.abc import Mapping -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 AdvancedCompactConfig -from ._formats import MemoryDocument, MemoryIndexEntry -from ._paths import AdvancedMemoryPaths - - -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 SqlTranscript(AdvancedMemorySqlBase): - __tablename__ = "advanced_memory_transcripts" - - 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) - session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) - record_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) - payload: Mapped[str] = mapped_column(Text) - recorded_at: Mapped[datetime] = mapped_column(PreciseTimestamp, default=func.now()) - expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) - - -class SqlTranscriptSeen(AdvancedMemorySqlBase): - __tablename__ = "advanced_memory_transcript_seen" - - dedupe_id: Mapped[str] = mapped_column(String(64), primary_key=True) - app_name: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) - user_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) - session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) - unique_key: Mapped[str] = mapped_column(String(DEFAULT_MAX_VARCHAR_LENGTH)) - unique_value: Mapped[str] = mapped_column(String(DEFAULT_MAX_VARCHAR_LENGTH)) - expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) - - -class SqlToolResult(AdvancedMemorySqlBase): - __tablename__ = "advanced_memory_tool_results" - - 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) - session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) - result_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_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: AdvancedCompactConfig, 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 - - async def _refresh_session_scope(self, db: Any, session_id: str) -> None: - expiry = self._expiry(self._config.session_ttl_seconds) - if expiry is None: - return - tables = ((SqlToolResult, (self._app_name, self._user_id, session_id)), ) - if self._config.session_ttl_delete_transcripts: - tables = ( - (SqlTranscript, (self._app_name, self._user_id, session_id)), - (SqlTranscriptSeen, (self._app_name, self._user_id, session_id)), - *tables, - ) - for model, key in tables: - rows = await self._storage.query( - db, - SqlKey(key=key, storage_cls=model), - SqlCondition(filters=[ - getattr(model, "app_name") == self._app_name, - getattr(model, "user_id") == self._user_id, - getattr(model, "session_id") == session_id, - getattr(model, "expires_at").is_(None) | (getattr(model, "expires_at") > self._now()), - ]), - ) - for row in rows: - row.expires_at = expiry - - async def delete_session(self, session_id: str) -> None: - """Delete all Advanced Memory rows for one session.""" - models = ( - SqlTranscript, - SqlTranscriptSeen, - SqlToolResult, - ) - filters = { - SqlTranscript: [ - SqlTranscript.app_name == self._app_name, - SqlTranscript.user_id == self._user_id, - SqlTranscript.session_id == session_id, - ], - SqlTranscriptSeen: [ - SqlTranscriptSeen.app_name == self._app_name, - SqlTranscriptSeen.user_id == self._user_id, - SqlTranscriptSeen.session_id == session_id, - ], - SqlToolResult: [ - SqlToolResult.app_name == self._app_name, - SqlToolResult.user_id == self._user_id, - SqlToolResult.session_id == session_id, - ], - } - async with self._storage.create_db_session() as db: - for model in models: - await self._storage.delete( - db, - SqlKey(key=tuple(), storage_cls=model), - SqlCondition(filters=filters[model]), - ) - await self._storage.commit(db) - - -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) - await self._storage.commit(db) - content = row.content - lines, used_bytes = [], 0 - for line in content.splitlines(keepends=True)[:self._config.memory_index_max_lines]: - size = len(line.encode(self._config.encoding)) - if used_bytes + size > self._config.memory_index_max_bytes: - break - lines.append(line) - used_bytes += size - return "".join(lines) - - 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 SqlToolResultStore(_SqlStore): - - async def write(self, session_id: str, result_id: str, serialized_result: str) -> Path: - async with self._storage.create_db_session() as db: - key = (self._app_name, self._user_id, session_id, result_id) - row = await self._storage.get(db, SqlKey(key=key, storage_cls=SqlToolResult)) - if row is None: - row = SqlToolResult( - app_name=key[0], - user_id=key[1], - session_id=key[2], - result_id=key[3], - ) - await self._storage.add(db, row) - row.content = serialized_result - row.updated_at = self._now() - row.expires_at = self._expiry(self._config.session_ttl_seconds) - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/tool/{result_id}") - - async def read(self, session_id: str, result_id: str) -> str | None: - async with self._storage.create_db_session() as db: - row = await self._storage.get( - db, - SqlKey(key=(self._app_name, self._user_id, session_id, result_id), storage_cls=SqlToolResult), - ) - if row is None or self._expired(row.expires_at): - return None - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return row.content - - -class SqlTranscriptStore(_SqlStore): - - @staticmethod - def _validate_record(record: Mapping[str, Any]) -> None: - """Reject Event and Session Memory duplication in SQL.""" - if record.get("kind") in {"event", "session-memory-checkpoint"}: - raise ValueError("SQL transcripts only store context-compression records") - - def _dedupe_id(self, session_id: str, unique_key: str, value: str) -> str: - raw = "\0".join((self._app_name, self._user_id, session_id, unique_key, value)) - return hashlib.sha256(raw.encode("utf-8")).hexdigest() - - async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: - self._validate_record(record) - payload = dict(record) - payload.setdefault("recorded_at", self._now().replace(tzinfo=timezone.utc).isoformat()) - async with self._storage.create_db_session() as db: - await self._storage.add( - db, - SqlTranscript( - app_name=self._app_name, - user_id=self._user_id, - session_id=session_id, - record_id=uuid.uuid4().hex, - payload=json.dumps(payload, ensure_ascii=False), - expires_at=(self._expiry(self._config.session_ttl_seconds) - if self._config.session_ttl_delete_transcripts else None), - )) - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript") - - async def append_unique( - self, - session_id: str, - record: Mapping[str, Any], - *, - unique_key: str, - ) -> tuple[Path, bool]: - self._validate_record(record) - payload = dict(record) - value = payload.get(unique_key) - if not isinstance(value, str) or not value: - raise ValueError(f"Transcript unique key {unique_key!r} must be a non-empty string") - async with self._storage.create_db_session() as db: - dedupe_id = self._dedupe_id(session_id, unique_key, value) - seen_key = (self._app_name, self._user_id, session_id, unique_key, value) - seen = await self._storage.get( - db, - SqlKey(key=(dedupe_id, ), storage_cls=SqlTranscriptSeen), - ) - if seen is not None and not self._expired(seen.expires_at): - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript"), False - if seen is not None: - await self._storage.delete( - db, - SqlKey(key=(dedupe_id, ), storage_cls=SqlTranscriptSeen), - SqlCondition(filters=[ - SqlTranscriptSeen.dedupe_id == dedupe_id, - ]), - ) - payload.setdefault("recorded_at", self._now().replace(tzinfo=timezone.utc).isoformat()) - await self._storage.add( - db, - SqlTranscriptSeen( - dedupe_id=dedupe_id, - app_name=seen_key[0], - user_id=seen_key[1], - session_id=seen_key[2], - unique_key=seen_key[3], - unique_value=seen_key[4], - expires_at=(self._expiry(self._config.session_ttl_seconds) - if self._config.session_ttl_delete_transcripts else None), - )) - await self._storage.add( - db, - SqlTranscript( - app_name=self._app_name, - user_id=self._user_id, - session_id=session_id, - record_id=uuid.uuid4().hex, - payload=json.dumps(payload, ensure_ascii=False), - expires_at=(self._expiry(self._config.session_ttl_seconds) - if self._config.session_ttl_delete_transcripts else None), - )) - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript"), True - - async def read_all(self, session_id: str) -> list[dict[str, Any]]: - async with self._storage.create_db_session() as db: - rows = await self._storage.query( - db, - SqlKey(key=(self._app_name, self._user_id, session_id), storage_cls=SqlTranscript), - SqlCondition( - filters=[ - SqlTranscript.app_name == self._app_name, - SqlTranscript.user_id == self._user_id, - SqlTranscript.session_id == session_id, - SqlTranscript.expires_at.is_(None) | (SqlTranscript.expires_at > self._now()), - ], - order_func=SqlTranscript.recorded_at.asc, - ), - ) - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return [json.loads(row.payload) for row in rows] - - -class SqlAdvancedMemoryCleanup: - """Periodically remove expired Advanced Memory SQL rows.""" - - _models = ( - SqlMemoryIndex, - SqlMemoryTopic, - SqlTranscript, - SqlTranscriptSeen, - SqlToolResult, - ) - - def __init__(self, config: AdvancedCompactConfig, 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 - and self._config.session_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: - models = self._models if self._config.session_ttl_delete_transcripts else tuple( - model for model in self._models if model is not SqlTranscript) - for model in models: - await self._storage.delete( - db, - SqlKey(key=tuple(), storage_cls=model), - SqlCondition(filters=[model.expires_at.is_not(None), model.expires_at <= now]), - ) - 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", - "SqlToolResultStore", - "SqlTranscriptStore", -] diff --git a/trpc_agent_sdk/sessions/compact/_storage.py b/trpc_agent_sdk/sessions/compact/_storage.py deleted file mode 100644 index 98b174872..000000000 --- a/trpc_agent_sdk/sessions/compact/_storage.py +++ /dev/null @@ -1,499 +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 shutil -import tempfile -import threading -import time -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 AdvancedCompactConfig -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 - - -def _is_expired(path: Path, ttl: int | None) -> bool: - if ttl is None or not path.exists(): - return False - return time.time() - path.stat().st_mtime >= ttl - - -def _touch(path: Path) -> None: - path.parent.mkdir(parents=True, exist_ok=True) - path.touch() - - -def _expire_memory_dir(memory_dir: Path, config: AdvancedCompactConfig) -> bool: - """Expire the whole long-term memory group using index activity time.""" - index_path = memory_dir / config.memory_index_name - if not _is_expired(index_path, config.memory_ttl_seconds): - return False - for path in memory_dir.glob("*.md"): - path.unlink(missing_ok=True) - return True - - -def _refresh_memory_dir(memory_dir: Path) -> None: - """Refresh activity for every file in the long-term memory group.""" - for path in memory_dir.glob("*.md"): - _touch(path) - - -def _session_activity_path(session_dir: Path) -> Path: - return session_dir / ".advanced-memory-activity" - - -def _expire_session_dir(session_dir: Path, config: AdvancedCompactConfig) -> bool: - """Expire all Advanced Memory data belonging to one local session.""" - if not session_dir.exists() or config.session_ttl_seconds is None: - return False - activity_path = _session_activity_path(session_dir) - if activity_path.exists(): - expired = _is_expired(activity_path, config.session_ttl_seconds) - else: - files = [path for path in session_dir.rglob("*") if path.is_file()] - expired = bool(files) and time.time() - max(path.stat().st_mtime - for path in files) >= config.session_ttl_seconds - if expired: - if config.session_ttl_delete_transcripts: - shutil.rmtree(session_dir, ignore_errors=True) - else: - transcript_path = session_dir / config.transcript_name - for child in session_dir.iterdir(): - if child == transcript_path: - continue - if child.is_dir(): - shutil.rmtree(child, ignore_errors=True) - else: - child.unlink(missing_ok=True) - return expired - - -def _refresh_session_dir(session_dir: Path) -> None: - _touch(_session_activity_path(session_dir)) - - -class LongTermMemoryStore: - """Manage MEMORY.md and its detail files in the same directory.""" - - def __init__(self, config: AdvancedCompactConfig, 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 _expire_memory_dir(self._paths.memory_dir, self._config) or not self.index_path.exists(): - return "" - _refresh_memory_dir(self._paths.memory_dir) - 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) - _refresh_memory_dir(self._paths.memory_dir) - - 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_topic_sync, path) - - def _read_topic_sync(self, path: Path) -> str | None: - if _expire_memory_dir(self._paths.memory_dir, self._config) or not path.exists(): - return None - _refresh_memory_dir(self._paths.memory_dir) - return path.read_text(encoding=self._config.encoding) - - 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_sync, path) - - def _read_frontmatter_sync(self, path: Path) -> str | None: - """Synchronously read a topic's bounded frontmatter block.""" - if _expire_memory_dir(self._paths.memory_dir, self._config) or not path.exists(): - return None - _refresh_memory_dir(self._paths.memory_dir) - 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(self._write_topic_sync, path, document.to_markdown()) - return path - - def _write_topic_sync(self, path: Path, content: str) -> None: - _expire_memory_dir(self._paths.memory_dir, self._config) - _atomic_write_text(path, content, encoding=self._config.encoding) - _refresh_memory_dir(self._paths.memory_dir) - - 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 _expire_memory_dir(self._paths.memory_dir, self._config): - return [] - if not self._paths.memory_dir.exists(): - return [] - _refresh_memory_dir(self._paths.memory_dir) - 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: AdvancedCompactConfig, 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, session_id, path) - - def _read_sync(self, session_id: str, path: Path) -> str | None: - """Synchronously read session memory.""" - session_dir = self._paths.session_dir(session_id) - if _expire_session_dir(session_dir, self._config) or not path.exists(): - return None - _refresh_session_dir(session_dir) - 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( - self._write_sync, - session_id, - path, - document.to_markdown(), - ) - return path - - def _write_sync(self, session_id: str, path: Path, content: str) -> None: - session_dir = self._paths.session_dir(session_id) - _expire_session_dir(session_dir, self._config) - _atomic_write_text(path, content, encoding=self._config.encoding) - _refresh_session_dir(session_dir) - - -class ToolResultStore: - """Persist complete tool results that exceed the context budget.""" - - def __init__(self, config: AdvancedCompactConfig, 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( - self._write_sync, - session_id, - path, - serialized_result, - ) - 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, session_id, path) - - def _read_sync(self, session_id: str, path: Path) -> str | None: - """Synchronously read an optional complete tool-result file.""" - session_dir = self._paths.session_dir(session_id) - if _expire_session_dir(session_dir, self._config) or not path.exists(): - return None - _refresh_session_dir(session_dir) - return path.read_text(encoding=self._config.encoding) - - def _write_sync(self, session_id: str, path: Path, content: str) -> None: - session_dir = self._paths.session_dir(session_id) - _expire_session_dir(session_dir, self._config) - _atomic_write_text(path, content, encoding=self._config.encoding) - _refresh_session_dir(session_dir) - - -class TranscriptStore: - """Store complete per-session records as append-only JSONL.""" - - def __init__(self, config: AdvancedCompactConfig, 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.""" - _expire_session_dir(path.parent, self._config) - path.parent.mkdir(parents=True, exist_ok=True) - with self._write_lock: - self._append_serialized_unlocked(path, serialized) - _refresh_session_dir(path.parent) - - 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.""" - with self._write_lock: - if _expire_session_dir(path.parent, self._config): - for cache_key in list(self._seen_unique_values): - if cache_key[0] == path: - self._seen_unique_values.pop(cache_key, None) - path.parent.mkdir(parents=True, exist_ok=True) - cache_key = (path, unique_key) - 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) - _refresh_session_dir(path.parent) - 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: - expired = _expire_session_dir(path.parent, self._config) - if expired and self._config.session_ttl_delete_transcripts: - return [] - if not path.exists(): - return [] - _refresh_session_dir(path.parent) - 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 - - -class LocalAdvancedMemoryCleanup: - """Periodically remove expired local Advanced Memory data.""" - - def __init__(self, config: AdvancedCompactConfig) -> 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: - return - if self._config.memory_ttl_seconds is None and self._config.session_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] - session_roots = [root / self._config.session_dir_name] - tenants_root = root / "tenants" - if tenants_root.exists(): - for app_dir in tenants_root.iterdir(): - if app_dir.is_dir(): - for user_dir in app_dir.iterdir(): - if user_dir.is_dir(): - memory_dirs.append(user_dir / self._config.memory_dir_name) - session_roots.append(user_dir / self._config.session_dir_name) - for memory_dir in memory_dirs: - _expire_memory_dir(memory_dir, self._config) - for session_root in session_roots: - if session_root.exists(): - for session_dir in session_root.iterdir(): - if session_dir.is_dir(): - _expire_session_dir(session_dir, self._config) - - async def _run(self) -> None: - if self._stop_event is None: - return - ttls = [ - ttl for ttl in ( - self._config.memory_ttl_seconds, - self._config.session_ttl_seconds, - ) if ttl is not None - ] - interval = min(ttls) if ttls else 60 - try: - while not self._stop_event.is_set(): - try: - await asyncio.wait_for(self._stop_event.wait(), timeout=interval) - break - except asyncio.TimeoutError: - await self.cleanup_once() - except asyncio.CancelledError: - raise - - async def close(self) -> None: - if self._task is not None: - await self.cleanup_once() - 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/sessions/compact/_tool_result_budget.py b/trpc_agent_sdk/sessions/compact/_tool_result_budget.py index 7181584e5..3ad6cda17 100644 --- a/trpc_agent_sdk/sessions/compact/_tool_result_budget.py +++ b/trpc_agent_sdk/sessions/compact/_tool_result_budget.py @@ -12,12 +12,11 @@ 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 @@ -41,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 @@ -49,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 @@ -67,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, @@ -92,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]: @@ -116,7 +112,7 @@ 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] = {} @@ -124,7 +120,7 @@ def __init__(self, memory_runtime: AdvancedMemoryRuntime) -> None: 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 @@ -138,40 +134,36 @@ def _session_lock(self, session_id: str) -> asyncio.Lock: return lock async def _load_state(self, session_id: str) -> ToolResultBudgetState: - """Restore frozen results and historical replacements from the transcript.""" + """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[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] = [] @@ -195,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), @@ -206,21 +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 = Path( - self._runtime.paths.storage_reference( - "tool_result", - session_id=session_id, - result_id=candidate.result_id, - )) - persisted_path_text = str(persisted_path).replace( - "advanced-memory:/", - "advanced-memory://", - 1, - ) + """Build an event reference and model-visible preview.""" preview, truncated = _preview_text( candidate.serialized_result, self._runtime.config.tool_result_preview_chars, @@ -230,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 persisted.", - "path": persisted_path_text, - "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]: @@ -262,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] = [] @@ -281,67 +259,13 @@ 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( - self, - session_id: str, - replacement: ToolResultReplacement, - ) -> None: - """Persist the full result before appending its replacement record.""" - candidate = replacement.candidate - persisted_path = await self._runtime.tool_results.write( - session_id, - candidate.result_id, - candidate.serialized_result, - ) - persisted_path_text = str(persisted_path).replace( - "advanced-memory:/", - "advanced-memory://", - 1, - ) - replacement.replacement_response["persisted_output"]["path"] = persisted_path_text - 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": persisted_path_text, - "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", @@ -353,8 +277,7 @@ async def apply( if not self._runtime.config.enabled: return ToolResultBudgetResult(0, 0, 0) if ctx is None or hasattr(self._runtime, "scope"): - await self._runtime.initialize() - return await self._apply_scoped(request, session_id) + 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: @@ -365,22 +288,26 @@ async def apply( 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) -> ToolResultBudgetResult: + 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) @@ -389,7 +316,6 @@ async def _apply_scoped(self, request: "LlmRequest", session_id: str) -> ToolRes 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) @@ -436,7 +362,7 @@ async def __call__(self, ctx: "InvocationContext", request: "LlmRequest") -> Non 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/sessions/compact/_transcript.py b/trpc_agent_sdk/sessions/compact/_transcript.py deleted file mode 100644 index 6cfe2192b..000000000 --- a/trpc_agent_sdk/sessions/compact/_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/tools/_advanced_memory_tool.py b/trpc_agent_sdk/tools/_advanced_memory_tool.py index 1601b41eb..76bce5249 100644 --- a/trpc_agent_sdk/tools/_advanced_memory_tool.py +++ b/trpc_agent_sdk/tools/_advanced_memory_tool.py @@ -16,7 +16,7 @@ from trpc_agent_sdk.sessions.compact._formats import MemoryType from trpc_agent_sdk.sessions.compact._formats import memory_freshness from trpc_agent_sdk.sessions.compact._formats import parse_memory_updated_at -from trpc_agent_sdk.sessions.compact._runtime import AdvancedMemoryRuntime +from trpc_agent_sdk.advanced_memory._runtime import AdvancedMemoryRuntime from ._function_tool import FunctionTool From 4116be24b3d1f373418963317d5e355fd5ab0b6b Mon Sep 17 00:00:00 2001 From: raychen <815315825@qq.com> Date: Fri, 11 Sep 2026 16:19:36 +0800 Subject: [PATCH 5/6] =?UTF-8?q?feature:=20=E6=9B=B4=E6=96=B0=E4=BA=86?= =?UTF-8?q?=E4=BB=A3=E7=A0=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../run_agent.py | 134 ---- .../run_agent.py | 38 +- .../run_agent.py | 38 +- .../session_summarizer_with_advanced/.env | 4 + .../README.md | 91 +++ .../agent/__init__.py | 5 + .../agent/agent.py | 38 ++ .../agent/config.py | 19 + .../agent/prompts.py | 15 + .../agent/tools.py | 6 + .../run_agent.py | 173 +++++ tests/sessions/compact/test_coordination.py | 2 +- .../sessions/compact/test_session_compact.py | 198 +++++- tests/sessions/compact/test_token_budget.py | 29 +- tests/sessions/replay/backends.py | 4 +- tests/sessions/test_base_session_service.py | 18 +- .../test_in_memory_session_service.py | 2 +- tests/sessions/test_redis_session_service.py | 4 +- tests/sessions/test_session_summarizer.py | 6 +- tests/sessions/test_sql_session_service.py | 2 +- tests/sessions/test_summarizer_checker.py | 8 +- tests/sessions/test_summarizer_manager.py | 10 +- trpc_agent_sdk/abc/__init__.py | 6 + trpc_agent_sdk/abc/_compact.py | 187 ++++++ trpc_agent_sdk/abc/_session_service.py | 13 +- trpc_agent_sdk/advanced_memory/__init__.py | 57 -- trpc_agent_sdk/advanced_memory/_config.py | 83 --- .../advanced_memory/_integration.py | 106 --- .../advanced_memory/_memory_context.py | 126 ---- trpc_agent_sdk/advanced_memory/_paths.py | 110 --- .../advanced_memory/_preload_memory.py | 300 --------- .../advanced_memory/_redis_stores.py | 302 --------- trpc_agent_sdk/advanced_memory/_runtime.py | 206 ------ trpc_agent_sdk/advanced_memory/_sql_stores.py | 533 --------------- trpc_agent_sdk/advanced_memory/_storage.py | 189 ------ .../advanced_memory/_storage_backend.py | 30 - .../evaluation/_eval_session_service.py | 28 - trpc_agent_sdk/memory/__init__.py | 11 - .../memory/_advanced_memory_service.py | 113 ---- trpc_agent_sdk/runners.py | 10 - trpc_agent_sdk/sessions/__init__.py | 70 +- .../sessions/_base_session_service.py | 56 +- .../sessions/_in_memory_session_service.py | 52 +- .../sessions/_redis_session_service.py | 102 +-- trpc_agent_sdk/sessions/_session.py | 6 +- .../sessions/_sql_session_service.py | 67 +- trpc_agent_sdk/sessions/compact/__init__.py | 143 ++-- .../sessions/compact/_base_manager.py | 52 -- trpc_agent_sdk/sessions/compact/_callbacks.py | 51 -- trpc_agent_sdk/sessions/compact/_config.py | 130 ---- trpc_agent_sdk/sessions/compact/_manager.py | 135 ---- trpc_agent_sdk/sessions/compact/_runtime.py | 53 -- .../sessions/compact/advanced/__init__.py | 50 ++ .../_auto_compact.py} | 629 +++++++++--------- .../sessions/compact/advanced/_base.py | 52 ++ .../_compaction_memory_extractor.py} | 348 ++++------ .../sessions/compact/advanced/_config.py | 174 +++++ .../compact/{ => advanced}/_coordination.py | 0 .../sessions/compact/advanced/_filters.py | 80 +++ .../compact/{ => advanced}/_formats.py | 116 ---- .../compact/{ => advanced}/_history_snip.py | 96 ++- .../sessions/compact/advanced/_manager.py | 101 +++ .../_micro_compact.py} | 148 ++--- .../sessions/compact/advanced/_runtime.py | 30 + .../compact/{ => advanced}/_token_budget.py | 146 ++-- .../{ => advanced}/_tool_result_budget.py | 137 ++-- .../sessions/compact/advanced/_utils.py | 74 +++ .../sessions/compact/default/__init__.py | 34 + .../default/_checker.py} | 7 +- .../default/_summarizer.py} | 23 +- .../default}/_summarizer_manager.py | 54 +- trpc_agent_sdk/tools/__init__.py | 22 +- trpc_agent_sdk/tools/_advanced_memory_tool.py | 171 ----- 73 files changed, 2352 insertions(+), 4311 deletions(-) create mode 100644 examples/session_summarizer_with_advanced/.env create mode 100644 examples/session_summarizer_with_advanced/README.md create mode 100644 examples/session_summarizer_with_advanced/agent/__init__.py create mode 100644 examples/session_summarizer_with_advanced/agent/agent.py create mode 100644 examples/session_summarizer_with_advanced/agent/config.py create mode 100644 examples/session_summarizer_with_advanced/agent/prompts.py create mode 100644 examples/session_summarizer_with_advanced/agent/tools.py create mode 100644 examples/session_summarizer_with_advanced/run_agent.py create mode 100644 trpc_agent_sdk/abc/_compact.py delete mode 100644 trpc_agent_sdk/advanced_memory/__init__.py delete mode 100644 trpc_agent_sdk/advanced_memory/_config.py delete mode 100644 trpc_agent_sdk/advanced_memory/_integration.py delete mode 100644 trpc_agent_sdk/advanced_memory/_memory_context.py delete mode 100644 trpc_agent_sdk/advanced_memory/_paths.py delete mode 100644 trpc_agent_sdk/advanced_memory/_preload_memory.py delete mode 100644 trpc_agent_sdk/advanced_memory/_redis_stores.py delete mode 100644 trpc_agent_sdk/advanced_memory/_runtime.py delete mode 100644 trpc_agent_sdk/advanced_memory/_sql_stores.py delete mode 100644 trpc_agent_sdk/advanced_memory/_storage.py delete mode 100644 trpc_agent_sdk/advanced_memory/_storage_backend.py delete mode 100644 trpc_agent_sdk/sessions/compact/_base_manager.py delete mode 100644 trpc_agent_sdk/sessions/compact/_callbacks.py delete mode 100644 trpc_agent_sdk/sessions/compact/_config.py delete mode 100644 trpc_agent_sdk/sessions/compact/_manager.py delete mode 100644 trpc_agent_sdk/sessions/compact/_runtime.py create mode 100644 trpc_agent_sdk/sessions/compact/advanced/__init__.py rename trpc_agent_sdk/sessions/compact/{_autocompact.py => advanced/_auto_compact.py} (57%) create mode 100644 trpc_agent_sdk/sessions/compact/advanced/_base.py rename trpc_agent_sdk/sessions/compact/{_session_memory.py => advanced/_compaction_memory_extractor.py} (73%) create mode 100644 trpc_agent_sdk/sessions/compact/advanced/_config.py rename trpc_agent_sdk/sessions/compact/{ => advanced}/_coordination.py (100%) create mode 100644 trpc_agent_sdk/sessions/compact/advanced/_filters.py rename trpc_agent_sdk/sessions/compact/{ => advanced}/_formats.py (54%) rename trpc_agent_sdk/sessions/compact/{ => advanced}/_history_snip.py (81%) create mode 100644 trpc_agent_sdk/sessions/compact/advanced/_manager.py rename trpc_agent_sdk/sessions/compact/{_microcompact.py => advanced/_micro_compact.py} (67%) create mode 100644 trpc_agent_sdk/sessions/compact/advanced/_runtime.py rename trpc_agent_sdk/sessions/compact/{ => advanced}/_token_budget.py (68%) rename trpc_agent_sdk/sessions/compact/{ => advanced}/_tool_result_budget.py (77%) create mode 100644 trpc_agent_sdk/sessions/compact/advanced/_utils.py create mode 100644 trpc_agent_sdk/sessions/compact/default/__init__.py rename trpc_agent_sdk/sessions/{_summarizer_checker.py => compact/default/_checker.py} (96%) rename trpc_agent_sdk/sessions/{_session_summarizer.py => compact/default/_summarizer.py} (96%) rename trpc_agent_sdk/sessions/{ => compact/default}/_summarizer_manager.py (80%) diff --git a/examples/memory_service_with_advanced_memory/run_agent.py b/examples/memory_service_with_advanced_memory/run_agent.py index efd8014e0..e69de29bb 100644 --- a/examples/memory_service_with_advanced_memory/run_agent.py +++ b/examples/memory_service_with_advanced_memory/run_agent.py @@ -1,134 +0,0 @@ -#!/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 the two-session Advanced Memory demonstration.""" - -import asyncio -import os -from pathlib import Path - -from dotenv import load_dotenv -from trpc_agent_sdk.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(Path(__file__).with_name(".env")) - - -def create_services(agent) -> tuple[InMemorySessionService, AdvancedMemoryService]: - """Create standard Session storage with Advanced Compact and Memory.""" - memory_ttl = os.getenv("M_TTL") - session_ttl = os.getenv("SESSION_TTL") - session_ttl_seconds = int(session_ttl) if session_ttl else 0 - config = AdvancedMemoryServiceConfig( - root_dir=Path(__file__).resolve().parent, - memory_ttl_seconds=int(memory_ttl) if memory_ttl else None, - session_ttl_seconds=session_ttl_seconds or None, - memory_focus_instruction=("特别关注并主动记住用户长期稳定的兴趣爱好、" - "编程语言偏好、开发习惯和测试习惯。"), - ) - compact_config = AdvancedCompactConfig() - compact_manager = AdvancedSessionCompactManager(config=compact_config) - session_service = InMemorySessionService( - session_config=SessionServiceConfig( - ttl=SessionServiceConfig.create_ttl_config( - enable=bool(session_ttl), - ttl_seconds=session_ttl_seconds, - cleanup_interval_seconds=5, - ), - store_historical_events=True, - ), - session_compact_manager=compact_manager, - ) - return session_service, AdvancedMemoryService(config=config) - - -async def run_turn(runner, *, user_id: str, session_id: str, prompt: str) -> None: - """Run one turn and print tool activity and the final response.""" - print(f"\n👤 [{session_id}] {prompt}") - content = Content(parts=[Part.from_text(text=prompt)]) - async for event in runner.run_async( - user_id=user_id, - session_id=session_id, - new_message=content, - ): - if not event.content or not event.content.parts: - continue - for part in event.content.parts: - if part.function_call: - print(f"🔧 {part.function_call.name}({part.function_call.args})") - elif part.function_response: - print(f"📊 {part.function_response.response}") - elif part.text and not part.thought and not event.partial: - print(f"🤖 {part.text}") - - -async def main() -> None: - """Run two independent sessions sharing Advanced Memory.""" - agent = create_agent() - session_service, memory_service = create_services(agent) - - from trpc_agent_sdk.runners import Runner - runner = Runner( - app_name="advanced_memory_demo", - agent=agent, - session_service=session_service, - memory_service=memory_service, - ) - memory_ttl = os.getenv("M_TTL") - memory_ttl_seconds = int(memory_ttl) if memory_ttl else 0 - session_ttl = os.getenv("SESSION_TTL") - session_ttl_seconds = int(session_ttl) if session_ttl else 0 - try: - session_one_prompts = [ - ("Please remember that my favorite programming language is Python. " - "Save this as a user preference."), - "I use Python mainly for backend services and data processing.", - "I prefer typed Python code with clear dataclasses and small modules.", - "For testing Python code, I usually prefer pytest and focused unit tests.", - "When documenting projects, I prefer concise examples with runnable commands.", - ] - for prompt in session_one_prompts: - await run_turn( - runner, - user_id="demo-user", - session_id="session-1", - prompt=prompt, - ) - - await run_turn( - runner, - user_id="demo-user", - session_id="session-1", - prompt="Summarize what you learned about my Python development preferences.", - ) - - await run_turn( - runner, - user_id="demo-user", - session_id="session-2", - prompt="What do you remember about my favorite programming language?", - ) - - wait_seconds = max(memory_ttl_seconds, session_ttl_seconds) - if wait_seconds: - print(f"\n⏳ Waiting for TTL cleanup ({wait_seconds + 5}s)...") - await asyncio.sleep(wait_seconds + 5) - print("🧹 Expired Advanced Memory data should now be removed.") - finally: - await runner.close() - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/examples/session_service_with_advanced_memory_redis/run_agent.py b/examples/session_service_with_advanced_memory_redis/run_agent.py index 77efbaca2..828eab31e 100644 --- a/examples/session_service_with_advanced_memory_redis/run_agent.py +++ b/examples/session_service_with_advanced_memory_redis/run_agent.py @@ -14,11 +14,15 @@ 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.sessions.compact import AdvancedAutoCompactSummarizer +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerConfig +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerManager +from trpc_agent_sdk.sessions.compact import AutoCompactSummarizerConfig +from trpc_agent_sdk.sessions.compact import SessionMemoryExtractorConfig +from trpc_agent_sdk.sessions.compact import TokenContextTrackerConfig from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import Part @@ -40,17 +44,21 @@ def redis_url() -> str: return f"redis://{db_host}:{db_port}/{db_name}" -def create_compact_config() -> AdvancedCompactConfig: +def create_compact_config() -> AdvancedAutoCompactSummarizerConfig: """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, + return AdvancedAutoCompactSummarizerConfig( + token_context_tracker=TokenContextTrackerConfig( + model_context_window_tokens=4096, + max_output_tokens=256, + warning_ratio=0.25, + auto_compact_ratio=0.30, + blocking_ratio=0.95, + ), + session_memory=SessionMemoryExtractorConfig( + initial_tokens=500, + update_tokens=500, + ), + auto_compact=AutoCompactSummarizerConfig(keep_recent_contents=2), ) @@ -63,13 +71,15 @@ async def main() -> None: agent = create_agent() compact_config = create_compact_config() - compact_manager = AdvancedSessionCompactManager(config=compact_config) + compact_manager = AdvancedAutoCompactSummarizerManager( + AdvancedAutoCompactSummarizer(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, + summarizer_manager=compact_manager, ) runner = Runner( app_name=app_name, diff --git a/examples/session_service_with_advanced_memory_sql/run_agent.py b/examples/session_service_with_advanced_memory_sql/run_agent.py index ffee4b3df..5c1bb99b7 100644 --- a/examples/session_service_with_advanced_memory_sql/run_agent.py +++ b/examples/session_service_with_advanced_memory_sql/run_agent.py @@ -14,11 +14,15 @@ 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.sessions.compact import AdvancedAutoCompactSummarizer +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerConfig +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerManager +from trpc_agent_sdk.sessions.compact import AutoCompactSummarizerConfig +from trpc_agent_sdk.sessions.compact import SessionMemoryExtractorConfig +from trpc_agent_sdk.sessions.compact import TokenContextTrackerConfig from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import Part @@ -36,17 +40,21 @@ def sql_url() -> str: f"{db_host}:{db_port}/{db_name}?charset=utf8mb4") -def create_compact_config() -> AdvancedCompactConfig: +def create_compact_config() -> AdvancedAutoCompactSummarizerConfig: """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, + return AdvancedAutoCompactSummarizerConfig( + token_context_tracker=TokenContextTrackerConfig( + model_context_window_tokens=4096, + max_output_tokens=256, + warning_ratio=0.25, + auto_compact_ratio=0.30, + blocking_ratio=0.95, + ), + session_memory=SessionMemoryExtractorConfig( + initial_tokens=500, + update_tokens=500, + ), + auto_compact=AutoCompactSummarizerConfig(keep_recent_contents=2), ) @@ -59,13 +67,15 @@ async def main() -> None: agent = create_agent() compact_config = create_compact_config() - compact_manager = AdvancedSessionCompactManager(config=compact_config) + compact_manager = AdvancedAutoCompactSummarizerManager( + AdvancedAutoCompactSummarizer(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, + summarizer_manager=compact_manager, ) runner = Runner( app_name=app_name, diff --git a/examples/session_summarizer_with_advanced/.env b/examples/session_summarizer_with_advanced/.env new file mode 100644 index 000000000..dc791393a --- /dev/null +++ b/examples/session_summarizer_with_advanced/.env @@ -0,0 +1,4 @@ +# Set TRPC_AGENT_API_KEY、TRPC_AGENT_BASE_URL、TRPC_AGENT_MODEL_NAME +TRPC_AGENT_API_KEY=your-api-key +TRPC_AGENT_BASE_URL=your-base-url +TRPC_AGENT_MODEL_NAME=your-model-name diff --git a/examples/session_summarizer_with_advanced/README.md b/examples/session_summarizer_with_advanced/README.md new file mode 100644 index 000000000..77cd9d96e --- /dev/null +++ b/examples/session_summarizer_with_advanced/README.md @@ -0,0 +1,91 @@ +# Advanced Session Summarizer 示例 + +本示例验证 `AdvancedAutoCompactSummarizer` 的模型调用前压缩流程。示例执行 5 轮真实多轮对话,随着历史增长,内置 Model Filter 会在每次请求模型前检查并压缩上下文。 + +## 验证内容 + +- 模型显式配置 `AdvancedAutoCompactSummarizerFilter` +- 使用较小的字符阈值,在前几轮内触发自动压缩 +- 压缩后的模型可见窗口以 summary event 开头 +- 被替换的原始 events 移入 `historical_events` +- Session Memory 写入 `session.state` +- 压缩后对话继续进行,模型仍可依据摘要回答历史事实 + +## 为什么必须用真实多轮对话 + +压缩边界基于模型请求中的 content 数量计算(`keep_recent_contents`)。框架会把非当前 Agent 产出的历史事件转换为 user 角色,并合并相邻同角色 content。手工塞入 `author="assistant"` 的伪造事件会被合并成单条 user content,导致找不到压缩边界并报 `Not enough model contents to compact`。因此示例通过真实回合累积 user/model 交替历史。 + +## 组件关系 + +```text +OpenAIModel +└── AdvancedAutoCompactSummarizerFilter(模型调用前触发) + +InMemorySessionService +└── AdvancedAutoCompactSummarizerManager + └── AdvancedAutoCompactSummarizer +``` + +Filter 必须显式安装到模型上,Manager 不会自动修改 Agent 或 Model。 + +## 环境要求 + +- Python3.10+,推荐 Python3.12 + +## 构建步骤 + +```bash +git clone https://github.com/trpc-group/trpc-agent-python.git +cd trpc-agent-python +./build.sh +source .venv/bin/activate +``` + +## 运行步骤 + +### 配置环境变量 + +在当前目录的 `.env` 中配置(或通过 `export` 设置): + +```dotenv +TRPC_AGENT_API_KEY=your-api-key +TRPC_AGENT_BASE_URL=https://your-openai-compatible-endpoint/v1 +TRPC_AGENT_MODEL_NAME=your-model-name +``` + +### 运行命令 + +```bash +cd examples/session_summarizer_with_advanced +python3 run_agent.py +``` + +## 预期结果 + +```text +After turn 2 + active events: 4 + historical events: 0 + active window starts with summary: False + session memory persisted: True + +After turn 3 + active events: 5 + historical events: 3 + active window starts with summary: True + session memory persisted: True + +After turn 5 + active events: 5 + historical events: 10 + active window starts with summary: True + session memory persisted: True + +PASS: compaction ran on turn(s) [3, 4, 5]. +``` + +关键现象是 `active events` 稳定在一个小窗口,而 `historical events` 持续增长——说明模型可见上下文被压缩,原始事件仍可完整追溯。具体轮次和数量取决于模型回复长度。若模型未按压缩提示词返回 `` 块,示例会在结束时明确报错。 + +## 调整阈值 + +示例为便于观察而关闭 token 模式,使用 `trigger_chars=4000`。生产环境可启用 `TokenContextTrackerConfig` 并设置模型上下文窗口,或根据业务规模提高字符阈值。 diff --git a/examples/session_summarizer_with_advanced/agent/__init__.py b/examples/session_summarizer_with_advanced/agent/__init__.py new file mode 100644 index 000000000..bc6e483f9 --- /dev/null +++ b/examples/session_summarizer_with_advanced/agent/__init__.py @@ -0,0 +1,5 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. diff --git a/examples/session_summarizer_with_advanced/agent/agent.py b/examples/session_summarizer_with_advanced/agent/agent.py new file mode 100644 index 000000000..13bb361a2 --- /dev/null +++ b/examples/session_summarizer_with_advanced/agent/agent.py @@ -0,0 +1,38 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Agent used by the advanced session compaction example.""" + +from trpc_agent_sdk.agents import LlmAgent +from trpc_agent_sdk.models import LLMModel +from trpc_agent_sdk.models import OpenAIModel +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerFilter + +from .config import get_model_config +from .prompts import INSTRUCTION + + +def _create_model() -> LLMModel: + """Create a model with explicit before-model compaction filtering.""" + api_key, url, model_name = get_model_config() + return OpenAIModel( + model_name=model_name, + api_key=api_key, + base_url=url, + filters=[AdvancedAutoCompactSummarizerFilter()], + ) + + +def create_agent() -> LlmAgent: + """Create the Python tutor agent.""" + return LlmAgent( + name="python_tutor", + description="Python programming tutor that helps users learn Python", + model=_create_model(), + instruction=INSTRUCTION, + ) + + +root_agent = create_agent() diff --git a/examples/session_summarizer_with_advanced/agent/config.py b/examples/session_summarizer_with_advanced/agent/config.py new file mode 100644 index 000000000..db0d491b8 --- /dev/null +++ b/examples/session_summarizer_with_advanced/agent/config.py @@ -0,0 +1,19 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +""" Agent config module""" + +import os + + +def get_model_config() -> tuple[str, str, str]: + """Get model config from environment variables""" + api_key = os.getenv('TRPC_AGENT_API_KEY', '') + url = os.getenv('TRPC_AGENT_BASE_URL', '') + model_name = os.getenv('TRPC_AGENT_MODEL_NAME', '') + if not api_key or not url or not model_name: + raise ValueError('''TRPC_AGENT_API_KEY, TRPC_AGENT_BASE_URL, + and TRPC_AGENT_MODEL_NAME must be set in environment variables''') + return api_key, url, model_name diff --git a/examples/session_summarizer_with_advanced/agent/prompts.py b/examples/session_summarizer_with_advanced/agent/prompts.py new file mode 100644 index 000000000..c64e864f5 --- /dev/null +++ b/examples/session_summarizer_with_advanced/agent/prompts.py @@ -0,0 +1,15 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +""" prompts for agent""" + +INSTRUCTION = """You are a professional Python programming tutor. Your tasks are: +1. Answer users' Python-related questions patiently +2. Provide clear explanations and example code +3. Adjust teaching difficulty based on the user's progress +4. Encourage practice and questions + +Communicate in a friendly, professional manner. +""" diff --git a/examples/session_summarizer_with_advanced/agent/tools.py b/examples/session_summarizer_with_advanced/agent/tools.py new file mode 100644 index 000000000..16e188b49 --- /dev/null +++ b/examples/session_summarizer_with_advanced/agent/tools.py @@ -0,0 +1,6 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Tools for the advanced session compaction example.""" diff --git a/examples/session_summarizer_with_advanced/run_agent.py b/examples/session_summarizer_with_advanced/run_agent.py new file mode 100644 index 000000000..bc82216e3 --- /dev/null +++ b/examples/session_summarizer_with_advanced/run_agent.py @@ -0,0 +1,173 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Demonstrate Advanced session compaction running before each model call.""" + +from __future__ import annotations + +import asyncio +import os +import uuid +from pathlib import Path + +from dotenv import load_dotenv + +from trpc_agent_sdk.models import LLMModel +from trpc_agent_sdk.runners import Runner +from trpc_agent_sdk.sessions import InMemorySessionService +from trpc_agent_sdk.sessions import Session +from trpc_agent_sdk.sessions import SessionServiceConfig +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizer +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerConfig +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerManager +from trpc_agent_sdk.sessions.compact import AutoCompactSummarizerConfig +from trpc_agent_sdk.sessions.compact import SessionMemoryExtractorConfig +from trpc_agent_sdk.sessions.compact import TokenContextTrackerConfig +from trpc_agent_sdk.types import Content +from trpc_agent_sdk.types import Part + +load_dotenv(Path(__file__).with_name(".env")) + +SESSION_MEMORY_STATE_KEY = "_trpc_agent:summary" + +PROJECT_BRIEF = """I am building Project Apollo and I want you to remember its constraints: +- Runtime: Python 3.12 with asyncio everywhere, no blocking calls in request paths. +- Web layer: FastAPI with Pydantic models for every request and response body. +- Storage: SQLite through SQLAlchemy, and every failed write must roll back. +- Testing: pytest with async tests covering each persistence path. +- Style: full type hints, small functions, no bare except. +- Deployment: Linux containers, released every Friday afternoon. +""" + +# Multi-turn conversation. Real turns are what build alternating user/model +# history, which is the shape the compaction boundary is computed from. +CONVERSATIONS = ( + PROJECT_BRIEF + "\nAcknowledge the constraints and outline the module layout you would use.", + "Show me the SQLAlchemy session helper for Project Apollo, " + "including how rollback is handled on a failed write.", + "Now show the pytest fixtures and one async test that proves the rollback path works.", + "Explain how I should structure the FastAPI routers and dependency injection for this project.", + "Recap Project Apollo: its stack, storage rules, testing rules, and release cadence.", +) + + +def create_compact_config() -> AdvancedAutoCompactSummarizerConfig: + """Use small character budgets so the example compacts within a few turns.""" + return AdvancedAutoCompactSummarizerConfig( + # Character thresholds keep the demo independent of any model's + # context window. Production setups usually enable the token tracker. + token_context_tracker=TokenContextTrackerConfig(enabled=False), + session_memory=SessionMemoryExtractorConfig( + initial_chars=1_000, + update_chars=500, + ), + auto_compact=AutoCompactSummarizerConfig( + trigger_chars=4_000, + target_chars=2_000, + blocking_chars=20_000, + keep_recent_contents=2, + summary_input_max_chars=12_000, + ), + ) + + +def create_summarizer_manager(model: LLMModel) -> AdvancedAutoCompactSummarizerManager: + """Create the Advanced summarizer that the model filter drives.""" + summarizer = AdvancedAutoCompactSummarizer( + config=create_compact_config(), + model=model, + ) + return AdvancedAutoCompactSummarizerManager(summarizer=summarizer) + + +def print_session_state(label: str, session: Session) -> None: + """Print the state needed to verify compaction behavior.""" + summary_anchor = bool(session.events and session.events[0].is_summary_event()) + print(f"\n{label}") + print(f" active events: {len(session.events)}") + print(f" historical events: {len(session.historical_events)}") + print(f" active window starts with summary: {summary_anchor}") + print(f" session memory persisted: {SESSION_MEMORY_STATE_KEY in session.state}") + + +async def run_turn(runner: Runner, user_id: str, session_id: str, prompt: str) -> None: + """Send one user message and stream the answer.""" + print(f"\nUser: {prompt.splitlines()[0]}") + print("Assistant: ", end="", flush=True) + async for event in runner.run_async( + user_id=user_id, + session_id=session_id, + new_message=Content(parts=[Part.from_text(text=prompt)]), + ): + if not event.content or not event.content.parts: + continue + for part in event.content.parts: + if part.text and not part.thought: + print(part.text, end="" if event.partial else "\n", flush=True) + + +async def main() -> None: + """Run a multi-turn conversation and verify before-model Advanced AutoCompact.""" + app_name = "advanced-session-summarizer-demo" + user_id = "demo-user" + session_id = os.getenv("SESSION_ID", str(uuid.uuid4())) + + # Import after load_dotenv so the module-level Agent can read model settings. + from agent.agent import root_agent + + manager = create_summarizer_manager(root_agent.model) + session_service = InMemorySessionService( + summarizer_manager=manager, + session_config=SessionServiceConfig(store_historical_events=True), + ) + runner = Runner( + app_name=app_name, + agent=root_agent, + session_service=session_service, + ) + + try: + await session_service.create_session( + app_name=app_name, + user_id=user_id, + session_id=session_id, + ) + print(f"Session: {app_name}/{user_id}/{session_id}") + + compaction_turns: list[int] = [] + historical_count = 0 + for index, prompt in enumerate(CONVERSATIONS, start=1): + await run_turn(runner, user_id, session_id, prompt) + stored = await session_service.get_session( + app_name=app_name, + user_id=user_id, + session_id=session_id, + ) + if stored is None: + raise RuntimeError("Session disappeared during the conversation") + print_session_state(f"After turn {index}", stored) + if len(stored.historical_events) > historical_count: + compaction_turns.append(index) + historical_count = len(stored.historical_events) + + if not compaction_turns: + raise RuntimeError( + "Compaction never ran. Confirm that the model installs " + "AdvancedAutoCompactSummarizerFilter, that the conversation " + "exceeds auto_compact.trigger_chars, and that the model returns " + "the requested block." + ) + if not stored.events[0].is_summary_event(): + raise RuntimeError("Compaction ran but the active window lost its summary anchor") + + print(f"\nPASS: compaction ran on turn(s) {compaction_turns}.") + print(f" archived events: {len(stored.historical_events)}") + print(f" model-visible events: {len(stored.events)}") + finally: + await runner.close() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/sessions/compact/test_coordination.py b/tests/sessions/compact/test_coordination.py index 2b2fb7ee1..3dba41f86 100644 --- a/tests/sessions/compact/test_coordination.py +++ b/tests/sessions/compact/test_coordination.py @@ -6,7 +6,7 @@ import pytest -from trpc_agent_sdk.sessions.compact._coordination import CrossLoopLock +from trpc_agent_sdk.sessions.compact.advanced._coordination import CrossLoopLock @pytest.mark.asyncio diff --git a/tests/sessions/compact/test_session_compact.py b/tests/sessions/compact/test_session_compact.py index f8de41c30..f9ef740ef 100644 --- a/tests/sessions/compact/test_session_compact.py +++ b/tests/sessions/compact/test_session_compact.py @@ -3,34 +3,60 @@ from __future__ import annotations from types import SimpleNamespace +from unittest.mock import AsyncMock import pytest +from trpc_agent_sdk.abc import CompactSummarizerABC +from trpc_agent_sdk.abc import CompactTrigger +from trpc_agent_sdk.context import new_agent_context from trpc_agent_sdk.events import Event from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.models import LlmResponse from trpc_agent_sdk.sessions import InMemorySessionService from trpc_agent_sdk.sessions import SessionServiceConfig -from trpc_agent_sdk.sessions.compact import 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 AdvancedAutoCompactSummarizer +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerConfig +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerManager +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerRuntime +from trpc_agent_sdk.sessions.compact import AutoCompactSummarizerConfig from trpc_agent_sdk.sessions.compact import SessionMemoryExtractor -from trpc_agent_sdk.sessions.compact import ToolResultBudget +from trpc_agent_sdk.sessions.compact.advanced import SessionMemoryExtractorConfig +from trpc_agent_sdk.sessions.compact.advanced import ToolResultBudgetConfig +from trpc_agent_sdk.sessions.compact.advanced._tool_result_budget import ToolResultBudget from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import FunctionResponse from trpc_agent_sdk.types import Part -class _MemoryGenerator: +class _MemoryModel: + name = "memory-model" - async def generate(self, extraction_input, ctx) -> SessionMemoryDocument: - del ctx - return SessionMemoryDocument( - session_title="Test session", - current_state=extraction_input.last_event_id, + async def generate_async(self, request, *, stream, ctx): + del request, stream, ctx + yield LlmResponse( + content=Content( + role="model", + parts=[Part.from_text(text="# Session Title\nTest session\n\n# Current State\nUpdated")], + ), ) +class _SummaryModel: + name = "summary-model" + + async def generate_async(self, request, *, stream, ctx): + del request, stream, ctx + yield LlmResponse(content=Content( + role="model", + parts=[ + Part.from_text( + text="covered" + "# Session Title\nCompact summary") + ], + )) + + def _event(event_id: str, content: Content) -> Event: return Event( id=event_id, @@ -41,22 +67,121 @@ def _event(event_id: str, content: Content) -> Event: def test_compact_runtime_has_no_external_storage() -> None: - runtime = SessionCompactRuntime.create(AdvancedCompactConfig()) + runtime = AdvancedAutoCompactSummarizerRuntime(AdvancedAutoCompactSummarizerConfig()) assert not hasattr(runtime, "transcripts") assert not hasattr(runtime, "tool_results") assert not hasattr(runtime, "paths") +@pytest.mark.asyncio +async def test_disabled_auto_compact_returns_without_recursion() -> None: + config = AdvancedAutoCompactSummarizerConfig( + auto_compact=AutoCompactSummarizerConfig(enabled=False), + ) + summarizer = AdvancedAutoCompactSummarizer(config) + session = SimpleNamespace(id="session", app_name="app", user_id="user") + ctx = SimpleNamespace(session=session, session_id=session.id) + + result = await summarizer.apply(LlmRequest(), ctx=ctx) + + assert result.compacted is False + assert result.blocked is False + + @pytest.mark.asyncio async def test_session_service_accepts_a_configured_compact_manager() -> None: - manager = AdvancedSessionCompactManager(config=AdvancedCompactConfig()) + summarizer = AdvancedAutoCompactSummarizer(AdvancedAutoCompactSummarizerConfig()) + manager = AdvancedAutoCompactSummarizerManager(summarizer) + service = InMemorySessionService( + session_config=SessionServiceConfig(store_historical_events=True), + summarizer_manager=manager, + ) + + assert service.summarizer_manager is manager + await service.close() + + +@pytest.mark.asyncio +async def test_advanced_summarizer_implements_compact_abc_and_timing() -> None: + summarizer = AdvancedAutoCompactSummarizer(AdvancedAutoCompactSummarizerConfig( + session_memory=SessionMemoryExtractorConfig(enabled=False), + )) + before_model_manager = AdvancedAutoCompactSummarizerManager(summarizer) + after_turn_manager = AdvancedAutoCompactSummarizerManager( + summarizer, + compact_trigger=CompactTrigger.AFTER_TURN, + ) + + assert isinstance(summarizer, CompactSummarizerABC) + assert before_model_manager.compact_trigger == CompactTrigger.BEFORE_MODEL + assert after_turn_manager.compact_trigger == CompactTrigger.AFTER_TURN + + session = SimpleNamespace(events=[]) + ctx = SimpleNamespace(session=session) + summarizer.should_summarize = AsyncMock(return_value=True) + summarizer.create_session_summary = AsyncMock(return_value="summary") + await before_model_manager.create_session_summary(session, ctx=ctx) + summarizer.should_summarize.assert_not_awaited() + + await after_turn_manager.create_session_summary(session, ctx=ctx) + summarizer.should_summarize.assert_awaited_once_with(session) + summarizer.create_session_summary.assert_awaited_once_with( + session, + ctx=ctx, + store_historical_events=True, + ) + + +@pytest.mark.asyncio +async def test_advanced_end_of_turn_compaction_moves_old_events_to_history() -> None: service = InMemorySessionService( session_config=SessionServiceConfig(store_historical_events=True), - session_compact_manager=manager, + ) + session = await service.create_session( + app_name="app", + user_id="user", + session_id="session", + ) + for index in range(3): + await service.append_event( + session, + _event( + f"event-{index}", + Content(role="user", parts=[Part.from_text(text=f"{index}:" + "x" * 2_000)]), + ), + ) + original_ids = [event.id for event in session.events] + config = AdvancedAutoCompactSummarizerConfig( + session_memory=SessionMemoryExtractorConfig(enabled=False), + auto_compact=AutoCompactSummarizerConfig( + trigger_chars=500, + target_chars=250, + blocking_chars=10_000, + keep_recent_contents=1, + ), + ) + summarizer = AdvancedAutoCompactSummarizer(config, model=_SummaryModel()) + ctx = SimpleNamespace( + session=session, + session_id=session.id, + session_service=service, + agent_context=new_agent_context(), + agent=SimpleNamespace(model=_SummaryModel()), ) - assert service.session_compact_manager is manager + assert await summarizer.should_summarize(session) is True + summary = await summarizer.create_session_summary( + session, + ctx=ctx, + store_historical_events=True, + ) + + assert summary is not None and "Compact summary" in summary + assert len(session.events) == 2 + assert session.events[0].is_summary_event() + assert session.events[1].id == original_ids[-1] + assert [event.id for event in session.historical_events] == original_ids[:-1] await service.close() @@ -72,24 +197,36 @@ async def test_session_memory_is_written_to_session_state() -> None: session, _event("event-1", Content(parts=[Part.from_text(text="hello")])), ) + config = AdvancedAutoCompactSummarizerConfig( + session_memory=SessionMemoryExtractorConfig( + initial_chars=1, + update_chars=1, + ), + ) extractor = SessionMemoryExtractor( - SessionCompactRuntime.create( - AdvancedCompactConfig( - session_memory_initial_chars=1, - session_memory_update_chars=1, - )), - _MemoryGenerator(), - session_service=service, + AdvancedAutoCompactSummarizerRuntime(config), + model=_MemoryModel(), ) result = await extractor.extract_if_needed( - session, - SimpleNamespace(session=session, agent=SimpleNamespace(model="test")), + SimpleNamespace( + session=session, + session_service=service, + agent=SimpleNamespace(model=_MemoryModel(), generate_content_config=None), + override_messages=None, + ), force=True, ) assert result.extracted is True assert "_trpc_agent:summary" in session.state + stored = await service.get_session( + app_name="app", + user_id="user", + session_id="session", + ) + assert stored is not None + assert "_trpc_agent:summary" in stored.state await service.close() @@ -110,15 +247,18 @@ async def test_tool_result_budget_keeps_the_session_event_id() -> None: ]) await service.append_event(session, _event("event-tool", content)) request = LlmRequest(model="test", contents=[content.model_copy(deep=True)]) + config = AdvancedAutoCompactSummarizerConfig( + tool_result_budget=ToolResultBudgetConfig( + max_chars=100, + preview_chars=20, + ), + ) budget = ToolResultBudget( - SessionCompactRuntime.create(AdvancedCompactConfig( - tool_result_max_chars=100, - tool_result_preview_chars=20, - ))) + AdvancedAutoCompactSummarizerRuntime(config), + ) await budget.apply( request, - session_id=session.id, ctx=SimpleNamespace(session=session), ) diff --git a/tests/sessions/compact/test_token_budget.py b/tests/sessions/compact/test_token_budget.py index 9e2c918a6..ead1cd6cc 100644 --- a/tests/sessions/compact/test_token_budget.py +++ b/tests/sessions/compact/test_token_budget.py @@ -4,9 +4,9 @@ from types import SimpleNamespace -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.sessions.compact.advanced import TokenContextTrackerConfig +from trpc_agent_sdk.sessions.compact.advanced._token_budget import TokenContextTracker from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import Part @@ -37,7 +37,7 @@ def test_usage_baseline_adds_only_contents_after_matching_event(tmp_path) -> Non agent=SimpleNamespace(model="test-model"), ) tracker = TokenContextTracker( - AdvancedCompactConfig( + TokenContextTrackerConfig( enabled=True, model_context_window_tokens=1_000, max_output_tokens=100, @@ -62,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(AdvancedCompactConfig(enabled=True)) + tracker = TokenContextTracker(TokenContextTrackerConfig(enabled=True)) estimate = tracker.estimate(request, ctx) @@ -84,7 +84,7 @@ def test_changed_recorded_system_or_tool_fingerprint_falls_back(tmp_path) -> Non agent=SimpleNamespace(model="test-model"), ) - estimate = TokenContextTracker(AdvancedCompactConfig(enabled=True)).estimate(request, ctx) + estimate = TokenContextTracker(TokenContextTrackerConfig(enabled=True)).estimate(request, ctx) assert estimate.source == "estimated" assert estimate.tokens < 999_999 @@ -93,23 +93,34 @@ def test_changed_recorded_system_or_tool_fingerprint_falls_back(tmp_path) -> Non def test_budget_reserves_max_output_and_calculates_three_thresholds(tmp_path) -> None: """Ensure thresholds use the window after reserving max output.""" tracker = TokenContextTracker( - AdvancedCompactConfig( + TokenContextTrackerConfig( enabled=True, model_context_window_tokens=10_000, max_output_tokens=2_000, )) - budget = tracker.budget(_request("测试请求")) + ctx = SimpleNamespace( + session=SimpleNamespace(events=[]), + agent=SimpleNamespace(model="test-model"), + ) + budget = tracker.budget(_request("测试请求"), ctx) assert budget.effective_window_tokens == 8_000 assert budget.warning_threshold_tokens == 6_800 - assert budget.autocompact_threshold_tokens == 7_200 + assert budget.auto_compact_threshold_tokens == 7_200 assert budget.blocking_threshold_tokens == 7_600 def test_no_window_keeps_compatibility_mode(tmp_path) -> None: """Ensure token decisions remain disabled without a model window.""" - budget = TokenContextTracker(AdvancedCompactConfig(enabled=True)).budget(_request("compatibility request")) + ctx = SimpleNamespace( + session=SimpleNamespace(events=[]), + agent=SimpleNamespace(model="test-model"), + ) + budget = TokenContextTracker(TokenContextTrackerConfig(enabled=True)).budget( + _request("compatibility request"), + ctx, + ) assert not budget.token_mode_enabled assert budget.estimate.source == "estimated" diff --git a/tests/sessions/replay/backends.py b/tests/sessions/replay/backends.py index 5c638a0fe..c5b5a7462 100644 --- a/tests/sessions/replay/backends.py +++ b/tests/sessions/replay/backends.py @@ -25,7 +25,9 @@ from trpc_agent_sdk.sessions import SessionServiceConfig from trpc_agent_sdk.sessions import SessionSummarizer from trpc_agent_sdk.sessions import SqlSessionService -from trpc_agent_sdk.sessions._summarizer_manager import SummarizerSessionManager +from trpc_agent_sdk.sessions.compact.default._summarizer_manager import ( + DefaultSessionSummarizerManager as SummarizerSessionManager, +) from .harness import ReplayBackend from .report import BackendStatus diff --git a/tests/sessions/test_base_session_service.py b/tests/sessions/test_base_session_service.py index bcb63037a..bc6fb733f 100644 --- a/tests/sessions/test_base_session_service.py +++ b/tests/sessions/test_base_session_service.py @@ -13,15 +13,15 @@ from __future__ import annotations import time -from unittest.mock import AsyncMock, MagicMock, patch - -import pytest +from unittest.mock import AsyncMock, MagicMock from trpc_agent_sdk.abc import ListSessionsResponse from trpc_agent_sdk.events import Event from trpc_agent_sdk.sessions._base_session_service import BaseSessionService from trpc_agent_sdk.sessions._session import Session -from trpc_agent_sdk.sessions._summarizer_manager import SummarizerSessionManager +from trpc_agent_sdk.sessions.compact.default._summarizer_manager import ( + DefaultSessionSummarizerManager as SummarizerSessionManager, +) from trpc_agent_sdk.sessions._types import SessionServiceConfig from trpc_agent_sdk.types import Content, EventActions, Part, State @@ -344,6 +344,16 @@ async def test_update_session_default_noop(self): session = _make_session() await svc.update_session(session) + async def test_update_session_state_falls_back_to_full_update(self): + svc = ConcreteSessionService() + svc.update_session = AsyncMock() + session = _make_session() + + await svc.update_session_state(session, {"summary": {"version": 1}}) + + assert session.state["summary"] == {"version": 1} + svc.update_session.assert_awaited_once_with(session) + async def test_close(self): svc = ConcreteSessionService() await svc.close() diff --git a/tests/sessions/test_in_memory_session_service.py b/tests/sessions/test_in_memory_session_service.py index 51b14af33..08e93953b 100644 --- a/tests/sessions/test_in_memory_session_service.py +++ b/tests/sessions/test_in_memory_session_service.py @@ -429,7 +429,7 @@ async def test_patch_state_preserves_stored_events(self): stale = session.model_copy(deep=True) stale.events = [] - await svc.patch_session_state(stale, {"_trpc_agent:summary": {"v": 1}}) + await svc.update_session_state(stale, {"_trpc_agent:summary": {"v": 1}}) stored = await svc.get_session( app_name="app", diff --git a/tests/sessions/test_redis_session_service.py b/tests/sessions/test_redis_session_service.py index e5f154bf5..a2e1f49d2 100644 --- a/tests/sessions/test_redis_session_service.py +++ b/tests/sessions/test_redis_session_service.py @@ -377,7 +377,7 @@ async def test_patch_state_preserves_stored_events(self): stale = session.model_copy(deep=True) stale.events = [] - await svc.patch_session_state(stale, {"_trpc_agent:summary": {"v": 1}}) + await svc.update_session_state(stale, {"_trpc_agent:summary": {"v": 1}}) stored = await svc.get_session( app_name="app", @@ -409,7 +409,7 @@ async def test_patch_state_repairs_lua_empty_array_encoding(self): assert loaded is not None assert loaded.historical_events == [] - await svc.patch_session_state(loaded, {"_trpc_agent:summary": {"v": 1}}) + await svc.update_session_state(loaded, {"_trpc_agent:summary": {"v": 1}}) assert loaded.state["_trpc_agent:summary"] == {"v": 1} await svc.close() diff --git a/tests/sessions/test_session_summarizer.py b/tests/sessions/test_session_summarizer.py index ef46db4bb..472d31676 100644 --- a/tests/sessions/test_session_summarizer.py +++ b/tests/sessions/test_session_summarizer.py @@ -21,10 +21,10 @@ from trpc_agent_sdk.events import Event from trpc_agent_sdk.sessions._session import Session -from trpc_agent_sdk.sessions._session_summarizer import ( +from trpc_agent_sdk.sessions.compact.default._summarizer import ( DEFAULT_SUMMARIZER_PROMPT, - SessionSummarizer, - SessionSummary, + DefaultSessionSummarizer as SessionSummarizer, + DefaultSessionSummary as SessionSummary, ) from trpc_agent_sdk.types import Content, EventActions, FunctionCall, FunctionResponse, Part diff --git a/tests/sessions/test_sql_session_service.py b/tests/sessions/test_sql_session_service.py index 1eaa27853..8a0d7c136 100644 --- a/tests/sessions/test_sql_session_service.py +++ b/tests/sessions/test_sql_session_service.py @@ -461,7 +461,7 @@ async def test_patch_state_preserves_stored_events(self): stale = session.model_copy(deep=True) stale.events = [] - await svc.patch_session_state(stale, {"_trpc_agent:summary": {"v": 1}}) + await svc.update_session_state(stale, {"_trpc_agent:summary": {"v": 1}}) stored = await svc.get_session( app_name="app", diff --git a/tests/sessions/test_summarizer_checker.py b/tests/sessions/test_summarizer_checker.py index 62613ed40..8d12360ef 100644 --- a/tests/sessions/test_summarizer_checker.py +++ b/tests/sessions/test_summarizer_checker.py @@ -24,7 +24,7 @@ from trpc_agent_sdk.events import Event from trpc_agent_sdk.sessions._session import Session -from trpc_agent_sdk.sessions._summarizer_checker import ( +from trpc_agent_sdk.sessions.compact.default._checker import ( set_summarizer_check_functions_by_and, set_summarizer_check_functions_by_or, set_summarizer_conversation_threshold, @@ -284,7 +284,7 @@ def test_above_threshold(self): checker = set_summarizer_conversation_threshold(10) session = _make_session(conversation_count=15) assert checker(session) is True - assert session.conversation_count == 0 + assert session.conversation_count == 15 def test_below_threshold(self): checker = set_summarizer_conversation_threshold(10) @@ -301,12 +301,12 @@ def test_default_threshold(self): session = _make_session(conversation_count=101) assert checker(session) is True - def test_resets_count_on_true(self): + def test_does_not_mutate_count(self): checker = set_summarizer_conversation_threshold(5) session = _make_session(conversation_count=10) result = checker(session) assert result is True - assert session.conversation_count == 0 + assert session.conversation_count == 10 class TestCheckFunctionsByAnd: diff --git a/tests/sessions/test_summarizer_manager.py b/tests/sessions/test_summarizer_manager.py index fa3292983..0086f5aed 100644 --- a/tests/sessions/test_summarizer_manager.py +++ b/tests/sessions/test_summarizer_manager.py @@ -20,8 +20,11 @@ from trpc_agent_sdk.events import Event from trpc_agent_sdk.sessions._session import Session -from trpc_agent_sdk.sessions._session_summarizer import SessionSummarizer, SessionSummary -from trpc_agent_sdk.sessions._summarizer_manager import SummarizerSessionManager +from trpc_agent_sdk.sessions.compact.default._summarizer import DefaultSessionSummarizer as SessionSummarizer +from trpc_agent_sdk.sessions.compact.default._summarizer import DefaultSessionSummary as SessionSummary +from trpc_agent_sdk.sessions.compact.default._summarizer_manager import ( + DefaultSessionSummarizerManager as SummarizerSessionManager, +) from trpc_agent_sdk.types import Content, Part @@ -139,11 +142,12 @@ async def test_summary_when_should_summarize(self): mock_service = AsyncMock() manager.set_session_service(mock_service) - session = _make_session(events=[_make_event()]) + session = _make_session(events=[_make_event()], conversation_count=15) await manager.create_session_summary(session) manager._summarizer.create_session_summary.assert_called_once() mock_service.update_session.assert_called_once() + assert session.conversation_count == 0 async def test_no_summary_when_should_not_summarize(self): model = _make_model() diff --git a/trpc_agent_sdk/abc/__init__.py b/trpc_agent_sdk/abc/__init__.py index 5a3f75125..f9e3cf650 100644 --- a/trpc_agent_sdk/abc/__init__.py +++ b/trpc_agent_sdk/abc/__init__.py @@ -16,6 +16,9 @@ from ._artifact_service import ArtifactId from ._artifact_service import ArtifactServiceABC from ._artifact_service import ArtifactVersion +from ._compact import CompactSummarizerABC +from ._compact import CompactSummarizerManagerABC +from ._compact import CompactTrigger from ._filter import FilterABC from ._filter import FilterAsyncGenHandleType from ._filter import FilterAsyncGenReturnType @@ -43,6 +46,9 @@ "ArtifactId", "ArtifactServiceABC", "ArtifactVersion", + "CompactSummarizerABC", + "CompactSummarizerManagerABC", + "CompactTrigger", "FilterABC", "FilterAsyncGenHandleType", "FilterAsyncGenReturnType", diff --git a/trpc_agent_sdk/abc/_compact.py b/trpc_agent_sdk/abc/_compact.py new file mode 100644 index 000000000..b96db4808 --- /dev/null +++ b/trpc_agent_sdk/abc/_compact.py @@ -0,0 +1,187 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""The base class for compact summarizers.""" + +from abc import ABC +from abc import abstractmethod +from enum import Enum +from typing import List +from typing import Optional +from typing import Dict +from typing import Any +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from trpc_agent_sdk.context import InvocationContext + +from ._session import SessionABC +from ._response import ResponseABC +from ._request import RequestABC +from ._session_service import SessionServiceABC + + +class CompactTrigger(str, Enum): + """Select when session compaction is evaluated.""" + + AFTER_TURN = "after_turn" + BEFORE_MODEL = "before_model" + + +class CompactSummarizerABC(ABC): + """The base class for compact summarizers.""" + + @abstractmethod + async def should_summarize(self, session: SessionABC) -> bool: + """Check if the session should be summarized. + + Args: + session: The session to check. + + Returns: + True if the session should be summarized, False otherwise. + """ + + @abstractmethod + async def create_session_summary_by_events( + self, + events: List[ResponseABC], + session_id: str, + keep_recent_count: int = 10, + ctx: Optional["InvocationContext"] = None, + historical_events: Optional[List[ResponseABC]] = None, + store_historical_events: bool = False) -> tuple[Optional[str], List[ResponseABC]]: + """Create a session summary by events. + + Args: + events: The events to summarize. + session_id: The session ID. + keep_recent_count: The number of recent events to keep. + ctx: The invocation context. + historical_events: The historical events. + store_historical_events: Whether to store the historical events. + + Returns: + A tuple containing the session summary and the historical events. + """ + + @abstractmethod + async def create_session_summary(self, + session: SessionABC, + ctx: Optional["InvocationContext"] = None, + store_historical_events: bool = False) -> Optional[str]: + """Create a session summary. + + Args: + session: The session to summarize. + ctx: The invocation context. + store_historical_events: Whether to store the historical events. + + Returns: + The session summary. + """ + + def get_summary_metadata(self) -> Dict[str, Any]: + """Get the summary metadata. + + Returns: + The summary metadata. + """ + return {} + + async def create_session_summary_by_request( + self, + request: RequestABC, + ctx: Optional["InvocationContext"] = None, + force: bool = False, + ) -> Optional[ResponseABC]: + """Compact one model request before generation. + + The default implementation is intentionally a no-op so existing + summarizers only implementing end-of-turn compaction remain + compatible. + """ + del request, ctx, force + return None + + +class CompactSummarizerManagerABC(ABC): + """Coordinate one CompactSummarizer implementation with a SessionService.""" + + def __init__( + self, + summarizer: CompactSummarizerABC, + compact_trigger: CompactTrigger = CompactTrigger.AFTER_TURN, + ): + self._summarizer = summarizer + self._base_service = None + self._compact_trigger = compact_trigger + + @property + def summarizer(self) -> CompactSummarizerABC: + """Get the CompactSummarizer implementation.""" + return self._summarizer + + @property + def session_service(self) -> SessionServiceABC: + """Get the base session service.""" + return self._base_service + + @property + def compact_trigger(self) -> CompactTrigger: + """Return when this manager evaluates compaction.""" + return self._compact_trigger + + def set_session_service(self, session_service: SessionServiceABC, force: bool = False) -> None: + """Set the session service to use. + + Args: + session_service: The session service to use. + force: Whether to force update even if already set. + """ + if not self._base_service or force: + self._base_service = session_service + + def set_summarizer(self, summarizer: CompactSummarizerABC, force: bool = False) -> None: + """Set the summarizer to use. + + Args: + summarizer: The summarizer to use + force: Whether to force update even if already set + """ + if not self._summarizer or force: + self._summarizer = summarizer + + @abstractmethod + async def create_session_summary( + self, + session: SessionABC, + force: bool = False, + ctx: Optional["InvocationContext"] = None, + ) -> None: + """Update compact state through the SessionService post-turn hook.""" + + @abstractmethod + async def get_session_summary(self, session: SessionABC) -> Optional[str]: + """Return the compact representation exposed as a session summary.""" + + async def create_session_summary_before_model( + self, + request: RequestABC, + ctx: "InvocationContext", + force: bool = False, + ) -> Optional[ResponseABC]: + """Run request compaction when configured for the before-model phase.""" + if self._compact_trigger != CompactTrigger.BEFORE_MODEL: + return None + return await self._summarizer.create_session_summary_by_request( + request, + ctx=ctx, + force=force, + ) + + async def close(self) -> None: + """Release resources owned by this manager.""" + return None diff --git a/trpc_agent_sdk/abc/_session_service.py b/trpc_agent_sdk/abc/_session_service.py index d26b7f397..f1cbd1923 100644 --- a/trpc_agent_sdk/abc/_session_service.py +++ b/trpc_agent_sdk/abc/_session_service.py @@ -125,18 +125,19 @@ async def update_session(self, session: SessionABC) -> None: session: The session to update """ - async def patch_session_state( + async def update_session_state( self, session: SessionABC, state_delta: dict[str, Any], ) -> None: - """Atomically merge session-scoped state without replacing Events. + """Persist a session-scoped state delta. - 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. + Backends may override this method with an efficient partial update. + The default implementation preserves compatibility with existing + SessionService implementations by falling back to ``update_session``. """ - raise NotImplementedError(f"{type(self).__name__} does not support atomic session state patches") + session.state.update(state_delta) + await self.update_session(session) @abstractmethod async def create_session_summary(self, session: SessionABC, ctx: "InvocationContext" = None) -> None: diff --git a/trpc_agent_sdk/advanced_memory/__init__.py b/trpc_agent_sdk/advanced_memory/__init__.py deleted file mode 100644 index 0f7fe7ca5..000000000 --- a/trpc_agent_sdk/advanced_memory/__init__.py +++ /dev/null @@ -1,57 +0,0 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. -# -# Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Optional long-term memory APIs.""" - -from ._config import AdvancedMemoryServiceConfig -from trpc_agent_sdk.sessions.compact._formats import MemoryDocument -from trpc_agent_sdk.sessions.compact._formats import MemoryIndexEntry -from trpc_agent_sdk.sessions.compact._formats import MemoryType -from trpc_agent_sdk.sessions.compact._formats import memory_freshness -from trpc_agent_sdk.sessions.compact._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 -from ._storage_backend import AdvancedMemoryStorageBackend -from ._storage_backend import LocalAdvancedMemoryStorageBackend - -__all__ = [ - "AdvancedMemoryStorageBackend", - "AdvancedMemoryServiceConfig", - "LongTermMemoryIntegration", - "AdvancedMemoryPaths", - "AdvancedMemoryRuntime", - "ScopedAdvancedMemoryRuntime", - "LongTermMemoryStore", - "LocalAdvancedMemoryStorageBackend", - "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/advanced_memory/_config.py b/trpc_agent_sdk/advanced_memory/_config.py deleted file mode 100644 index 99588568c..000000000 --- a/trpc_agent_sdk/advanced_memory/_config.py +++ /dev/null @@ -1,83 +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 - -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/advanced_memory/_integration.py b/trpc_agent_sdk/advanced_memory/_integration.py deleted file mode 100644 index 3c01e02f6..000000000 --- a/trpc_agent_sdk/advanced_memory/_integration.py +++ /dev/null @@ -1,106 +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 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/advanced_memory/_memory_context.py deleted file mode 100644 index b3e62a9df..000000000 --- a/trpc_agent_sdk/advanced_memory/_memory_context.py +++ /dev/null @@ -1,126 +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. -"""Inject the long-term memory index into model system instructions.""" - -from __future__ import annotations - -from typing import TYPE_CHECKING - -from trpc_agent_sdk.sessions.compact._callbacks import install_staged_callback -from ._runtime import AdvancedMemoryRuntime - -if TYPE_CHECKING: - from trpc_agent_sdk.agents import LlmAgent - from trpc_agent_sdk.context import InvocationContext - from trpc_agent_sdk.models import LlmRequest - -LONG_TERM_MEMORY_MARKER = "" - - -class LongTermMemoryContext: - """Load a bounded MEMORY.md index for each model request.""" - - def __init__(self, memory_runtime: AdvancedMemoryRuntime) -> None: - """Store the runtime bound to this long-term memory context.""" - self._runtime = memory_runtime - - @property - def runtime(self) -> AdvancedMemoryRuntime: - """Return the runtime bound to this long-term memory context.""" - return self._runtime - - async def apply(self, request: "LlmRequest", ctx: "InvocationContext | None" = None) -> bool: - """Append the MEMORY.md index and on-demand read guidance.""" - 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 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 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 " - "lead, not a complete fact. Use it only when relevant to the current task. For exact details, prefer " - "read_memory on the referenced file; do not infer details from one index line.\n\n" - "Memory records are point-in-time observations and may become stale. Before relying on a memory for " - "current code, configuration, or external state, verify it against the current source or resource. " - "If a memory is incorrect or outdated, update the existing memory instead of creating a duplicate.\n\n" - "## Proactively maintain long-term memory\n" - "If the save_memory tool is available, proactively save information that is sufficiently certain and " - "useful across sessions; do not wait for the user to say \"remember this\". Prefer saving:\n" - "- user: stable identity, role, preferences, skill level, work habits, or explicit personal constraints;\n" - "- feedback: corrections, confirmations, or preferences about collaboration, format, and quality;\n" - "- project: goals, confirmed technical decisions, architecture/process conventions, important state, or " - "deadlines that cannot be reliably inferred from code alone;\n" - "- reference: locations, purposes, and usage constraints for external systems, docs, APIs, repositories, " - "or resources.\n" - "For corrections, replacements, or important additions, update the existing memory with the same " - "filename instead of creating a duplicate. Keep each memory focused on one stable, concrete, actionable " - "topic; inspect the index first and reuse an existing topic when possible.\n\n" - "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 " - 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: " - 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]) - return True - - -class LongTermMemoryContextCallback: - """Adapt the long-term memory index injector to before_model_callback.""" - - advanced_memory_stage = 5 - - def __init__(self, memory_context: LongTermMemoryContext) -> None: - """Store the injector executed before each model request.""" - self._memory_context = memory_context - - @property - def memory_context(self) -> LongTermMemoryContext: - """Return the memory context used by this callback.""" - return self._memory_context - - async def __call__(self, ctx: "InvocationContext", request: "LlmRequest") -> None: - """Inject the long-term memory index before a model request.""" - await self._memory_context.apply(request, ctx) - return None - - -def setup_long_term_memory_context( - agent: "LlmAgent", - memory_runtime: AdvancedMemoryRuntime, -) -> LongTermMemoryContext: - """Install the index callback while preserving pipeline stage order.""" - memory_context = LongTermMemoryContext(memory_runtime) - callback = LongTermMemoryContextCallback(memory_context) - existing_context = install_staged_callback( - agent, - callback, - callback_type=LongTermMemoryContextCallback, - component_attribute="memory_context", - memory_runtime=memory_runtime, - conflict_message="Long-term memory context is already configured with another runtime", - ) - return existing_context or memory_context diff --git a/trpc_agent_sdk/advanced_memory/_paths.py b/trpc_agent_sdk/advanced_memory/_paths.py deleted file mode 100644 index 768c86ff4..000000000 --- a/trpc_agent_sdk/advanced_memory/_paths.py +++ /dev/null @@ -1,110 +0,0 @@ -# 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/advanced_memory/_preload_memory.py deleted file mode 100644 index d032a8c42..000000000 --- a/trpc_agent_sdk/advanced_memory/_preload_memory.py +++ /dev/null @@ -1,300 +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. -"""Select and preload long-term memories for the current model request.""" - -from __future__ import annotations - -import json -import re -from dataclasses import dataclass -from datetime import datetime -from datetime import timezone -from html import escape -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.sessions.compact._formats import memory_freshness -from trpc_agent_sdk.sessions.compact._formats import parse_memory_updated_at -from ._runtime import AdvancedMemoryRuntime -from trpc_agent_sdk.types import Content -from trpc_agent_sdk.types import Part - -if TYPE_CHECKING: - from trpc_agent_sdk.context import InvocationContext - -_FRONTMATTER_FIELD = re.compile(r"^(?P[A-Za-z_]+):\s*(?P.*)$", re.MULTILINE) - - -@dataclass(frozen=True) -class MemoryCandidate: - """Describe one topic file using only its frontmatter metadata.""" - - filename: str - name: str - description: str - memory_type: str - updated_at: datetime | None - - def to_selector_dict(self) -> dict[str, str]: - """Render the metadata passed to the relevance selector.""" - return { - "filename": self.filename, - "name": self.name, - "description": self.description, - "type": self.memory_type, - "freshness": memory_freshness(self.updated_at), - } - - -class MemoryRelevanceSelector(Protocol): - """Select relevant topic filenames for a user query.""" - - async def select( - self, - query: str, - candidates: list[MemoryCandidate], - ctx: "InvocationContext", - *, - limit: int, - ) -> list[str]: - """Return at most ``limit`` filenames from the candidate list.""" - - -def _frontmatter(content: str) -> dict[str, str]: - """Parse the simple single-line frontmatter used by MemoryDocument.""" - if not content.startswith("---\n"): - return {} - end = content.find("\n---", 4) - if end < 0: - return {} - return {match.group("field"): match.group("value").strip() for match in _FRONTMATTER_FIELD.finditer(content[4:end])} - - -def _candidate_from_content(filename: str, content: str) -> MemoryCandidate: - """Build candidate metadata from a topic document.""" - metadata = _frontmatter(content) - return MemoryCandidate( - filename=filename, - name=metadata.get("name", filename), - description=metadata.get("description", ""), - memory_type=metadata.get("type", ""), - updated_at=parse_memory_updated_at(content), - ) - - -class ModelMemoryRelevanceSelector: - """Use an isolated lightweight Agent 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) - if model is None: - raise ValueError("Memory relevance selector cannot resolve an LLM model") - return model - - @staticmethod - def _build_prompt(query: str, candidates: list[MemoryCandidate], limit: int) -> str: - """Build the strict JSON selection prompt.""" - candidate_payload = [candidate.to_selector_dict() for candidate in candidates] - return ("Select the long-term memory files that are clearly relevant to the user's query.\n" - f"Return at most {limit} filenames. If none are clearly relevant, return an empty list.\n" - "Use freshness as one relevance signal, but do not discard an older memory solely because it is old.\n" - "Only return filenames from the candidate list. Do not explain your choices.\n" - 'Return exactly one JSON object: {"selected_memories": ["filename.md"]}\n\n' - f"User query:\n{query}\n\n" - f"Candidate memories:\n{json.dumps(candidate_payload, ensure_ascii=False, indent=2)}") - - @staticmethod - def _parse_selection( - text: str, - candidates: list[MemoryCandidate], - limit: int, - ) -> list[str]: - """Parse and validate the selector's JSON response.""" - payload = None - decoder = json.JSONDecoder() - for index, character in enumerate(text): - if character != "{": - continue - try: - value, _ = decoder.raw_decode(text[index:]) - except json.JSONDecodeError: - continue - if isinstance(value, dict): - payload = value - break - if payload is None: - raise ValueError("Memory relevance selector returned no JSON object") - selected = payload.get("selected_memories") - if not isinstance(selected, list): - raise ValueError("Memory relevance selector returned an invalid selected_memories list") - valid_filenames = {candidate.filename for candidate in candidates} - result: list[str] = [] - for filename in selected: - if isinstance(filename, str) and filename in valid_filenames and filename not in result: - result.append(filename) - if len(result) >= limit: - break - return result - - async def select( - self, - query: str, - candidates: list[MemoryCandidate], - ctx: "InvocationContext", - *, - 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, - ) - 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() - - -async def select_relevant_memory_filenames( - query: str, - candidates: list[MemoryCandidate], - ctx: "InvocationContext", - *, - selector: MemoryRelevanceSelector, - limit: int, -) -> list[str]: - """Select relevant memory filenames behind a replaceable screening boundary.""" - selected = await selector.select(query, candidates, ctx, limit=limit) - valid_filenames = {candidate.filename for candidate in candidates} - return list(dict.fromkeys(filename for filename in selected if filename in valid_filenames))[:limit] - - -class MemoryPreloader: - """Find and render relevant topic files for automatic prompt injection.""" - - def __init__( - self, - runtime: AdvancedMemoryRuntime, - selector: MemoryRelevanceSelector | None = None, - ) -> None: - """Store the runtime and replaceable relevance selector.""" - self._runtime = runtime - self._selector = selector or ModelMemoryRelevanceSelector() - - 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 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[: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.""" - config = self._runtime.config - if not config.enabled or not config.preload_memory_enabled or not query.strip(): - return None - try: - candidates = await self._candidates(ctx) - except Exception as exc: # noqa: BLE001 - logger.warning("Advanced Memory preload candidate loading failed: %s", exc) - return None - if not candidates: - return None - try: - selected = await select_relevant_memory_filenames( - query, - candidates, - ctx, - selector=self._selector, - limit=config.preload_memory_max_topics, - ) - except Exception as exc: # noqa: BLE001 - logger.warning("Advanced Memory preload selection failed: %s", exc) - return None - by_filename = {candidate.filename: candidate for candidate in candidates} - sections: list[str] = [] - used_chars = 0 - for filename in selected: - candidate = by_filename.get(filename) - if candidate is None: - continue - try: - 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 - if full_content is None: - continue - remaining = config.preload_memory_max_chars - used_chars - if remaining <= 0: - break - truncated = len(full_content) > remaining - content = full_content[:remaining] - safe_filename = escape(filename, quote=True) - sections.append(f'\n' - f"{content}\n" - "") - used_chars += len(content) - if not sections: - return None - return ( - "\n" - "The following memories were automatically selected for the current request. " - "They are historical observations, not guaranteed current facts. Verify them when necessary " - "and update them if they are outdated or incorrect. Each memory includes its source filename " - "and has already been read for this request. You may read or update these files again when needed.\n\n" + - "\n\n".join(sections) + "\n") diff --git a/trpc_agent_sdk/advanced_memory/_redis_stores.py b/trpc_agent_sdk/advanced_memory/_redis_stores.py deleted file mode 100644 index 07f4f73ae..000000000 --- a/trpc_agent_sdk/advanced_memory/_redis_stores.py +++ /dev/null @@ -1,302 +0,0 @@ -"""Redis implementations of the Advanced Memory storage contracts.""" - -from __future__ import annotations - -import asyncio -import json -from collections.abc import Mapping -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 trpc_agent_sdk.sessions.compact._formats import MemoryDocument, MemoryIndexEntry -from ._paths import AdvancedMemoryPaths - -_APPEND_UNIQUE_SCRIPT = """ -if redis.call('SADD', KEYS[2], ARGV[1]) == 0 then return 0 end -redis.call('XADD', KEYS[1], '*', 'data', ARGV[2]) -return 1 -""" - -_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}}}" - self._app_base = f"{config.redis_key_prefix}:{{{app_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 _session_base(self, session_id: str) -> str: - safe_session_id = self._paths.session_dir(session_id).name - tenant = f"{self._paths.tenant_root_dir.parent.name}:{self._paths.tenant_root_dir.name}" - return f"{self._config.redis_key_prefix}:{{{tenant}:{safe_session_id}}}" - - def _session_registry(self, session_id: str) -> str: - return f"{self._session_base(session_id)}:keys" - - 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_session_ttl(self, session_id: str, *keys: str) -> None: - skip_prefixes: tuple[str, ...] = () - if not self._config.session_ttl_delete_transcripts: - skip_prefixes = (f"{self._session_base(session_id)}:transcript", ) - await self._refresh_ttl_group( - self._session_registry(session_id), - list(keys), - self._config.session_ttl_seconds, - skip_prefixes=skip_prefixes, - ) - - async def _refresh_memory_ttl(self, *keys: str) -> None: - await self._refresh_ttl_group( - self._memory_registry(), - list(keys), - self._config.memory_ttl_seconds, - ) - - async def delete_session(self, session_id: str) -> None: - """Delete all Advanced Memory keys for one session.""" - session_base = self._session_base(session_id) - registry = self._session_registry(session_id) - keys: set[str] = {registry} - tracked = await self._command("smembers", registry) or [] - keys.update(value for value in (self._text(item) for item in tracked) if value) - - cursor: Any = 0 - pattern = f"{session_base}:*" - while True: - cursor, scanned = await self._command( - "scan", - cursor, - match=pattern, - count=100, - ) - keys.update(value for value in (self._text(item) for item in scanned) if value) - if int(cursor) == 0: - break - if keys: - await self._command("delete", *keys) - - @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() - lines, used_bytes = [], 0 - for line in value.splitlines(keepends=True)[:self._config.memory_index_max_lines]: - size = len(line.encode(self._config.encoding)) - if used_bytes + size > self._config.memory_index_max_bytes: - break - lines.append(line) - used_bytes += size - return "".join(lines) - - 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] - - -class RedisToolResultStore(_RedisStore): - - async def write(self, session_id: str, result_id: str, serialized_result: str) -> Path: - key = f"{self._session_base(session_id)}:tool:{result_id}" - await self._command("set", key, serialized_result) - await self._refresh_session_ttl(session_id, key) - return Path(f"advanced-memory://{key}") - - async def read(self, session_id: str, result_id: str) -> str | None: - key = f"{self._session_base(session_id)}:tool:{result_id}" - value = await self._command("get", key) - await self._refresh_session_ttl(session_id, key) - return self._text(value) - - -class RedisTranscriptStore(_RedisStore): - - @staticmethod - def _validate_record(record: Mapping[str, Any]) -> None: - """Reject Event and Session Memory duplication in Redis.""" - if record.get("kind") in {"event", "session-memory-checkpoint"}: - raise ValueError("Redis transcripts only store context-compression records") - - async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: - self._validate_record(record) - payload = dict(record) - payload.setdefault("recorded_at", datetime.now(timezone.utc).isoformat()) - stream = f"{self._session_base(session_id)}:transcript" - await self._command("xadd", stream, {"data": json.dumps(payload)}) - await self._refresh_session_ttl(session_id, stream) - return Path(f"advanced-memory://{stream}") - - async def append_unique(self, session_id: str, record: Mapping[str, Any], *, unique_key: str) -> tuple[Path, bool]: - self._validate_record(record) - payload = dict(record) - value = payload.get(unique_key) - if not isinstance(value, str) or not 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()) - stream = f"{self._session_base(session_id)}:transcript" - seen = f"{stream}:seen:{unique_key}" - async with self._storage.create_db_session() as connection: - added = await self._storage.execute_command( - connection, - RedisCommand( - method="eval", - args=(_APPEND_UNIQUE_SCRIPT, 2, stream, seen, value, json.dumps(payload, ensure_ascii=False)), - )) - await self._refresh_session_ttl(session_id, stream, seen) - return Path(f"advanced-memory://{stream}"), bool(added) - - async def read_all(self, session_id: str) -> list[dict[str, Any]]: - stream = f"{self._session_base(session_id)}:transcript" - entries = await self._command("xrange", stream, "-", "+") - await self._refresh_session_ttl(session_id, stream) - records: list[dict[str, Any]] = [] - for _, fields in entries: - value = fields.get(b"data") if isinstance(fields, dict) else None - value = value or fields.get("data") - text = self._text(value) - if text: - records.append(json.loads(text)) - return records diff --git a/trpc_agent_sdk/advanced_memory/_runtime.py b/trpc_agent_sdk/advanced_memory/_runtime.py deleted file mode 100644 index 31bd28134..000000000 --- a/trpc_agent_sdk/advanced_memory/_runtime.py +++ /dev/null @@ -1,206 +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 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 trpc_agent_sdk.sessions.compact._coordination import SessionOperationCoordinator -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 - coordination: SessionOperationCoordinator - 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) - _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 - 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 - sql_storage = SqlStorage( - is_async=resolved_config.sql_is_async, - db_url=resolved_config.sql_url, - metadata=AdvancedMemorySqlBase.metadata, - expire_on_commit=False, - ) - else: - local_cleanup = LocalAdvancedMemoryCleanup(resolved_config) - return cls( - config=resolved_config, - paths=paths, - coordination=SessionOperationCoordinator(), - long_term_memory=LongTermMemoryStore(resolved_config, paths), - _redis_storage=redis_storage, - _sql_storage=sql_storage, - _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") - 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_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 - - @property - def coordination(self) -> SessionOperationCoordinator: - """Return the shared coordinator.""" - return self.root.coordination - - def session_key(self, session_id: str) -> str: - """Return a lock/cache key unique across all tenants.""" - return f"{self.scope.storage_key}\0{session_id}" - - async def initialize(self) -> bool: - """Initialize only this tenant's local directories.""" - if not self.config.enabled: - return False - if self.config.storage_backend == "sql" and self.root._sql_cleanup is not None: - await self.root._sql_cleanup.start() - if self.config.storage_backend == "local" and self.root._local_cleanup is not None: - await self.root._local_cleanup.start() - await self.long_term_memory.initialize() - return True diff --git a/trpc_agent_sdk/advanced_memory/_sql_stores.py b/trpc_agent_sdk/advanced_memory/_sql_stores.py deleted file mode 100644 index 11862c982..000000000 --- a/trpc_agent_sdk/advanced_memory/_sql_stores.py +++ /dev/null @@ -1,533 +0,0 @@ -"""SQL implementations of the Advanced Memory storage contracts.""" - -from __future__ import annotations - -import json -import asyncio -import hashlib -import uuid -from datetime import datetime, timedelta, timezone -from dataclasses import replace -from pathlib import Path -from collections.abc import Mapping -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 trpc_agent_sdk.sessions.compact._formats import MemoryDocument, MemoryIndexEntry -from ._paths import AdvancedMemoryPaths - - -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 SqlTranscript(AdvancedMemorySqlBase): - __tablename__ = "advanced_memory_transcripts" - - 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) - session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) - record_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) - payload: Mapped[str] = mapped_column(Text) - recorded_at: Mapped[datetime] = mapped_column(PreciseTimestamp, default=func.now()) - expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) - - -class SqlTranscriptSeen(AdvancedMemorySqlBase): - __tablename__ = "advanced_memory_transcript_seen" - - dedupe_id: Mapped[str] = mapped_column(String(64), primary_key=True) - app_name: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) - user_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) - session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), index=True) - unique_key: Mapped[str] = mapped_column(String(DEFAULT_MAX_VARCHAR_LENGTH)) - unique_value: Mapped[str] = mapped_column(String(DEFAULT_MAX_VARCHAR_LENGTH)) - expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) - - -class SqlToolResult(AdvancedMemorySqlBase): - __tablename__ = "advanced_memory_tool_results" - - 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) - session_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) - result_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_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 - - async def _refresh_session_scope(self, db: Any, session_id: str) -> None: - expiry = self._expiry(self._config.session_ttl_seconds) - if expiry is None: - return - tables = ((SqlToolResult, (self._app_name, self._user_id, session_id)), ) - if self._config.session_ttl_delete_transcripts: - tables = ( - (SqlTranscript, (self._app_name, self._user_id, session_id)), - (SqlTranscriptSeen, (self._app_name, self._user_id, session_id)), - *tables, - ) - for model, key in tables: - rows = await self._storage.query( - db, - SqlKey(key=key, storage_cls=model), - SqlCondition(filters=[ - getattr(model, "app_name") == self._app_name, - getattr(model, "user_id") == self._user_id, - getattr(model, "session_id") == session_id, - getattr(model, "expires_at").is_(None) | (getattr(model, "expires_at") > self._now()), - ]), - ) - for row in rows: - row.expires_at = expiry - - async def delete_session(self, session_id: str) -> None: - """Delete all Advanced Memory rows for one session.""" - models = ( - SqlTranscript, - SqlTranscriptSeen, - SqlToolResult, - ) - filters = { - SqlTranscript: [ - SqlTranscript.app_name == self._app_name, - SqlTranscript.user_id == self._user_id, - SqlTranscript.session_id == session_id, - ], - SqlTranscriptSeen: [ - SqlTranscriptSeen.app_name == self._app_name, - SqlTranscriptSeen.user_id == self._user_id, - SqlTranscriptSeen.session_id == session_id, - ], - SqlToolResult: [ - SqlToolResult.app_name == self._app_name, - SqlToolResult.user_id == self._user_id, - SqlToolResult.session_id == session_id, - ], - } - async with self._storage.create_db_session() as db: - for model in models: - await self._storage.delete( - db, - SqlKey(key=tuple(), storage_cls=model), - SqlCondition(filters=filters[model]), - ) - await self._storage.commit(db) - - -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) - await self._storage.commit(db) - content = row.content - lines, used_bytes = [], 0 - for line in content.splitlines(keepends=True)[:self._config.memory_index_max_lines]: - size = len(line.encode(self._config.encoding)) - if used_bytes + size > self._config.memory_index_max_bytes: - break - lines.append(line) - used_bytes += size - return "".join(lines) - - 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 SqlToolResultStore(_SqlStore): - - async def write(self, session_id: str, result_id: str, serialized_result: str) -> Path: - async with self._storage.create_db_session() as db: - key = (self._app_name, self._user_id, session_id, result_id) - row = await self._storage.get(db, SqlKey(key=key, storage_cls=SqlToolResult)) - if row is None: - row = SqlToolResult( - app_name=key[0], - user_id=key[1], - session_id=key[2], - result_id=key[3], - ) - await self._storage.add(db, row) - row.content = serialized_result - row.updated_at = self._now() - row.expires_at = self._expiry(self._config.session_ttl_seconds) - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/tool/{result_id}") - - async def read(self, session_id: str, result_id: str) -> str | None: - async with self._storage.create_db_session() as db: - row = await self._storage.get( - db, - SqlKey(key=(self._app_name, self._user_id, session_id, result_id), storage_cls=SqlToolResult), - ) - if row is None or self._expired(row.expires_at): - return None - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return row.content - - -class SqlTranscriptStore(_SqlStore): - - @staticmethod - def _validate_record(record: Mapping[str, Any]) -> None: - """Reject Event and Session Memory duplication in SQL.""" - if record.get("kind") in {"event", "session-memory-checkpoint"}: - raise ValueError("SQL transcripts only store context-compression records") - - def _dedupe_id(self, session_id: str, unique_key: str, value: str) -> str: - raw = "\0".join((self._app_name, self._user_id, session_id, unique_key, value)) - return hashlib.sha256(raw.encode("utf-8")).hexdigest() - - async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: - self._validate_record(record) - payload = dict(record) - payload.setdefault("recorded_at", self._now().replace(tzinfo=timezone.utc).isoformat()) - async with self._storage.create_db_session() as db: - await self._storage.add( - db, - SqlTranscript( - app_name=self._app_name, - user_id=self._user_id, - session_id=session_id, - record_id=uuid.uuid4().hex, - payload=json.dumps(payload, ensure_ascii=False), - expires_at=(self._expiry(self._config.session_ttl_seconds) - if self._config.session_ttl_delete_transcripts else None), - )) - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript") - - async def append_unique( - self, - session_id: str, - record: Mapping[str, Any], - *, - unique_key: str, - ) -> tuple[Path, bool]: - self._validate_record(record) - payload = dict(record) - value = payload.get(unique_key) - if not isinstance(value, str) or not value: - raise ValueError(f"Transcript unique key {unique_key!r} must be a non-empty string") - async with self._storage.create_db_session() as db: - dedupe_id = self._dedupe_id(session_id, unique_key, value) - seen_key = (self._app_name, self._user_id, session_id, unique_key, value) - seen = await self._storage.get( - db, - SqlKey(key=(dedupe_id, ), storage_cls=SqlTranscriptSeen), - ) - if seen is not None and not self._expired(seen.expires_at): - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript"), False - if seen is not None: - await self._storage.delete( - db, - SqlKey(key=(dedupe_id, ), storage_cls=SqlTranscriptSeen), - SqlCondition(filters=[ - SqlTranscriptSeen.dedupe_id == dedupe_id, - ]), - ) - payload.setdefault("recorded_at", self._now().replace(tzinfo=timezone.utc).isoformat()) - await self._storage.add( - db, - SqlTranscriptSeen( - dedupe_id=dedupe_id, - app_name=seen_key[0], - user_id=seen_key[1], - session_id=seen_key[2], - unique_key=seen_key[3], - unique_value=seen_key[4], - expires_at=(self._expiry(self._config.session_ttl_seconds) - if self._config.session_ttl_delete_transcripts else None), - )) - await self._storage.add( - db, - SqlTranscript( - app_name=self._app_name, - user_id=self._user_id, - session_id=session_id, - record_id=uuid.uuid4().hex, - payload=json.dumps(payload, ensure_ascii=False), - expires_at=(self._expiry(self._config.session_ttl_seconds) - if self._config.session_ttl_delete_transcripts else None), - )) - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return Path(f"advanced-memory://sql/{self._app_name}/{self._user_id}/{session_id}/transcript"), True - - async def read_all(self, session_id: str) -> list[dict[str, Any]]: - async with self._storage.create_db_session() as db: - rows = await self._storage.query( - db, - SqlKey(key=(self._app_name, self._user_id, session_id), storage_cls=SqlTranscript), - SqlCondition( - filters=[ - SqlTranscript.app_name == self._app_name, - SqlTranscript.user_id == self._user_id, - SqlTranscript.session_id == session_id, - SqlTranscript.expires_at.is_(None) | (SqlTranscript.expires_at > self._now()), - ], - order_func=SqlTranscript.recorded_at.asc, - ), - ) - await self._refresh_session_scope(db, session_id) - await self._storage.commit(db) - return [json.loads(row.payload) for row in rows] - - -class SqlAdvancedMemoryCleanup: - """Periodically remove expired Advanced Memory SQL rows.""" - - _models = ( - SqlMemoryIndex, - SqlMemoryTopic, - SqlTranscript, - SqlTranscriptSeen, - SqlToolResult, - ) - - 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 - and self._config.session_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: - models = self._models if self._config.session_ttl_delete_transcripts else tuple( - model for model in self._models if model is not SqlTranscript) - for model in models: - await self._storage.delete( - db, - SqlKey(key=tuple(), storage_cls=model), - SqlCondition(filters=[model.expires_at.is_not(None), model.expires_at <= now]), - ) - 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", - "SqlToolResultStore", - "SqlTranscriptStore", -] diff --git a/trpc_agent_sdk/advanced_memory/_storage.py b/trpc_agent_sdk/advanced_memory/_storage.py deleted file mode 100644 index d4d17daea..000000000 --- a/trpc_agent_sdk/advanced_memory/_storage.py +++ /dev/null @@ -1,189 +0,0 @@ -# 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 tempfile -import time -from dataclasses import replace -from datetime import datetime -from datetime import timezone -from pathlib import Path - -from trpc_agent_sdk.sessions.compact._formats import MemoryDocument -from trpc_agent_sdk.sessions.compact._formats import MemoryIndexEntry - -from ._config import AdvancedMemoryServiceConfig -from ._paths import AdvancedMemoryPaths - - -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 "" - lines: list[str] = [] - used_bytes = 0 - with self.index_path.open(encoding=self._config.encoding) as source: - for _ in range(self._config.memory_index_max_lines): - line = source.readline() - if not line: - break - size = len(line.encode(self._config.encoding)) - if used_bytes + size > self._config.memory_index_max_bytes: - break - lines.append(line) - used_bytes += size - return "".join(lines) - - 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/advanced_memory/_storage_backend.py b/trpc_agent_sdk/advanced_memory/_storage_backend.py deleted file mode 100644 index 4e10a2f0c..000000000 --- a/trpc_agent_sdk/advanced_memory/_storage_backend.py +++ /dev/null @@ -1,30 +0,0 @@ -"""Storage boundary for Advanced Memory tenant namespaces. - -Backends expose logical records rather than filesystem paths so a future Redis -implementation can preserve the same tenant and session semantics. -""" - -from __future__ import annotations - -from typing import Protocol - -from ._paths import MemoryScope -from ._runtime import ScopedAdvancedMemoryRuntime - - -class AdvancedMemoryStorageBackend(Protocol): - """Create storage views isolated to an application user.""" - - def for_scope(self, scope: MemoryScope) -> ScopedAdvancedMemoryRuntime: - """Return the tenant-bound storage view.""" - - -class LocalAdvancedMemoryStorageBackend: - """Adapt the file-backed runtime to the storage backend boundary.""" - - def __init__(self, runtime: object) -> None: - self._runtime = runtime - - def for_scope(self, scope: MemoryScope) -> ScopedAdvancedMemoryRuntime: - """Return a file-backed scope without exposing local path mechanics.""" - return self._runtime.for_scope(scope.app_name, scope.user_id) diff --git a/trpc_agent_sdk/evaluation/_eval_session_service.py b/trpc_agent_sdk/evaluation/_eval_session_service.py index 6e8e5e8fd..a63712a0d 100644 --- a/trpc_agent_sdk/evaluation/_eval_session_service.py +++ b/trpc_agent_sdk/evaluation/_eval_session_service.py @@ -9,16 +9,12 @@ 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, @@ -34,19 +30,6 @@ 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, @@ -108,17 +91,6 @@ 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 a4b768e07..db9e1c012 100644 --- a/trpc_agent_sdk/memory/__init__.py +++ b/trpc_agent_sdk/memory/__init__.py @@ -13,7 +13,6 @@ from trpc_agent_sdk.abc import MemoryServiceABC as BaseMemoryService from trpc_agent_sdk.abc import MemoryServiceConfig -from ._advanced_memory_service import AdvancedMemoryService from ._in_memory_memory_service import EventTtl from ._in_memory_memory_service import InMemoryMemoryService from ._redis_memory_service import RedisMemoryService @@ -27,8 +26,6 @@ __all__ = [ "BaseMemoryService", "MemoryServiceConfig", - "AdvancedMemoryServiceConfig", - "AdvancedMemoryService", "EventTtl", "InMemoryMemoryService", "RedisMemoryService", @@ -40,11 +37,3 @@ "format_timestamp", ] - -def __getattr__(name: str): - """Lazily expose Advanced Memory configuration without import cycles.""" - if name == "AdvancedMemoryServiceConfig": - from trpc_agent_sdk.advanced_memory import AdvancedMemoryServiceConfig - - return AdvancedMemoryServiceConfig - 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 c626256cd..e69de29bb 100644 --- a/trpc_agent_sdk/memory/_advanced_memory_service.py +++ b/trpc_agent_sdk/memory/_advanced_memory_service.py @@ -1,113 +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. -"""Runner-compatible facade for the Advanced Memory mechanism.""" - -from __future__ import annotations - -from typing import Any -from typing import Optional -from typing import TYPE_CHECKING - -from trpc_agent_sdk.abc import MemoryServiceABC as BaseMemoryService -from trpc_agent_sdk.abc import MemoryServiceConfig -from trpc_agent_sdk.abc import SearchMemoryResponse -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 AdvancedMemoryServiceConfig - from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime - from trpc_agent_sdk.advanced_memory import LongTermMemoryIntegration - - -class AdvancedMemoryService(BaseMemoryService): - """Expose user-scoped long-term Memory through the Runner memory API. - - ``Runner`` calls :meth:`bind` automatically. Session compression is - configured independently through ``SessionService.session_compact_manager``. - """ - - def __init__( - self, - config: AdvancedMemoryServiceConfig | None = None, - *, - runtime: AdvancedMemoryRuntime | 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 AdvancedMemoryServiceConfig - from trpc_agent_sdk.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 AdvancedMemoryServiceConfig()) - super().__init__(MemoryServiceConfig(enabled=resolved_config.enabled)) - self._runtime = runtime or AdvancedMemoryRuntime.create(resolved_config) - self._preload_memory_model = preload_memory_model - self._install_long_term_memory_tools = install_long_term_memory_tools - self._integration: LongTermMemoryIntegration | None = None - self._bound_agent: Any | None = None - - @property - def config(self) -> AdvancedMemoryServiceConfig: - """Return the Advanced Memory configuration.""" - return self._runtime.config - - @property - def runtime(self) -> AdvancedMemoryRuntime: - """Return the Advanced Memory runtime.""" - return self._runtime - - @property - 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 long-term Memory and return the unchanged SessionService.""" - from trpc_agent_sdk.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") - return session_service - - self._integration = setup_long_term_memory( - agent, - self._runtime, - preload_memory_model=self._preload_memory_model, - install_tools=self._install_long_term_memory_tools, - ) - self._bound_agent = agent - return session_service - - async def store_session( - self, - session: Session, - agent_context: Optional[AgentContext] = None, - ) -> None: - """Long-term Memory is updated explicitly through its tools.""" - return None - - async def search_memory( - self, - key: str, - query: str, - limit: int = 10, - agent_context: Optional[AgentContext] = None, - ) -> SearchMemoryResponse: - """Return an empty legacy-style response. - - Advanced long-term memory is intentionally accessed through its - ``save_memory``, ``read_memory``, and ``list_memory_index`` tools. - """ - return SearchMemoryResponse() - - async def close(self) -> None: - """Release service-owned local or external storage resources.""" - await self._runtime.close() diff --git a/trpc_agent_sdk/runners.py b/trpc_agent_sdk/runners.py index 418d09519..21afe7dae 100644 --- a/trpc_agent_sdk/runners.py +++ b/trpc_agent_sdk/runners.py @@ -223,16 +223,6 @@ def __init__( the memory service. Set to False when the service is managed outside the runner. """ - # Advanced Memory needs the agent and session service in addition to - # the traditional memory-service hook. Bind it here so callers can - # use the same construction pattern as Redis/Mem0 memory services. - from trpc_agent_sdk.memory import AdvancedMemoryService - - if isinstance(memory_service, AdvancedMemoryService): - session_service = memory_service.bind(agent, session_service) - 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 9ce67dc43..8cb7f9f6f 100644 --- a/trpc_agent_sdk/sessions/__init__.py +++ b/trpc_agent_sdk/sessions/__init__.py @@ -16,29 +16,39 @@ from ._base_session_service import BaseSessionService from ._history_record import HistoryRecord +from .compact.default import DefaultSessionSummarizer +from .compact.default import DefaultSessionSummary +from .compact.default import DefaultSessionSummarizerManager +from .compact.default import CheckSummarizerFunction +from .compact.default import set_summarizer_check_functions_by_and +from .compact.default import set_summarizer_check_functions_by_or +from .compact.default import set_summarizer_conversation_threshold +from .compact.default import set_summarizer_events_count_threshold +from .compact.default import set_summarizer_important_content_threshold +from .compact.default import set_summarizer_time_interval_threshold +from .compact.default import set_summarizer_token_threshold +from .compact.advanced import AdvancedAutoCompactSummarizer +from .compact.advanced import AdvancedAutoCompactSummarizerManager +from .compact.advanced import BaseCompactSummarizerHandler +from .compact.advanced import BaseTokenEstimator +from .compact.advanced import BaseModelContextWindowResolver +from .compact.advanced import AutoCompactSummarizerConfig +from .compact.advanced import HistorySnipConfig +from .compact.advanced import TokenContextTrackerConfig +from .compact.advanced import MicroCompactConfig +from .compact.advanced import AdvancedAutoCompactSummarizerConfig from ._in_memory_session_service import InMemorySessionService from ._in_memory_session_service import SessionWithTTL from ._in_memory_session_service import StateWithTTL from ._redis_session_service import RedisSessionService from ._redis_cluster_session_service import RedisClusterSessionService from ._session import Session -from ._session_summarizer import SessionSummarizer -from ._session_summarizer import SessionSummary from ._sql_session_service import SessionStorageBase from ._sql_session_service import SessionStorageEvent from ._sql_session_service import SqlSessionService from ._sql_session_service import StorageAppState from ._sql_session_service import StorageSession from ._sql_session_service import StorageUserState -from ._summarizer_checker import CheckSummarizerFunction -from ._summarizer_checker import set_summarizer_check_functions_by_and -from ._summarizer_checker import set_summarizer_check_functions_by_or -from ._summarizer_checker import set_summarizer_conversation_threshold -from ._summarizer_checker import set_summarizer_events_count_threshold -from ._summarizer_checker import set_summarizer_important_content_threshold -from ._summarizer_checker import set_summarizer_time_interval_threshold -from ._summarizer_checker import set_summarizer_token_threshold -from ._summarizer_manager import SummarizerSessionManager from ._types import SessionServiceConfig from ._utils import StateStorageEntry from ._utils import app_state_key @@ -49,14 +59,15 @@ from ._utils import session_key from ._utils import user_state_key +# Default compact session summarizer for backward compatibility +SessionSummary = DefaultSessionSummary +SessionSummarizer = DefaultSessionSummarizer +SummarizerSessionManager = DefaultSessionSummarizerManager + __all__ = [ "ListSessionsResponse", "State", "BaseSessionService", - "BaseSessionCompactManager", - "AdvancedCompactConfig", - "AdvancedSessionCompactManager", - "AutoCompact", "HistoryRecord", "InMemorySessionService", "SessionWithTTL", @@ -64,8 +75,6 @@ "RedisSessionService", "RedisClusterSessionService", "Session", - "SessionSummarizer", - "SessionSummary", "SessionStorageBase", "SessionStorageEvent", "SqlSessionService", @@ -80,6 +89,9 @@ "set_summarizer_important_content_threshold", "set_summarizer_time_interval_threshold", "set_summarizer_token_threshold", + "DefaultSessionSummarizer", + "DefaultSessionSummary", + "DefaultSessionSummarizerManager", "SummarizerSessionManager", "SessionServiceConfig", "StateStorageEntry", @@ -90,18 +102,14 @@ "is_summary_anchor", "session_key", "user_state_key", + "AdvancedAutoCompactSummarizer", + "AdvancedAutoCompactSummarizerManager", + "BaseCompactSummarizerHandler", + "BaseTokenEstimator", + "BaseModelContextWindowResolver", + "AutoCompactSummarizerConfig", + "HistorySnipConfig", + "TokenContextTrackerConfig", + "MicroCompactConfig", + "AdvancedAutoCompactSummarizerConfig", ] - - -def __getattr__(name: str): - """Lazily expose Advanced Memory without creating an import cycle.""" - if name in { - "AdvancedCompactConfig", - "AdvancedSessionCompactManager", - "AutoCompact", - "BaseSessionCompactManager", - }: - from . import compact - - return getattr(compact, name) - raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/trpc_agent_sdk/sessions/_base_session_service.py b/trpc_agent_sdk/sessions/_base_session_service.py index 6cbd8fbed..84c02881d 100644 --- a/trpc_agent_sdk/sessions/_base_session_service.py +++ b/trpc_agent_sdk/sessions/_base_session_service.py @@ -24,22 +24,19 @@ """Base session service interface.""" 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 +from trpc_agent_sdk.abc import CompactSummarizerManagerABC from trpc_agent_sdk.context import InvocationContext from trpc_agent_sdk.events import Event from trpc_agent_sdk.types import State from ._session import Session -from ._summarizer_manager import SummarizerSessionManager from ._types import SessionServiceConfig -if TYPE_CHECKING: - from .compact import BaseSessionCompactManager - class BaseSessionService(SessionServiceABC): """Abstract base class for session management services. @@ -48,18 +45,15 @@ class BaseSessionService(SessionServiceABC): """ def __init__(self, - summarizer_manager: Optional[SummarizerSessionManager] = None, - session_config: Optional[SessionServiceConfig] = None, - session_compact_manager: Optional["BaseSessionCompactManager"] = None): + summarizer_manager: Optional[CompactSummarizerManagerABC] = None, + session_config: Optional[SessionServiceConfig] = 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 @@ -67,11 +61,9 @@ 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]: + def summarizer_manager(self) -> Optional[CompactSummarizerManagerABC]: """Get the summarizer manager.""" return self._summarizer_manager @@ -80,38 +72,16 @@ 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: + def set_summarizer_manager(self, summarizer_manager: CompactSummarizerManagerABC, force: bool = False) -> None: """Set the summarizer manager to use. Args: 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) + self._summarizer_manager.set_session_service(self, force) @override async def append_event(self, session: Session, event: Event) -> Event: @@ -205,8 +175,6 @@ 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]: @@ -220,10 +188,10 @@ async def get_session_summary(self, session: Session) -> Optional[str]: """ if self._summarizer_manager: summary = await self._summarizer_manager.get_session_summary(session) - if summary: + if isinstance(summary, str): + return summary + if summary is not None: return summary.summary_text - 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: @@ -246,5 +214,5 @@ def filter_events(self, session: Session, need_copy: bool = False) -> Session: @override async def close(self) -> None: """Closes the session service and releases any resources.""" - if self._session_compact_manager: - await self._session_compact_manager.close() + if self._summarizer_manager: + await self._summarizer_manager.close() diff --git a/trpc_agent_sdk/sessions/_in_memory_session_service.py b/trpc_agent_sdk/sessions/_in_memory_session_service.py index fdd1d1ce3..c19d56ab1 100644 --- a/trpc_agent_sdk/sessions/_in_memory_session_service.py +++ b/trpc_agent_sdk/sessions/_in_memory_session_service.py @@ -31,13 +31,13 @@ import uuid from typing import Any from typing import Optional -from typing import TYPE_CHECKING from typing_extensions import override from pydantic import BaseModel from pydantic import Field from trpc_agent_sdk.abc import ListSessionsResponse +from trpc_agent_sdk.abc import CompactSummarizerManagerABC from trpc_agent_sdk.context import AgentContext from trpc_agent_sdk.events import Event from trpc_agent_sdk.log import logger @@ -46,15 +46,11 @@ from ._base_session_service import BaseSessionService from ._session import Session -from ._summarizer_manager import SummarizerSessionManager from ._types import SessionServiceConfig from ._utils import StateStorageEntry from ._utils import extract_state_delta from ._utils import merge_state -if TYPE_CHECKING: - from .compact._base_manager import BaseSessionCompactManager - class SessionWithTTL(BaseModel): """Wrapper for session with TTL support.""" @@ -111,14 +107,9 @@ class InMemorySessionService(BaseSessionService): """ def __init__(self, - summarizer_manager: Optional[SummarizerSessionManager] = None, - 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, - ) + summarizer_manager: Optional[CompactSummarizerManagerABC] = None, + session_config: Optional[SessionServiceConfig] = None): + super().__init__(summarizer_manager=summarizer_manager, session_config=session_config) # Storage with TTL support # Map: app_name -> user_id -> session_id -> SessionWithTTL self._sessions: dict[str, dict[str, dict[str, SessionWithTTL]]] = {} @@ -278,6 +269,26 @@ def _warning(message: str) -> None: return event + @override + async def update_session_state( + self, + session: Session, + state_delta: dict[str, Any], + ) -> None: + """Patch stored session state without replacing its Event window.""" + if not state_delta: + return + session.state.update(state_delta) + + app_sessions = self._sessions.get(session.app_name) + user_sessions = app_sessions.get(session.user_id) if app_sessions else None + stored = user_sessions.get(session.id) if user_sessions else None + if stored is None: + logger.warning("Session %s not found while updating state", session.id) + return + stored.session.state.update(state_delta) + stored.ttl.update_expired_at() + @override async def update_session(self, session: Session) -> None: """Update a session in storage. @@ -302,21 +313,6 @@ 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 2fce605a8..9190f898f 100644 --- a/trpc_agent_sdk/sessions/_redis_session_service.py +++ b/trpc_agent_sdk/sessions/_redis_session_service.py @@ -13,10 +13,10 @@ 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 +from trpc_agent_sdk.abc import CompactSummarizerManagerABC from trpc_agent_sdk.context import AgentContext from trpc_agent_sdk.events import Event from trpc_agent_sdk.log import logger @@ -28,7 +28,6 @@ from ._base_session_service import BaseSessionService from ._session import Session -from ._summarizer_manager import SummarizerSessionManager from ._types import SessionServiceConfig from ._utils import StateStorageEntry from ._utils import app_state_key @@ -37,9 +36,6 @@ 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. @@ -90,10 +86,9 @@ class RedisSessionService(BaseSessionService): def __init__(self, db_url: str, - summarizer_manager: Optional[SummarizerSessionManager] = None, + summarizer_manager: Optional[CompactSummarizerManagerABC] = None, session_config: Optional[SessionServiceConfig] = None, is_async: bool = False, - session_compact_manager: BaseSessionCompactManager | None = None, **kwargs: Any): self._db_url = db_url self._is_async = is_async @@ -101,7 +96,6 @@ def __init__(self, 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. @@ -266,6 +260,29 @@ def _warning(message: str) -> None: return event + @override + async def update_session_state( + self, + session: Session, + state_delta: dict[str, Any], + ) -> None: + """Persist session-scoped state without replacing the caller's Event window.""" + if not state_delta: + return + session.state.update(state_delta) + + async with self._redis_storage.create_db_session() as redis_session: + key = session_key(session.app_name, session.user_id, session.id) + storage_session = await self._get_session(redis_session, key) + if not storage_session: + logger.warning( + "Session %s not found in Redis while updating state", + session.id, + ) + return + storage_session.state.update(state_delta) + await self._set_session(redis_session, storage_session) + @override async def update_session(self, session: Session) -> None: """Update a session in storage. @@ -282,75 +299,6 @@ 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.""" diff --git a/trpc_agent_sdk/sessions/_session.py b/trpc_agent_sdk/sessions/_session.py index 41335af34..061c9e2b6 100644 --- a/trpc_agent_sdk/sessions/_session.py +++ b/trpc_agent_sdk/sessions/_session.py @@ -160,10 +160,8 @@ def compact_events( None, ) if boundary_index is None: - raise ValueError( - f"Session compaction boundary Event {boundary_event_id!r} " - "is not in the active event window" - ) + 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: diff --git a/trpc_agent_sdk/sessions/_sql_session_service.py b/trpc_agent_sdk/sessions/_sql_session_service.py index 1f998ba9b..e4f5225a0 100644 --- a/trpc_agent_sdk/sessions/_sql_session_service.py +++ b/trpc_agent_sdk/sessions/_sql_session_service.py @@ -34,7 +34,6 @@ 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 @@ -51,6 +50,7 @@ from sqlalchemy.types import Integer from trpc_agent_sdk.abc import ListSessionsResponse +from trpc_agent_sdk.abc import CompactSummarizerManagerABC from trpc_agent_sdk.context import AgentContext from trpc_agent_sdk.events import Event from trpc_agent_sdk.log import logger @@ -72,15 +72,11 @@ from ._base_session_service import BaseSessionService from ._session import Session -from ._summarizer_manager import SummarizerSessionManager from ._types import SessionServiceConfig from ._utils import StateStorageEntry from ._utils import extract_state_delta 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.""" @@ -392,10 +388,9 @@ class SqlSessionService(BaseSessionService): def __init__(self, db_url: str, - summarizer_manager: Optional[SummarizerSessionManager] = None, + summarizer_manager: Optional[CompactSummarizerManagerABC] = None, is_async: bool = False, session_config: Optional[SessionServiceConfig] = None, - session_compact_manager: BaseSessionCompactManager | None = None, **kwargs: Any): self._db_url = db_url self._is_async = is_async @@ -403,7 +398,6 @@ def __init__(self, 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. @@ -644,6 +638,40 @@ async def append_event(self, session: Session, event: Event) -> Event: return event + @override + async def update_session_state( + self, + session: Session, + state_delta: dict[str, Any], + ) -> None: + """Persist session-scoped state without rewriting Event rows.""" + if not state_delta: + return + session.state.update(state_delta) + + async with self._sql_storage.create_db_session() as sql_session: + session_key = SqlKey( + key=(session.app_name, session.user_id, session.id), + storage_cls=StorageSession, + ) + storage_session: Optional[StorageSession] = await self._sql_storage.get_for_update( + sql_session, + session_key, + ) + if storage_session is None: + logger.warning( + "Session %s not found in storage while updating state", + session.id, + ) + return + + persisted_state = dict(storage_session.state or {}) + persisted_state.update(state_delta) + storage_session.state = persisted_state # type: ignore + await self._sql_storage.commit(sql_session) + await self._sql_storage.refresh(sql_session, storage_session) + session.last_update_time = storage_session.update_timestamp_tz + @override async def update_session(self, session: Session) -> None: app_name = session.app_name @@ -679,29 +707,6 @@ 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() diff --git a/trpc_agent_sdk/sessions/compact/__init__.py b/trpc_agent_sdk/sessions/compact/__init__.py index 1fe406010..77a45bb99 100644 --- a/trpc_agent_sdk/sessions/compact/__init__.py +++ b/trpc_agent_sdk/sessions/compact/__init__.py @@ -5,92 +5,73 @@ # tRPC-Agent-Python is licensed under Apache-2.0. """Canonical context-compression package for session management.""" -from ._autocompact import AutoCompact -from ._autocompact import AutoCompactCallback -from ._autocompact import AutoCompactResult -from ._autocompact import content_signature -from ._autocompact import ForkedLegacySummaryGenerator -from ._autocompact import setup_autocompact -from ._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 ._manager import AdvancedSessionCompactManager -from ._microcompact import Microcompact -from ._microcompact import MicrocompactCallback -from ._microcompact import MicrocompactResult -from ._microcompact import setup_microcompact -from ._session_memory import build_session_memory_prompt -from ._session_memory import ForkedSessionMemoryGenerator -from ._session_memory import has_session_memory_content -from ._session_memory import limit_session_memory_document -from ._session_memory import SessionMemoryExtractionInput -from ._session_memory import SessionMemoryExtractionResult -from ._session_memory import SessionMemoryExtractor -from ._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 +from trpc_agent_sdk.abc import CompactTrigger + +from .advanced import AdvancedAutoCompactSummarizer +from .advanced import AdvancedAutoCompactSummarizerManager +from .advanced import BaseCompactSummarizerHandler +from .advanced import BaseTokenEstimator +from .advanced import BaseModelContextWindowResolver +from .advanced import AutoCompactSummarizerConfig +from .advanced import HistorySnipConfig +from .advanced import TokenContextTrackerConfig +from .advanced import MicroCompactConfig +from .advanced import ToolResultBudgetConfig +from .advanced import SessionMemoryExtractorConfig +from .advanced import AdvancedAutoCompactSummarizerConfig +from .advanced import AdvancedAutoCompactSummarizerFilter +from .advanced import HistorySnip +from .advanced import MicroCompact +from .advanced import SessionMemoryDocument +from .advanced import SessionMemoryExtractor +from .advanced import AdvancedAutoCompactSummarizerRuntime +from .advanced import TokenContextTracker +from .advanced import ToolResultBudget +from .default import DEFAULT_SUMMARIZER_PROMPT +from .default import DefaultSessionSummarizer +from .default import DefaultSessionSummarizerManager +from .default import DefaultSessionSummary +from .default import CheckSummarizerFunction +from .default import set_summarizer_token_threshold +from .default import set_summarizer_events_count_threshold +from .default import set_summarizer_time_interval_threshold +from .default import set_summarizer_important_content_threshold +from .default import set_summarizer_conversation_threshold +from .default import set_summarizer_check_functions_by_and +from .default import set_summarizer_check_functions_by_or __all__ = [ - "AdvancedCompactConfig", - "AutoCompact", - "AutoCompactCallback", - "AutoCompactResult", - "ContextBudget", - "ContextTokenEstimate", - "ForkedLegacySummaryGenerator", - "ForkedSessionMemoryGenerator", - "HeuristicTokenEstimator", + "CompactTrigger", + "AdvancedAutoCompactSummarizer", + "AdvancedAutoCompactSummarizerManager", + "BaseCompactSummarizerHandler", + "BaseTokenEstimator", + "BaseModelContextWindowResolver", + "AutoCompactSummarizerConfig", + "HistorySnipConfig", + "TokenContextTrackerConfig", + "MicroCompactConfig", + "ToolResultBudgetConfig", + "SessionMemoryExtractorConfig", + "AdvancedAutoCompactSummarizerConfig", + "AdvancedAutoCompactSummarizerFilter", "HistorySnip", - "HistorySnipCallback", - "HistorySnipResult", - "Microcompact", - "MicrocompactCallback", - "MicrocompactResult", - "ModelContextWindowResolver", - "SESSION_MEMORY_SECTION_DESCRIPTIONS", - "SESSION_MEMORY_SECTIONS", - "SESSION_MEMORY_STATE_KEY", + "MicroCompact", "SessionMemoryDocument", - "SessionMemoryExtractionInput", - "SessionMemoryExtractionResult", "SessionMemoryExtractor", - "BaseSessionCompactManager", - "AdvancedSessionCompactManager", + "AdvancedAutoCompactSummarizerRuntime", "TokenContextTracker", - "TokenEstimator", "ToolResultBudget", - "ToolResultBudgetCallback", - "ToolResultBudgetResult", - "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_microcompact", - "setup_tool_result_budget", + "DEFAULT_SUMMARIZER_PROMPT", + "DefaultSessionSummarizer", + "DefaultSessionSummarizerManager", + "DefaultSessionSummary", + "CheckSummarizerFunction", + "set_summarizer_token_threshold", + "set_summarizer_events_count_threshold", + "set_summarizer_time_interval_threshold", + "set_summarizer_important_content_threshold", + "set_summarizer_conversation_threshold", + "set_summarizer_check_functions_by_and", + "set_summarizer_check_functions_by_or", ] diff --git a/trpc_agent_sdk/sessions/compact/_base_manager.py b/trpc_agent_sdk/sessions/compact/_base_manager.py deleted file mode 100644 index 2a861ec8e..000000000 --- a/trpc_agent_sdk/sessions/compact/_base_manager.py +++ /dev/null @@ -1,52 +0,0 @@ -# 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/sessions/compact/_callbacks.py b/trpc_agent_sdk/sessions/compact/_callbacks.py deleted file mode 100644 index a678fffb1..000000000 --- a/trpc_agent_sdk/sessions/compact/_callbacks.py +++ /dev/null @@ -1,51 +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. -"""Shared callback installation and stage ordering for Advanced Memory.""" - -from __future__ import annotations - -from typing import Any - -from ._runtime import SessionCompactRuntime - - -def install_staged_callback( - agent: Any, - callback: Any, - *, - callback_type: type, - component_attribute: str, - memory_runtime: SessionCompactRuntime, - conflict_message: str, -) -> Any | None: - """Install a staged callback idempotently and validate runtime ownership.""" - existing = agent.before_model_callback - callbacks = existing if isinstance(existing, list) else ([existing] if existing else []) - for item in callbacks: - if not isinstance(item, callback_type): - continue - component = getattr(item, component_attribute) - if component.runtime is not memory_runtime: - raise ValueError(conflict_message) - return component - stage = getattr(callback, "advanced_memory_stage", None) - if not isinstance(stage, int): - raise TypeError("advanced_memory_stage must be an integer") - - def get_stage(item: Any) -> int: - item_stage = getattr(item, "advanced_memory_stage", 0) - return item_stage if isinstance(item_stage, int) else 0 - - insertion_index = next( - (index for index, item in enumerate(callbacks) if get_stage(item) > stage), - len(callbacks), - ) - agent.before_model_callback = [ - *callbacks[:insertion_index], - callback, - *callbacks[insertion_index:], - ] - return None diff --git a/trpc_agent_sdk/sessions/compact/_config.py b/trpc_agent_sdk/sessions/compact/_config.py deleted file mode 100644 index 1692bf2d6..000000000 --- a/trpc_agent_sdk/sessions/compact/_config.py +++ /dev/null @@ -1,130 +0,0 @@ -# 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/sessions/compact/_manager.py b/trpc_agent_sdk/sessions/compact/_manager.py deleted file mode 100644 index f969d0cda..000000000 --- a/trpc_agent_sdk/sessions/compact/_manager.py +++ /dev/null @@ -1,135 +0,0 @@ -# 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/sessions/compact/_runtime.py b/trpc_agent_sdk/sessions/compact/_runtime.py deleted file mode 100644 index f5734e2c5..000000000 --- a/trpc_agent_sdk/sessions/compact/_runtime.py +++ /dev/null @@ -1,53 +0,0 @@ -# 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/sessions/compact/advanced/__init__.py b/trpc_agent_sdk/sessions/compact/advanced/__init__.py new file mode 100644 index 000000000..6d20ea48c --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/__init__.py @@ -0,0 +1,50 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Advanced compact session manager.""" + +from ._auto_compact import AdvancedAutoCompactSummarizer +from ._base import BaseCompactSummarizerHandler +from ._base import BaseTokenEstimator +from ._base import BaseModelContextWindowResolver +from ._config import AutoCompactSummarizerConfig +from ._config import HistorySnipConfig +from ._config import TokenContextTrackerConfig +from ._config import MicroCompactConfig +from ._config import ToolResultBudgetConfig +from ._config import SessionMemoryExtractorConfig +from ._config import AdvancedAutoCompactSummarizerConfig +from ._filters import AdvancedAutoCompactSummarizerFilter +from ._formats import SessionMemoryDocument +from ._history_snip import HistorySnip +from ._micro_compact import MicroCompact +from ._manager import AdvancedAutoCompactSummarizerManager +from ._compaction_memory_extractor import SessionMemoryExtractor +from ._runtime import AdvancedAutoCompactSummarizerRuntime +from ._token_budget import TokenContextTracker +from ._tool_result_budget import ToolResultBudget + +__all__ = [ + "AdvancedAutoCompactSummarizer", + "AdvancedAutoCompactSummarizerManager", + "BaseCompactSummarizerHandler", + "BaseTokenEstimator", + "BaseModelContextWindowResolver", + "AutoCompactSummarizerConfig", + "HistorySnipConfig", + "TokenContextTrackerConfig", + "MicroCompactConfig", + "ToolResultBudgetConfig", + "SessionMemoryExtractorConfig", + "AdvancedAutoCompactSummarizerConfig", + "AdvancedAutoCompactSummarizerFilter", + "HistorySnip", + "MicroCompact", + "SessionMemoryDocument", + "SessionMemoryExtractor", + "AdvancedAutoCompactSummarizerRuntime", + "TokenContextTracker", + "ToolResultBudget", +] diff --git a/trpc_agent_sdk/sessions/compact/_autocompact.py b/trpc_agent_sdk/sessions/compact/advanced/_auto_compact.py similarity index 57% rename from trpc_agent_sdk/sessions/compact/_autocompact.py rename to trpc_agent_sdk/sessions/compact/advanced/_auto_compact.py index af4630b5a..4b43f0c0b 100644 --- a/trpc_agent_sdk/sessions/compact/_autocompact.py +++ b/trpc_agent_sdk/sessions/compact/advanced/_auto_compact.py @@ -9,43 +9,42 @@ import asyncio import copy -import hashlib import json import re import uuid from dataclasses import dataclass -from typing import Any -from typing import Protocol -from typing import TYPE_CHECKING +from typing_extensions import override -from trpc_agent_sdk.agents import LlmAgent +from trpc_agent_sdk.abc import CompactSummarizerABC +from trpc_agent_sdk.abc import RequestABC +from trpc_agent_sdk.abc import ResponseABC +from trpc_agent_sdk.abc import SessionABC +from trpc_agent_sdk.context import InvocationContext from trpc_agent_sdk.events import Event +from trpc_agent_sdk.models import LlmRequest from trpc_agent_sdk.models import LlmResponse -from trpc_agent_sdk.runners import Runner -from trpc_agent_sdk.sessions import InMemorySessionService +from trpc_agent_sdk.models import LLMModel from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import Part -from ._callbacks import install_staged_callback +from ._base import BaseCompactSummarizerHandler +from ._config import AdvancedAutoCompactSummarizerConfig from ._formats import SESSION_MEMORY_SECTIONS from ._formats import SESSION_MEMORY_STATE_KEY from ._formats import SessionMemoryDocument from ._formats import parse_session_memory_state from ._history_snip import estimate_request_chars -from ._runtime import SessionCompactRuntime +from ._runtime import AdvancedAutoCompactSummarizerRuntime from ._token_budget import TokenContextTracker +from ._compaction_memory_extractor import SessionMemoryExtractor +from ._utils import content_signature +from ._utils import internal_compaction_call -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_BLOCKED_MESSAGE = ( +ADVANCED_AUTOCOMPACT_BLOCKED_MESSAGE = ( "Automatic context compaction has failed repeatedly and the request is near the hard context limit. " "To avoid sending a request that will certainly fail, reduce the input, start a new session, " "or manually organize session memory before retrying.") -AUTOCOMPACT_SUMMARY_PREFIX = """This session is being continued from a compacted context. +ADVANCED_AUTOCOMPACT_SUMMARY_PREFIX = """This session is being continued from a compacted context. The following summary contains the important information from earlier messages. The complete original events remain available in the SessionService. @@ -65,7 +64,7 @@ @dataclass(frozen=True) -class AutoCompactRecord: +class AdvancedAutoCompactRecord: """Store stable replay information for the latest successful compaction.""" boundary_signature: str @@ -77,15 +76,15 @@ class AutoCompactRecord: @dataclass -class AutoCompactState: +class AdvancedAutoCompactState: """Store the latest compaction record and consecutive failure count.""" - latest_compaction: AutoCompactRecord | None + latest_compaction: AdvancedAutoCompactRecord | None consecutive_failures: int @dataclass(frozen=True) -class AutoCompactResult: +class AdvancedAutoCompactResult: """Summarize one compaction, replay, or hard-block result.""" compacted: bool @@ -99,52 +98,7 @@ class AutoCompactResult: request_tokens_before: int | None = None request_tokens_after: int | None = None token_source: str | None = None - - -class LegacySummaryGenerator(Protocol): - """Define the replaceable legacy compaction summary interface.""" - - async def generate(self, history: str, ctx: "InvocationContext") -> str: - """Return a workable Markdown summary for bounded old history.""" - - -def content_signature(content: Content) -> str: - """Generate a stable signature that preserves message identity.""" - parts: list[dict[str, Any]] = [] - for part in content.parts or []: - if part.text is not None: - parts.append({ - "type": "text", - "sha256": hashlib.sha256(part.text.encode("utf-8")).hexdigest(), - }) - elif part.function_call is not None: - parts.append({ - "type": "function_call", - "id": getattr(part.function_call, "id", None), - "name": part.function_call.name, - }) - elif part.function_response is not None: - parts.append({ - "type": "function_response", - "id": getattr(part.function_response, "id", None), - "name": part.function_response.name, - }) - elif part.executable_code is not None: - parts.append({"type": "executable_code"}) - elif part.code_execution_result is not None: - parts.append({"type": "code_execution_result"}) - else: - parts.append({"type": "other"}) - serialized = json.dumps( - { - "role": content.role, - "parts": parts - }, - ensure_ascii=False, - sort_keys=True, - separators=(",", ":"), - ) - return hashlib.sha256(serialized.encode("utf-8")).hexdigest() + summary: str | None = None def _content_text(content: Content) -> str: @@ -158,122 +112,108 @@ def _content_text(content: Content) -> str: ) -class ForkedLegacySummaryGenerator: - """Call a tool-free legacy summary Agent through an isolated Runner.""" +class AdvancedAutoCompactSummarizer(CompactSummarizerABC): + """Compact with session memory first, then fall back to a legacy summary.""" - def __init__(self, model: Any | None = None) -> None: - """Store an optional dedicated model, falling back to the parent model.""" + def __init__( + self, + config: AdvancedAutoCompactSummarizerConfig | None = None, + *, + model: LLMModel | None = None, + session_memory_extractor: SessionMemoryExtractor | None = None, + ) -> None: + """Initialize the compressor, summary generator, and session locks.""" self._model = model + self._runtime = AdvancedAutoCompactSummarizerRuntime(config=config or AdvancedAutoCompactSummarizerConfig()) + self._auto_compact_config = config.auto_compact + self._session_memory_extractor = self._create_extractor(session_memory_extractor) + self._states: dict[str, AdvancedAutoCompactState] = {} + self._session_locks: dict[str, asyncio.Lock] = {} + self._scoped_processors: dict[object, "AdvancedAutoCompactSummarizer"] = {} - def _resolve_model(self, ctx: "InvocationContext") -> Any: + def _resolve_model(self, ctx: InvocationContext) -> LLMModel: """Resolve the model used for legacy compaction.""" - model = self._model or getattr(ctx.agent, "model", None) - if not model: + if self._model is not None: + return self._model + if ctx.agent is None: raise ValueError("Autocompact summary generator cannot resolve an LLM model") - return model - - async def generate(self, history: str, ctx: "InvocationContext") -> str: - """Generate a summary in a temporary session without parent callbacks.""" - config = ctx.agent.generate_content_config if isinstance(ctx.agent, LlmAgent) else None - agent = LlmAgent( - name="advanced_autocompact_summarizer", - description="Generate an isolated context-compaction summary.", - instruction=LEGACY_SUMMARY_INSTRUCTION, - model=self._resolve_model(ctx), - tools=[], - generate_content_config=config, - add_name_to_instruction=False, - ) - app_name = f"{ctx.app_name}_advanced_autocompact" - runner = Runner( - app_name=app_name, - agent=agent, - session_service=InMemorySessionService(), - enable_post_turn_processing=False, - ) - last_event = None - try: - session = await runner.session_service.create_session( - app_name=app_name, - user_id="advanced-autocompact", - state={}, - ) - prompt = ("Compress the following old conversation. The input may contain JSON representations " - "of tool calls and results:\n\n" - f"\n{history}\n") - async for event in runner.run_async( - user_id=session.user_id, - session_id=session.id, - new_message=Content(role="user", parts=[Part.from_text(text=prompt)]), + return ctx.agent.model + + async def _generate_summary(self, history: str, ctx: InvocationContext | None = None) -> str: + """Generate a summary using the LLM model. + + Args: + history: The conversation text to summarize + + Returns: + Generated summary text + """ + request = LlmRequest() + request.append_instructions([LEGACY_SUMMARY_INSTRUCTION]) + prompt = ("Compress the following old conversation. The input may contain JSON representations " + "of tool calls and results:\n\n" + f"\n{history}\n") + request.contents.append(Content(role="user", parts=[Part.from_text(text=prompt)])) + + output = "" + with internal_compaction_call(getattr(ctx, "agent_context", None)): + async for llm_response in self._resolve_model(ctx).generate_async( + request, + stream=False, + ctx=ctx, ): - if not event.partial: - last_event = event - finally: - await runner.close() - if not last_event or not last_event.content or not last_event.content.parts: - raise ValueError("Autocompact summary generator returned no final content") - output = "\n".join(part.text for part in last_event.content.parts if part.text).strip() + if llm_response.content and llm_response.content.parts: + for part in llm_response.content.parts: + if part.text: + output += part.text + output = output.strip() + if not output: + raise ValueError("AdvancedAutoCompactSummarizer returned no final content") summary_match = re.search( r"\s*(.*?)\s*", output, flags=re.DOTALL | re.IGNORECASE, ) if summary_match is None or not summary_match.group(1).strip(): - raise ValueError("Autocompact summary generator returned no block") + raise ValueError("AdvancedAutoCompactSummarizer returned no block") return summary_match.group(1).strip() + def _create_extractor(self, + session_memory_extractor: SessionMemoryExtractor | None = None) -> SessionMemoryExtractor: + """Create the session memory extractor.""" + if session_memory_extractor is not None: + return session_memory_extractor + return SessionMemoryExtractor( + runtime=self._runtime, + model=self._model, + ) -class AutoCompact: - """Compact with session memory first, then fall back to a legacy summary.""" - - def __init__( - self, - 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 session_memory_extractor(self) -> SessionMemoryExtractor: + """Return the session memory extractor.""" + return self._session_memory_extractor @property - def runtime(self) -> SessionCompactRuntime: + def runtime(self) -> AdvancedAutoCompactSummarizerRuntime: """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.""" - key = self._runtime.session_key(session_id) if hasattr(self._runtime, "session_key") else session_id + key = self._runtime.session_key(session_id) lock = self._session_locks.get(key) if lock is None: lock = asyncio.Lock() self._session_locks[key] = lock return lock - async def _load_state(self, session_id: str) -> AutoCompactState: + async def _load_state(self, session_id: str) -> AdvancedAutoCompactState: """Restore process-local compaction state.""" - state_key = self._runtime.session_key(session_id) if hasattr(self._runtime, "session_key") else session_id + state_key = self._runtime.session_key(session_id) state = self._states.get(state_key) if state is not None: return state - state = AutoCompactState(latest_compaction=None, consecutive_failures=0) + state = AdvancedAutoCompactState(latest_compaction=None, consecutive_failures=0) self._states[state_key] = state return state @@ -281,10 +221,10 @@ def _summary_content(self, summary: str) -> Content: """Wrap a compaction summary in stable model-visible user content.""" return Content( role="user", - parts=[Part.from_text(text=AUTOCOMPACT_SUMMARY_PREFIX + summary)], + parts=[Part.from_text(text=ADVANCED_AUTOCOMPACT_SUMMARY_PREFIX + summary)], ) - def _summary_with_recovery_path(self, summary: str, session_id: str) -> str: + def _summary_with_recovery_path(self, summary: str) -> str: """Tell the model where the authoritative compacted data lives.""" return (f"{summary.rstrip()}\n\n" "For exact content from before compaction, read the original " @@ -354,7 +294,7 @@ def _compaction_start(self, contents: list[Content], boundary_index: int) -> int """Return the retained-content start for a legacy compaction.""" start = min( boundary_index + 1, - len(contents) - self._runtime.config.autocompact_keep_recent_contents, + len(contents) - self._auto_compact_config.keep_recent_contents, ) return self._adjust_start_for_tool_pairing(contents, start) @@ -362,7 +302,7 @@ def _session_memory_compaction_start(self, boundary_index: int) -> int: """Drop everything through the session-memory checkpoint boundary.""" return boundary_index + 1 - def _apply_record(self, request: "LlmRequest", record: AutoCompactRecord) -> bool: + def _apply_record(self, request: LlmRequest, record: AdvancedAutoCompactRecord) -> bool: """Replay a persisted compaction record into a rebuilt request.""" boundary_index = self._find_signature_index( request.contents, @@ -381,12 +321,11 @@ def _apply_record(self, request: "LlmRequest", record: AutoCompactRecord) -> boo async def _latest_session_memory_record( self, - session_id: str, - ctx: "InvocationContext", + 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) + state = ctx.session.state + parsed = parse_session_memory_state(state.get(SESSION_MEMORY_STATE_KEY)) if parsed is None: return None document, checkpoint, _ = parsed @@ -403,14 +342,14 @@ async def _latest_session_memory_record( def _compact_with_summary( self, - request: "LlmRequest", + request: LlmRequest, *, summary: str, boundary_index: int, source: str, strict_boundary: bool = False, boundary_event_id: str | None = None, - ) -> AutoCompactRecord: + ) -> AdvancedAutoCompactRecord: """Replace the old prefix with a summary and return a replay record.""" boundary_signature = content_signature(request.contents[boundary_index]) boundary_occurrence = self._signature_occurrence( @@ -421,24 +360,24 @@ def _compact_with_summary( start = (self._session_memory_compaction_start(boundary_index) if strict_boundary else self._compaction_start( request.contents, boundary_index)) request.contents = [self._summary_content(summary), *request.contents[start:]] - return AutoCompactRecord( + return AdvancedAutoCompactRecord( boundary_signature, boundary_occurrence, summary, source, boundary_event_id, - f"autocompact:{uuid.uuid4().hex}", + f"advanced_autocompact:{uuid.uuid4().hex}", ) def _resolve_boundary_event_id( self, - ctx: "InvocationContext", + 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 []: + for event in ctx.session.events: content = getattr(event, "content", None) if content is None or content_signature(content) != signature: continue @@ -448,15 +387,13 @@ def _resolve_boundary_event_id( 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: + 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 - ] + content_events = [event for event in ctx.session.events if event.content is not None] if len(content_events) <= 1: return None keep_count = min( - self._runtime.config.autocompact_keep_recent_contents, + self._auto_compact_config.keep_recent_contents, len(content_events) - 1, ) boundary_index = len(content_events) - keep_count - 1 @@ -469,13 +406,13 @@ def _legacy_boundary_event_id(self, ctx: "InvocationContext") -> str | None: async def _persist_session_compaction( self, - ctx: "InvocationContext", - record: AutoCompactRecord, + ctx: InvocationContext, + record: AdvancedAutoCompactRecord, ) -> None: """Persist the compacted active window through the original SessionService.""" - compact_events = getattr(ctx.session, "compact_events", None) + compact_events = ctx.session.compact_events if not callable(compact_events): - # AutoCompact remains usable as a request-only primitive in unit + # AdvancedAutoCompactSummarizer remains usable as a request-only primitive in unit # tests and custom integrations. The standard Manager supplies # the framework Session and persists the compacted window. return @@ -486,9 +423,9 @@ async def _persist_session_compaction( record.boundary_occurrence, ) if boundary_event_id is None: - raise ValueError("Cannot map the AutoCompact boundary to an active Session Event") + raise ValueError("Cannot map the AdvancedAutoCompactSummarizer boundary to an active Session Event") - compaction_id = record.compaction_id or f"autocompact:{uuid.uuid4().hex}" + compaction_id = record.compaction_id or f"advanced_autocompact:{uuid.uuid4().hex}" summary_event = Event( invocation_id="summary", author="system", @@ -519,7 +456,7 @@ async def _persist_session_compaction( def _bounded_history(self, contents: list[Content]) -> str: """Bound old history to the configured summary-input character limit.""" rendered = "\n".join(f"\n{_content_text(content)}\n" for content in contents) - limit = self._runtime.config.autocompact_summary_input_max_chars + limit = self._auto_compact_config.summary_input_max_chars if len(rendered) <= limit: return rendered marker = "\n...[middle of old history omitted due to the summary input limit]...\n" @@ -530,15 +467,15 @@ def _bounded_history(self, contents: list[Content]) -> str: async def _legacy_summary( self, contents: list[Content], - ctx: "InvocationContext", + ctx: InvocationContext, ) -> str: """Shrink old history across retries and generate a legacy summary.""" - retries = self._runtime.config.autocompact_summary_retries + retries = self._auto_compact_config.summary_retries_count working = list(contents) last_error: Exception | None = None for attempt in range(retries): try: - return await self._summary_generator.generate( + return await self._generate_summary( self._bounded_history(working), ctx, ) @@ -550,16 +487,124 @@ async def _legacy_summary( working = working[drop_count:] raise RuntimeError("Legacy autocompact summary failed after retries") from last_error + def _request_from_events(self, events: list[ResponseABC]) -> LlmRequest: + """Build the model-visible request view used by end-of-turn compaction.""" + contents: list[Content] = [] + for event in events: + is_model_visible = getattr(event, "is_model_visible", None) + if callable(is_model_visible) and not is_model_visible(): + continue + content = getattr(event, "content", None) + if content is not None: + contents.append(content.model_copy(deep=True)) + return LlmRequest(contents=contents) + + @override + async def should_summarize(self, session: SessionABC) -> bool: + """Check the character threshold without mutating the Session.""" + if not self._auto_compact_config.enabled: + return False + request = self._request_from_events(list(getattr(session, "events", []) or [])) + if len(request.contents) <= self._auto_compact_config.keep_recent_contents: + return False + + token_config = self._runtime.config.token_context_tracker + if token_config.enabled and token_config.model_context_window_tokens is not None: + effective_window = token_config.model_context_window_tokens - token_config.max_output_tokens + threshold = int(effective_window * token_config.auto_compact_ratio) + return TokenContextTracker(token_config).estimate_request_tokens(request) >= threshold + return estimate_request_chars(request) >= self._auto_compact_config.trigger_chars + + @override + async def create_session_summary_by_events( + self, + events: list[ResponseABC], + session_id: str, + keep_recent_count: int = 10, + ctx: InvocationContext | None = None, + historical_events: list[ResponseABC] | None = None, + store_historical_events: bool = False, + ) -> tuple[str | None, list[ResponseABC]]: + """Compact Events through the existing request-compaction algorithm.""" + del keep_recent_count + if ctx is None: + raise ValueError("Invocation context is required for advanced compaction") + if session_id != ctx.session_id: + raise ValueError("Session ID does not match the invocation context") + + request = self._request_from_events(events) + result = await self.apply(request, ctx=ctx, force=True) + if result.summary is not None: + events[:] = list(ctx.session.events) + if store_historical_events and historical_events is not None: + historical_events[:] = list(ctx.session.historical_events) + return result.summary, events + + @override + async def create_session_summary( + self, + session: SessionABC, + ctx: InvocationContext | None = None, + store_historical_events: bool = False, + ) -> str | None: + """Compact one Session and persist its active and historical Events.""" + if ctx is None: + raise ValueError("Invocation context is required for advanced compaction") + events = getattr(session, "events", None) + historical_events = getattr(session, "historical_events", None) + if not isinstance(events, list) or not isinstance(historical_events, list): + raise TypeError("Advanced compaction requires a Session with Event history") + summary, _ = await self.create_session_summary_by_events( + events, + session.id, + ctx=ctx, + historical_events=historical_events, + store_historical_events=store_historical_events, + ) + return summary + + @override + async def create_session_summary_by_request( + self, + request: RequestABC, + ctx: InvocationContext | None = None, + force: bool = False, + ) -> LlmResponse | None: + """Compact a built model request immediately before generation.""" + if ctx is None: + raise ValueError("Invocation context is required for advanced compaction") + if not isinstance(request, LlmRequest): + raise TypeError("Advanced compaction requires an LlmRequest") + + result = await self.apply(request, ctx=ctx, force=force) + if not result.blocked: + TokenContextTracker.record_request_context(request, ctx) + return None + return LlmResponse(content=Content( + role="model", + parts=[Part.from_text(text=ADVANCED_AUTOCOMPACT_BLOCKED_MESSAGE)], + )) + + @override + def get_summary_metadata(self) -> dict[str, object]: + """Return advanced compaction configuration metadata.""" + return { + "strategy": "advanced", + "auto_compact_enabled": self._auto_compact_config.enabled, + "trigger_chars": self._auto_compact_config.trigger_chars, + "keep_recent_contents": self._auto_compact_config.keep_recent_contents, + } + async def apply( self, - request: "LlmRequest", + request: LlmRequest, *, - session_id: str, - ctx: "InvocationContext", + ctx: InvocationContext, force: bool = False, - ) -> AutoCompactResult: + ) -> AdvancedAutoCompactResult: """Run compaction against the current session's tenant namespace.""" - if hasattr(self._runtime, "scope"): + session_id = ctx.session_id + if self._runtime.scope: return await self._apply_scoped(request, session_id=session_id, ctx=ctx, force=force) runtime = self._runtime.for_session(ctx.session) processor = self._scoped_processors.get(runtime.scope) @@ -569,22 +614,32 @@ async def apply( processor._states = {} processor._session_locks = {} self._scoped_processors[runtime.scope] = processor - return await processor.apply(request, session_id=session_id, ctx=ctx, force=force) + return await processor.apply(request, ctx=ctx, force=force) async def _apply_scoped( self, - request: "LlmRequest", + request: LlmRequest, *, session_id: str, - ctx: "InvocationContext", + ctx: InvocationContext, force: bool = False, - ) -> AutoCompactResult: + ) -> AdvancedAutoCompactResult: """Replay old compaction and compact again when pressure is high.""" config = self._runtime.config - tracker = TokenContextTracker(config) - if not config.enabled or not config.autocompact_enabled: + auto_compact_config = config.auto_compact + tracker = TokenContextTracker(config.token_context_tracker) + if not auto_compact_config.enabled: request_chars = estimate_request_chars(request) - return AutoCompactResult(False, False, False, None, request_chars, request_chars, 0) + return AdvancedAutoCompactResult(compacted=False, + reapplied=False, + blocked=False, + source=None, + request_chars_before=request_chars, + request_chars_after=request_chars, + consecutive_failures=0, + request_tokens_before=None, + request_tokens_after=None, + token_source=None) async with self._session_lock(session_id): request.contents = [content.model_copy(deep=True) for content in request.contents] state = await self._load_state(session_id) @@ -598,53 +653,49 @@ async def _apply_scoped( 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: - return AutoCompactResult( - False, - reapplied, - True, - None, - request_chars_before, - request_chars_before, - state.consecutive_failures, + if token_mode else request_chars_before >= self._auto_compact_config.blocking_chars) + if state.consecutive_failures >= auto_compact_config.max_failures and blocking_reached: + return AdvancedAutoCompactResult( + compacted=False, + reapplied=reapplied, + blocked=True, + source=None, + request_chars_before=request_chars_before, + request_chars_after=request_chars_before, + consecutive_failures=state.consecutive_failures, ) - if state.consecutive_failures >= config.autocompact_max_failures: - return AutoCompactResult( - False, - reapplied, - False, - None, - request_chars_before, - request_chars_before, - state.consecutive_failures, + if state.consecutive_failures >= auto_compact_config.max_failures: + return AdvancedAutoCompactResult( + compacted=False, + reapplied=reapplied, + blocked=False, + source=None, + request_chars_before=request_chars_before, + request_chars_after=request_chars_before, + consecutive_failures=state.consecutive_failures, ) - autocompact_reached = (request_tokens_before >= token_budget_before.autocompact_threshold_tokens - if token_mode else request_chars_before >= config.autocompact_trigger_chars) - if not force and not autocompact_reached: - return AutoCompactResult( - False, - reapplied, - False, - state.latest_compaction.source if reapplied and state.latest_compaction else None, - request_chars_before, - request_chars_before, - state.consecutive_failures, + auto_compact_reached = (request_tokens_before >= token_budget_before.auto_compact_threshold_tokens + if token_mode else request_chars_before >= auto_compact_config.trigger_chars) + if not force and not auto_compact_reached: + return AdvancedAutoCompactResult( + compacted=False, + reapplied=reapplied, + blocked=False, + source=state.latest_compaction.source if reapplied and state.latest_compaction else None, + request_chars_before=request_chars_before, + request_chars_after=request_chars_before, + consecutive_failures=state.consecutive_failures, ) original_contents = [content.model_copy(deep=True) for content in request.contents] try: - compact_record: AutoCompactRecord | None = None + compact_record: AdvancedAutoCompactRecord | None = None 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, - ) + session_memory = await self._latest_session_memory_record(ctx, ) if session_memory is not None: memory, boundary_signature, boundary_occurrence, boundary_event_id = session_memory boundary_index = self._find_signature_index( @@ -660,10 +711,7 @@ async def _apply_scoped( if boundary_index is not None: compact_record = self._compact_with_summary( request, - summary=self._summary_with_recovery_path( - memory, - session_id, - ), + summary=self._summary_with_recovery_path(memory), boundary_index=boundary_index, source="session-memory", strict_boundary=True, @@ -673,14 +721,14 @@ async def _apply_scoped( target_reached = (tracker.budget(request, ctx).estimate.tokens <= token_budget_before.warning_threshold_tokens) else: - target_reached = estimate_request_chars(request) <= config.autocompact_target_chars + target_reached = estimate_request_chars(request) <= auto_compact_config.target_chars if not target_reached: request.contents = [content.model_copy(deep=True) for content in original_contents] compact_record = None if compact_record is None: keep_count = min( - config.autocompact_keep_recent_contents, + auto_compact_config.keep_recent_contents, max(1, len(request.contents) - 1), ) @@ -693,10 +741,7 @@ async def _apply_scoped( ) compact_record = self._compact_with_summary( request, - summary=self._summary_with_recovery_path( - summary, - session_id, - ), + summary=self._summary_with_recovery_path(summary), boundary_index=boundary_index, source="legacy", boundary_event_id=self._legacy_boundary_event_id(ctx), @@ -707,100 +752,60 @@ async def _apply_scoped( 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") + raise ValueError("Advanced Auto Compact did not reduce request token estimate") elif request_chars_after >= request_chars_before: - raise ValueError("Autocompact did not reduce request size") + raise ValueError("Advanced Auto Compact did not reduce request size") await self._persist_session_compaction(ctx, compact_record) state.latest_compaction = compact_record state.consecutive_failures = 0 - return AutoCompactResult( - True, - reapplied, - False, - compact_record.source, - request_chars_before, - request_chars_after, - 0, + return AdvancedAutoCompactResult( + compacted=True, + reapplied=reapplied, + blocked=False, + source=compact_record.source, + request_chars_before=request_chars_before, + request_chars_after=request_chars_after, + consecutive_failures=0, request_tokens_before=comparison_tokens_before if token_mode else None, request_tokens_after=comparison_tokens_after if token_mode else None, token_source="estimated" if token_mode else None, + summary=compact_record.summary, ) except Exception as exc: # noqa: BLE001 request.contents = original_contents state.consecutive_failures += 1 - blocked = state.consecutive_failures >= config.autocompact_max_failures and blocking_reached - return AutoCompactResult( - False, - reapplied, - blocked, - None, - request_chars_before, - request_chars_before, - state.consecutive_failures, - error=str(exc), + blocked = state.consecutive_failures >= auto_compact_config.max_failures and blocking_reached + return AdvancedAutoCompactResult( + compacted=False, + reapplied=reapplied, + blocked=blocked, + source=None, + request_chars_before=request_chars_before, + request_chars_after=request_chars_before, + consecutive_failures=state.consecutive_failures, + error=str(exc) if exc else None if blocked else None, request_tokens_before=request_tokens_before if token_mode else None, - request_tokens_after=request_tokens_before if token_mode else None, + request_tokens_after=comparison_tokens_before if token_mode else None, token_source=token_budget_before.estimate.source if token_mode else None, ) -class AutoCompactCallback: - """Adapt the automatic compressor to before_model_callback.""" - - advanced_memory_stage = 40 +class AdvancedAutoCompactSummarizerHandler(BaseCompactSummarizerHandler): + """Advanced auto compact summarizer handler.""" - def __init__(self, autocompact: AutoCompact) -> None: - """Store the compressor executed before model requests.""" - self._autocompact = autocompact - - @property - def autocompact(self) -> AutoCompact: - """Return the compressor used by this callback.""" - return self._autocompact - - async def __call__( + @override + async def handle( self, - ctx: "InvocationContext", - request: "LlmRequest", + ctx: InvocationContext, + request: LlmRequest, + force: bool = False, ) -> LlmResponse | None: """Compact before each request and return a local block after failures.""" - result = await self._autocompact.apply( + summarizer = self.get_summarizer(ctx) + if not isinstance(summarizer, AdvancedAutoCompactSummarizer): + raise ValueError("Summarizer is not an AdvancedAutoCompactSummarizer") + return await summarizer.create_session_summary_by_request( request, - session_id=ctx.session_id, ctx=ctx, + force=force, ) - if not result.blocked: - TokenContextTracker(self._autocompact.runtime.config).record_request_context( - request, - ctx, - ) - return None - return LlmResponse(content=Content( - role="model", - parts=[Part.from_text(text=AUTOCOMPACT_BLOCKED_MESSAGE)], - )) - - -def setup_autocompact( - agent: "ParentLlmAgent", - memory_runtime: SessionCompactRuntime, - summary_generator: LegacySummaryGenerator | None = None, - *, - model: Any | None = None, -) -> AutoCompact: - """Install the automatic compaction callback in pipeline stage order.""" - autocompact = AutoCompact( - memory_runtime, - summary_generator, - model=model, - ) - callback = AutoCompactCallback(autocompact) - existing_autocompact = install_staged_callback( - agent, - callback, - callback_type=AutoCompactCallback, - component_attribute="autocompact", - memory_runtime=memory_runtime, - conflict_message="Autocompact is already configured with another runtime", - ) - return existing_autocompact or autocompact diff --git a/trpc_agent_sdk/sessions/compact/advanced/_base.py b/trpc_agent_sdk/sessions/compact/advanced/_base.py new file mode 100644 index 000000000..ec15b559a --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/_base.py @@ -0,0 +1,52 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. + +from __future__ import annotations + +from abc import ABC +from abc import abstractmethod +from typing import Any +from typing import Optional + +from trpc_agent_sdk.abc import CompactSummarizerABC +from trpc_agent_sdk.abc import CompactSummarizerManagerABC +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.models import LlmRequest + + +class BaseCompactSummarizerHandler(ABC): + """Base compact summarizer handler.""" + + def get_summarizer(self, ctx: InvocationContext) -> CompactSummarizerABC: + """Get the summarizer.""" + session_service = ctx.session_service + if session_service is None: + raise ValueError("Session service is not set") + summarizer_manager = getattr(session_service, "summarizer_manager", None) + if summarizer_manager is None or not isinstance(summarizer_manager, CompactSummarizerManagerABC): + raise ValueError("Summarizer manager is not an CompactSummarizerManagerABC") + return summarizer_manager.summarizer + + @abstractmethod + async def handle(self, ctx: InvocationContext, req: LlmRequest): + """Handle the compact summarizer.""" + pass + + +class BaseTokenEstimator(ABC): + """Define the replaceable token estimator interface.""" + + @abstractmethod + def estimate_payload_tokens(self, payload: Any) -> int: + """Estimate tokens for any JSON-compatible payload.""" + + +class BaseModelContextWindowResolver(ABC): + """Define the model-identifier context-window resolver interface.""" + + @abstractmethod + def resolve_context_window_tokens(self, model: Any) -> Optional[int]: + """Return the model context window, or None when unknown.""" diff --git a/trpc_agent_sdk/sessions/compact/_session_memory.py b/trpc_agent_sdk/sessions/compact/advanced/_compaction_memory_extractor.py similarity index 73% rename from trpc_agent_sdk/sessions/compact/_session_memory.py rename to trpc_agent_sdk/sessions/compact/advanced/_compaction_memory_extractor.py index 2e9961289..349c6ce0e 100644 --- a/trpc_agent_sdk/sessions/compact/_session_memory.py +++ b/trpc_agent_sdk/sessions/compact/advanced/_compaction_memory_extractor.py @@ -3,42 +3,40 @@ # Copyright (C) 2026 Tencent. All rights reserved. # # tRPC-Agent-Python is licensed under Apache-2.0. -"""Maintain structured session memory with an isolated sub-agent.""" +"""Maintain structured session memory with direct model generation.""" from __future__ import annotations import json +import re from collections import Counter from dataclasses import dataclass from dataclasses import fields +from dataclasses import field from datetime import datetime from datetime import timezone -import re from typing import Any -from typing import Protocol -from typing import TYPE_CHECKING +from typing import Optional -from trpc_agent_sdk.agents import LlmAgent +from trpc_agent_sdk.context import InvocationContext from trpc_agent_sdk.log import logger -from trpc_agent_sdk.memory import InMemoryMemoryService -from trpc_agent_sdk.runners import Runner -from trpc_agent_sdk.sessions import InMemorySessionService +from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.models import LLMModel from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import Part +from ..._session import Session + from ._formats import SESSION_MEMORY_SECTION_DESCRIPTIONS from ._formats import SESSION_MEMORY_SECTIONS from ._formats import SESSION_MEMORY_STATE_KEY from ._formats import SessionMemoryDocument from ._formats import build_session_memory_state from ._formats import parse_session_memory_state -from ._runtime import SessionCompactRuntime +from ._runtime import AdvancedAutoCompactSummarizerRuntime 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 +from ._utils import content_signature +from ._utils import internal_compaction_call _SESSION_MEMORY_FIELDS = tuple(field.name for field in fields(SessionMemoryDocument)) @@ -164,23 +162,12 @@ class SessionMemoryExtractionInput: class SessionMemoryExtractionResult: """Describe one incremental session-memory extraction.""" - extracted: bool - reason: str - processed_events: int = 0 - first_event_id: str | None = None - last_event_id: str | None = None - error: str | None = None - - -class SessionMemoryGenerator(Protocol): - """Define the replaceable session-memory generator interface.""" - - async def generate( - self, - extraction_input: SessionMemoryExtractionInput, - ctx: "InvocationContext", - ) -> SessionMemoryDocument: - """Generate a complete document from old memory and new context.""" + extracted: bool = field(default=False) + reason: str = field(default="") + processed_events: int = field(default=0) + first_event_id: Optional[str] = field(default=None) + last_event_id: Optional[str] = field(default=None) + error: Optional[str] = field(default=None) def has_session_memory_content(document: SessionMemoryDocument) -> bool: @@ -247,152 +234,99 @@ def limit(value: str, limit_chars: int = max_chars) -> str: return limited_document -class ForkedSessionMemoryGenerator: - """Call the isolated extraction Agent through a temporary Runner.""" +class SessionMemoryExtractor: + """Check thresholds and coordinate extraction, writes, and checkpoints.""" def __init__( self, - model: Any | None = None, - *, - section_max_chars: int = 8_000, - max_retries: int = 1, + runtime: AdvancedAutoCompactSummarizerRuntime, + model: LLMModel | None = None, ) -> None: - """Store an optional dedicated model, falling back to the parent model.""" - if max_retries < 0: - raise ValueError("max_retries must not be negative") + """Initialize extraction and per-session serialization locks.""" self._model = model - self._section_max_chars = section_max_chars - self._max_retries = max_retries + self._runtime = runtime + self._config = runtime.config.session_memory - def _resolve_model(self, ctx: "InvocationContext") -> Any: + def _resolve_model(self, ctx: InvocationContext) -> LLMModel: """Prefer the dedicated model, falling back to the parent Agent model.""" model = self._model or getattr(ctx.agent, "model", None) if not model: raise ValueError("Session memory extractor cannot resolve an LLM model") return model - async def generate( + async def _call_llm_model( self, extraction_input: SessionMemoryExtractionInput, - ctx: "InvocationContext", + ctx: InvocationContext, ) -> SessionMemoryDocument: - """Run extraction in a Runner isolated from the parent session and services.""" - config = ctx.agent.generate_content_config if isinstance(ctx.agent, LlmAgent) else None - agent = LlmAgent( - name="advanced_session_memory_extractor", - description="Update Markdown session memory in isolation.", - instruction=SESSION_MEMORY_INSTRUCTION, - model=self._resolve_model(ctx), - tools=[], - generate_content_config=config, - add_name_to_instruction=False, - ) - app_name = f"{ctx.app_name}_advanced_session_memory" - runner = Runner( - app_name=app_name, - agent=agent, - session_service=InMemorySessionService(), - memory_service=InMemoryMemoryService(), - enable_post_turn_processing=False, + """Generate session memory directly through the configured LLM model.""" + model = self._resolve_model(ctx) + prompt = build_session_memory_prompt( + extraction_input, + section_max_chars=self._config.section_max_chars, ) - try: - prompt = build_session_memory_prompt( - extraction_input, - section_max_chars=self._section_max_chars, - ) - parse_error: Exception | None = None - for attempt in range(self._max_retries + 1): - session = await runner.session_service.create_session( - app_name=app_name, - user_id="advanced-session-memory", - state={}, - ) - retry_instruction = "" - if parse_error is not None: - retry_instruction = ("\n\nThe previous response could not be parsed. " - f"Parser error: {parse_error}. Return the required short analysis " - "followed by all ten Markdown headings and their body text. " - "Do not return JSON, XML, or code fences.") - content = Content(role="user", parts=[Part.from_text(text=prompt + retry_instruction)]) - last_event = None - async for event in runner.run_async( - user_id=session.user_id, - session_id=session.id, - new_message=content, + parse_error: Exception | None = None + for attempt in range(self._config.max_retries + 1): + retry_instruction = "" + if parse_error is not None: + retry_instruction = ("\n\nThe previous response could not be parsed. " + f"Parser error: {parse_error}. Return the required short analysis " + "followed by all ten Markdown headings and their body text. " + "Do not return JSON, XML, or code fences.") + + request = LlmRequest( + contents=[Content( + role="user", + parts=[Part.from_text(text=prompt + retry_instruction)], + )], ) + request.append_instructions([SESSION_MEMORY_INSTRUCTION]) + + output = "" + with internal_compaction_call(getattr(ctx, "agent_context", None)): + async for response in model.generate_async( + request, + stream=False, + ctx=ctx, ): - if not event.partial: - last_event = event - - try: - if not last_event or not last_event.content or not last_event.content.parts: - raise ValueError("Session memory extractor returned no final content") - merged_text = "\n".join(part.text for part in last_event.content.parts if part.text) - return parse_session_memory_markdown(merged_text) - except Exception as exc: # noqa: BLE001 - parse_error = exc - if attempt >= self._max_retries: - raise - finally: - await runner.close() - - -class SessionMemoryExtractor: - """Check thresholds and coordinate extraction, writes, and checkpoints.""" - - def __init__( - self, - memory_runtime: 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: - raise ValueError("Provide either generator or model, not both") - self._runtime = memory_runtime - self._generator = generator or ForkedSessionMemoryGenerator( - model, - section_max_chars=memory_runtime.config.session_memory_section_max_chars, - ) - self._session_service = session_service - - @property - def runtime(self) -> SessionCompactRuntime: - """Return the runtime bound to this extractor.""" - return self._runtime + if response.content and response.content.parts: + output += "\n".join(part.text for part in response.content.parts if part.text) - 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 + try: + if not output.strip(): + raise ValueError("Session memory extractor returned no final content") + return parse_session_memory_markdown(output) + except Exception as exc: # noqa: BLE001 + parse_error = exc + if attempt >= self._config.max_retries: + raise - def _session_event_records(self, session: "SessionABC") -> list[dict[str, Any]]: + def _session_event_records(self, session: Session) -> list[dict[str, Any]]: """Convert the authoritative Session Events into extraction records.""" records: list[dict[str, Any]] = [] seen: set[str] = set() # Archived Events are no longer addressable in the active model # request. Their information is already represented by the active # summary Event included in the extraction context. - events = list(getattr(session, "events", None) or []) + events = list(session.events or []) for event in events: - is_summary_event = getattr(event, "is_summary_event", None) - if callable(is_summary_event) and is_summary_event(): + if event.is_summary_event and event.is_summary_event(): continue - event_id = getattr(event, "id", None) - if not isinstance(event_id, str) or event_id in seen: + if event.id in seen: continue - seen.add(event_id) - timestamp = float(getattr(event, "timestamp", 0.0) or 0.0) + seen.add(event.id) + timestamp = float(event.timestamp or 0.0) records.append({ - "kind": "event", - "event_id": event_id, - "recorded_at": datetime.fromtimestamp( + "kind": + "event", + "event_id": + event.id, + "recorded_at": + datetime.fromtimestamp( timestamp, tz=timezone.utc, ).isoformat(), - "event": event.model_dump( + "event": + event.model_copy(deep=True).model_dump( mode="json", by_alias=True, exclude_none=True, @@ -442,27 +376,25 @@ def _serialized_record(self, record: dict[str, Any]) -> str: default=str, ) - def _context_contents(self, ctx: "InvocationContext") -> list[Any]: + def _context_contents(self, ctx: InvocationContext) -> list[Content]: """Extract model-context Content without Event metadata.""" - override_messages = getattr(ctx, "override_messages", None) + override_messages = ctx.override_messages if isinstance(override_messages, list): return [content for content in override_messages if content is not None] - contents: list[Any] = [] - session = getattr(ctx, "session", None) - for event in getattr(session, "events", []) or []: - is_model_visible = getattr(event, "is_model_visible", None) + contents: list[Content] = [] + session = ctx.session + for event in session.events or []: + is_model_visible = event.is_model_visible if callable(is_model_visible) and not is_model_visible(): continue - content = getattr(event, "content", None) + content = event.content if content is not None: contents.append(content) return contents - def _serialized_context_content(self, content: Any) -> str | None: + def _serialized_context_content(self, content: Content) -> str | None: """Serialize visible message content while excluding hidden thoughts.""" - if not hasattr(content, "model_dump"): - return None payload = content.model_dump( mode="json", by_alias=True, @@ -493,7 +425,7 @@ def _excerpt_text(self, serialized: str, limit: int) -> str: side = max(1, (limit - len(marker)) // 2) return serialized[:side] + marker + serialized[-(limit - len(marker) - side):] - def _context_messages(self, ctx: "InvocationContext") -> list[str]: + def _context_messages(self, ctx: InvocationContext) -> list[str]: """Render the complete visible conversation context in order.""" messages: list[str] = [] for content in self._context_contents(ctx): @@ -545,24 +477,24 @@ def _fits_prompt_budget( self, extraction_input: SessionMemoryExtractionInput, tracker: TokenContextTracker, - ctx: "InvocationContext", + ctx: InvocationContext, ) -> bool: """Return whether the complete sub-agent prompt fits the input budget.""" prompt = build_session_memory_prompt( extraction_input, - section_max_chars=self._runtime.config.session_memory_section_max_chars, + section_max_chars=self._config.section_max_chars, ) effective_window = tracker.effective_context_window_tokens(ctx) if effective_window is not None: - limit = effective_window - self._runtime.config.session_memory_request_overhead_tokens + limit = effective_window - self._config.request_overhead_tokens return limit > 0 and tracker.estimate_payload_tokens(prompt) <= limit - return len(prompt) <= self._runtime.config.session_memory_prompt_max_chars + return len(prompt) <= self._config.prompt_max_chars def _build_extraction_input( self, current_memory: str, pending: list[dict[str, Any]], - ctx: "InvocationContext", + ctx: InvocationContext, tracker: TokenContextTracker, ) -> tuple[list[dict[str, Any]], SessionMemoryExtractionInput | None]: """Build the largest safe input that fits the extraction budget. @@ -642,14 +574,14 @@ def missing_context(end: int) -> list[str]: return [], None - async def _read_current_memory(self, session: "SessionABC") -> str: + async def _read_current_memory(self, session: Session) -> str: """Read Session Memory from the SessionService-owned state.""" parsed = parse_session_memory_state(session.state.get(SESSION_MEMORY_STATE_KEY)) return parsed[0].to_markdown() if parsed is not None else SessionMemoryDocument().to_markdown() def _state_checkpoint( self, - session: "SessionABC", + session: Session, ) -> tuple[dict[str, Any] | None, int | None]: """Read the checkpoint and token metric from Session.state.""" parsed = parse_session_memory_state(session.state.get(SESSION_MEMORY_STATE_KEY)) @@ -664,39 +596,38 @@ def _state_checkpoint( def _boundary_for_event( self, - session: "SessionABC", + session: Session, 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 []) + events = list(session.events 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: + if event.id in seen_ids: continue - seen_ids.add(current_id) - content = getattr(event, "content", None) + seen_ids.add(event.id) + content = event.content if content is None: continue signature = content_signature(content) signatures.append(signature) - if current_id == event_id: + if event.id == event_id: return signature, signatures.count(signature) return None async def _persist_checkpoint( self, - session: "SessionABC", + ctx: InvocationContext, included_records: list[dict[str, Any]], document: SessionMemoryDocument, context_tokens: int | None, ) -> None: """Persist the processed increment boundary after a successful write.""" + session = ctx.session first_event_id = included_records[0]["event_id"] last_event_id = included_records[-1]["event_id"] values = ( @@ -711,8 +642,6 @@ async def _persist_checkpoint( document.key_results, document.worklog, ) - 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") @@ -733,27 +662,40 @@ async def _persist_checkpoint( checkpoint=checkpoint, context_tokens=context_tokens, ) - await self._session_service.patch_session_state( - session, - {SESSION_MEMORY_STATE_KEY: payload}, - ) + state_delta = {SESSION_MEMORY_STATE_KEY: payload} + session_service = getattr(ctx, "session_service", None) + if session_service is None: + session.state.update(state_delta) + return + + update_state = getattr(session_service, "update_session_state", None) + if callable(update_state): + await update_state(session, state_delta) + return + + # Compatibility fallback for duck-typed SessionService implementations + # that do not inherit the latest SessionServiceABC. + session.state.update(state_delta) + await session_service.update_session(session) async def extract_if_needed( self, - session: "SessionABC", - ctx: "InvocationContext", - *, + ctx: InvocationContext, force: bool = False, ) -> SessionMemoryExtractionResult: """Update memory when the threshold or force flag is reached.""" - config = self._runtime.config - if not config.enabled or not config.session_memory_enabled: - return SessionMemoryExtractionResult(False, "disabled") + config = self._config + session = ctx.session + if not config.enabled: + return SessionMemoryExtractionResult(reason="disabled") runtime = self._runtime.for_session(session) session_key = runtime.session_key(session.id) - async with self._runtime.coordination.guard(session_key) as acquired: + async with self._runtime.coordination.guard( + session_key, + timeout=config.wait_timeout_seconds, + ) as acquired: if not acquired: - return SessionMemoryExtractionResult(False, "coordination-timeout") + return SessionMemoryExtractionResult(reason="coordination-timeout") records = self._session_event_records(session) checkpoint, checkpoint_context_tokens = self._state_checkpoint(session) checkpoint_event_id = checkpoint["last_event_id"] if checkpoint is not None else None @@ -764,26 +706,25 @@ async def extract_if_needed( checkpoint_recorded_at if isinstance(checkpoint_recorded_at, str) else None, ) if not pending: - return SessionMemoryExtractionResult(False, "no-new-events") + return SessionMemoryExtractionResult(reason="no-new-events") pending_chars = self._record_chars(pending) - tracker = TokenContextTracker(config) + tracker = TokenContextTracker(self._runtime.config.token_context_tracker) token_mode = tracker.token_mode_enabled(ctx) context_tokens = tracker.estimate_payload_tokens(self._context_contents(ctx)) - threshold = (config.session_memory_update_tokens if checkpoint_event_id is not None and token_mode else - (config.session_memory_initial_tokens if token_mode else - (config.session_memory_update_chars - if checkpoint_event_id is not None else config.session_memory_initial_chars))) + threshold = (config.update_tokens if checkpoint_event_id is not None and token_mode else + (config.initial_tokens if token_mode else + (config.update_chars if checkpoint_event_id is not None else config.initial_chars))) tool_calls = self._count_tool_calls(pending) natural_break = not self._last_event_has_tool_call(pending) if not natural_break: - return SessionMemoryExtractionResult(False, "unsafe-boundary") + return SessionMemoryExtractionResult(reason="unsafe-boundary") threshold_met = ((context_tokens >= threshold if checkpoint_context_tokens is None else (context_tokens < checkpoint_context_tokens or context_tokens - checkpoint_context_tokens >= threshold)) if token_mode else pending_chars >= threshold) - tool_condition_met = tool_calls >= config.session_memory_tool_calls_between_updates or natural_break + tool_condition_met = tool_calls >= config.tool_calls_between_updates or natural_break if not force and (not threshold_met or not tool_condition_met): - return SessionMemoryExtractionResult(False, "threshold-not-met") + return SessionMemoryExtractionResult(reason="threshold-not-met") included, extraction_input = self._build_extraction_input( await self._read_current_memory(session), @@ -792,18 +733,18 @@ async def extract_if_needed( tracker, ) if extraction_input is None: - return SessionMemoryExtractionResult(False, "context-unavailable") + return SessionMemoryExtractionResult(reason="context-unavailable") try: - document = await self._generator.generate(extraction_input, ctx) + document = await self._call_llm_model(extraction_input, ctx) if not has_session_memory_content(document): raise ValueError("Session memory generator returned an all-empty document") document = limit_session_memory_document( document, - max_chars=config.session_memory_section_max_chars, - total_max_chars=config.session_memory_total_max_chars, + max_chars=config.section_max_chars, + total_max_chars=config.total_max_chars, ) await self._persist_checkpoint( - session, + ctx, included, document, context_tokens, @@ -816,8 +757,7 @@ async def extract_if_needed( exc_info=True, ) return SessionMemoryExtractionResult( - False, - "extraction-failed", + reason="extraction-failed", processed_events=0, first_event_id=included[0]["event_id"], last_event_id=included[-1]["event_id"], @@ -825,8 +765,8 @@ async def extract_if_needed( ) return SessionMemoryExtractionResult( - True, - "forced" if force else "threshold-met", + extracted=True, + reason="forced" if force else "threshold-met", processed_events=len(included), first_event_id=included[0]["event_id"], last_event_id=included[-1]["event_id"], diff --git a/trpc_agent_sdk/sessions/compact/advanced/_config.py b/trpc_agent_sdk/sessions/compact/advanced/_config.py new file mode 100644 index 000000000..35f16b0d5 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/_config.py @@ -0,0 +1,174 @@ +# Tencent is pleased to support the open source ecosystem. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# Licensed under Apache-2.0. +"""Configuration for Session Compact.""" + +from __future__ import annotations + +from dataclasses import dataclass +from dataclasses import field + +from ._base import BaseTokenEstimator +from ._base import BaseModelContextWindowResolver + +DEFAULT_COMPACTABLE_TOOL_NAMES = ( + "Read", + "Bash", + "Grep", + "Glob", + "Search", + "CodeSearch", +) + + +def _require_positive(**values: int | float) -> None: + for name, value in values.items(): + if value <= 0: + raise ValueError(f"{name} must be greater than zero") + + +def _require_non_negative(**values: int | float) -> None: + for name, value in values.items(): + if value < 0: + raise ValueError(f"{name} must be non-negative") + + +def _require_non_empty_names(name: str, values: tuple[str, ...]) -> None: + if not values or any(not value.strip() for value in values): + raise ValueError(f"{name} must contain non-empty names") + + +@dataclass(frozen=True) +class AutoCompactSummarizerConfig: + """Configure auto compact summarizer.""" + enabled: bool = field(default=True) + trigger_chars: int = field(default=700_000) + target_chars: int = field(default=350_000) + blocking_chars: int = field(default=780_000) + keep_recent_contents: int = field(default=8) + max_failures: int = field(default=3) + summary_input_max_chars: int = field(default=600_000) + summary_retries_count: int = field(default=3) + + +@dataclass(frozen=True) +class HistorySnipConfig: + """Configure history snip.""" + enabled: bool = field(default=True) + trigger_chars: int = field(default=600_000) + target_chars: int = field(default=400_000) + keep_recent: int = field(default=5) + tool_names: tuple[str, ...] = field(default=DEFAULT_COMPACTABLE_TOOL_NAMES) + + +@dataclass(frozen=True) +class TokenContextTrackerConfig: + """Configure token context tracker.""" + enabled: bool = field(default=True) + warning_ratio: float = field(default=0.85) + auto_compact_ratio: float = field(default=0.90) + blocking_ratio: float = field(default=0.95) + model_context_window_tokens: int | None = field(default=None) + max_output_tokens: int = field(default=0) + estimator: BaseTokenEstimator | None = field(default=None) + context_window_resolver: BaseModelContextWindowResolver | None = field(default=None) + + +@dataclass(frozen=True) +class MicroCompactConfig: + """Configure micro compact.""" + enabled: bool = field(default=True) + gap_seconds: float = field(default=3_600.0) + trigger_count: int = field(default=20) + keep_recent: int = field(default=5) + tool_names: tuple[str, ...] = field(default=DEFAULT_COMPACTABLE_TOOL_NAMES) + + +@dataclass(frozen=True) +class ToolResultBudgetConfig: + """Configure tool result budget.""" + enabled: bool = field(default=True) + max_chars: int = field(default=50_000) + per_message_max_chars: int = field(default=200_000) + preview_chars: int = field(default=2_000) + + +@dataclass(frozen=True) +class SessionMemoryExtractorConfig: + """Configure session memory.""" + enabled: bool = field(default=True) + initial_chars: int = field(default=40_000) + update_chars: int = field(default=20_000) + initial_tokens: int = field(default=10_000) + update_tokens: int = field(default=5_000) + tool_calls_between_updates: int = field(default=3) + prompt_max_chars: int = field(default=200_000) + request_overhead_tokens: int = field(default=2_048) + section_max_chars: int = field(default=8_000) + total_max_chars: int = field(default=54_000) + wait_timeout_seconds: float = field(default=15.0) + max_retries: int = field(default=1) + + +@dataclass(frozen=True) +class AdvancedAutoCompactSummarizerConfig: + """Configure advanced compact.""" + history_snip: HistorySnipConfig = field(default_factory=HistorySnipConfig) + token_context_tracker: TokenContextTrackerConfig = field(default_factory=TokenContextTrackerConfig) + session_memory: SessionMemoryExtractorConfig = field(default_factory=SessionMemoryExtractorConfig) + tool_result_budget: ToolResultBudgetConfig = field(default_factory=ToolResultBudgetConfig) + micro_compact: MicroCompactConfig = field(default_factory=MicroCompactConfig) + auto_compact: AutoCompactSummarizerConfig = field(default_factory=AutoCompactSummarizerConfig) + + def __post_init__(self) -> None: + """Validate compression limits and token thresholds.""" + _require_positive( + tool_result_budget_max_chars=self.tool_result_budget.max_chars, + tool_result_budget_per_message_max_chars=self.tool_result_budget.per_message_max_chars, + tool_result_budget_preview_chars=self.tool_result_budget.preview_chars, + history_snip_trigger_chars=self.history_snip.trigger_chars, + history_snip_target_chars=self.history_snip.target_chars, + history_snip_keep_recent=self.history_snip.keep_recent, + session_memory_initial_chars=self.session_memory.initial_chars, + session_memory_update_chars=self.session_memory.update_chars, + session_memory_initial_tokens=self.session_memory.initial_tokens, + session_memory_update_tokens=self.session_memory.update_tokens, + session_memory_tool_calls_between_updates=self.session_memory.tool_calls_between_updates, + session_memory_prompt_max_chars=self.session_memory.prompt_max_chars, + session_memory_section_max_chars=self.session_memory.section_max_chars, + session_memory_total_max_chars=self.session_memory.total_max_chars, + session_memory_wait_timeout_seconds=self.session_memory.wait_timeout_seconds, + auto_compact_trigger_chars=self.auto_compact.trigger_chars, + auto_compact_target_chars=self.auto_compact.target_chars, + auto_compact_blocking_chars=self.auto_compact.blocking_chars, + auto_compact_keep_recent_contents=self.auto_compact.keep_recent_contents, + auto_compact_max_failures=self.auto_compact.max_failures, + auto_compact_summary_input_max_chars=self.auto_compact.summary_input_max_chars, + auto_compact_summary_retries_count=self.auto_compact.summary_retries_count, + micro_compact_gap_seconds=self.micro_compact.gap_seconds, + micro_compact_trigger_count=self.micro_compact.trigger_count, + micro_compact_keep_recent=self.micro_compact.keep_recent, + ) + _require_non_negative( + max_output_tokens=self.token_context_tracker.max_output_tokens, + session_memory_request_overhead_tokens=self.session_memory.request_overhead_tokens, + ) + context_window_tokens = self.token_context_tracker.model_context_window_tokens + if context_window_tokens is not None: + _require_positive(model_context_window_tokens=context_window_tokens) + if self.token_context_tracker.max_output_tokens >= context_window_tokens: + raise ValueError("max_output_tokens must be smaller than model_context_window_tokens") + if not (0 < self.token_context_tracker.warning_ratio < self.token_context_tracker.auto_compact_ratio < + self.token_context_tracker.blocking_ratio < 1): + raise ValueError("token context tracker ratios must satisfy 0 < warning < auto compact < blocking < 1") + if self.tool_result_budget.preview_chars >= self.tool_result_budget.max_chars: + raise ValueError("tool_result_budget.preview_chars must be smaller than tool_result_budget.max_chars") + if self.history_snip.target_chars >= self.history_snip.trigger_chars: + raise ValueError("history_snip.target_chars must be smaller than history_snip.trigger_chars") + if self.auto_compact.trigger_chars <= self.auto_compact.target_chars: + raise ValueError("auto_compact.trigger_chars must be greater than auto_compact.target_chars") + if self.auto_compact.blocking_chars <= self.auto_compact.trigger_chars: + raise ValueError("auto_compact.blocking_chars must be greater than auto_compact.trigger_chars") + _require_non_empty_names("history_snip.tool_names", self.history_snip.tool_names) + _require_non_empty_names("micro_compact.tool_names", self.micro_compact.tool_names) diff --git a/trpc_agent_sdk/sessions/compact/_coordination.py b/trpc_agent_sdk/sessions/compact/advanced/_coordination.py similarity index 100% rename from trpc_agent_sdk/sessions/compact/_coordination.py rename to trpc_agent_sdk/sessions/compact/advanced/_coordination.py diff --git a/trpc_agent_sdk/sessions/compact/advanced/_filters.py b/trpc_agent_sdk/sessions/compact/advanced/_filters.py new file mode 100644 index 000000000..fb7f1a0cc --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/_filters.py @@ -0,0 +1,80 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. + +from typing import Any +from typing_extensions import override + +from trpc_agent_sdk.abc import CompactTrigger +from trpc_agent_sdk.context import AgentContext +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.context import get_invocation_ctx +from trpc_agent_sdk.filter import BaseFilter +from trpc_agent_sdk.filter import FilterResult +from trpc_agent_sdk.filter import FilterType + +from ._manager import AdvancedAutoCompactSummarizerManager +from ._utils import INTERNAL_COMPACTION_METADATA_KEY + + +class AdvancedAutoCompactSummarizerFilter(BaseFilter): + """Advanced auto compact summarizer filter.""" + + def __init__(self) -> None: + """Initialize the advanced auto compact summarizer filter.""" + super().__init__() + self.name = "advanced_auto_compact_summarizer_filter" + self.type = FilterType.MODEL + + def get_summarizer_manager(self, ctx: InvocationContext) -> AdvancedAutoCompactSummarizerManager: + """Get the summarizer.""" + session_service = ctx.session_service + if session_service is None: + raise ValueError("Session service is not set") + summarizer_manager = getattr(session_service, "summarizer_manager", None) + if summarizer_manager is None or not isinstance(summarizer_manager, AdvancedAutoCompactSummarizerManager): + raise ValueError("Summarizer manager is not an AdvancedAutoCompactSummarizerManager") + return summarizer_manager + + @override + async def _before(self, ctx: AgentContext, req: Any, rsp: FilterResult): + """Run the advanced auto compact summarizer filter.""" + if ctx.get_metadata(INTERNAL_COMPACTION_METADATA_KEY, False): + return None + invocation_ctx: InvocationContext = get_invocation_ctx() + summarizer_manager = self.get_summarizer_manager(invocation_ctx) + result = await summarizer_manager.create_session_summary_before_model( + req, + invocation_ctx, + ) + if not result: + return None + invocation_ctx.end_invocation = True + rsp.rsp = result + rsp.is_continue = False + rsp.error = None + session_memory_extractor = summarizer_manager.get_session_memory_extractor() + if session_memory_extractor is not None: + await session_memory_extractor.extract_if_needed( + invocation_ctx, + force=False, + ) + return + + @override + async def _after(self, ctx: AgentContext, req: Any, rsp: FilterResult): + """Run the advanced auto compact summarizer filter.""" + if ctx.get_metadata(INTERNAL_COMPACTION_METADATA_KEY, False): + return None + invocation_ctx: InvocationContext = get_invocation_ctx() + summarizer_manager = self.get_summarizer_manager(invocation_ctx) + if summarizer_manager.compact_trigger != CompactTrigger.BEFORE_MODEL: + return None + session_memory_extractor = summarizer_manager.get_session_memory_extractor() + if session_memory_extractor is not None: + await session_memory_extractor.extract_if_needed( + invocation_ctx, + force=False, + ) diff --git a/trpc_agent_sdk/sessions/compact/_formats.py b/trpc_agent_sdk/sessions/compact/advanced/_formats.py similarity index 54% rename from trpc_agent_sdk/sessions/compact/_formats.py rename to trpc_agent_sdk/sessions/compact/advanced/_formats.py index 33c3841b6..05a19c965 100644 --- a/trpc_agent_sdk/sessions/compact/_formats.py +++ b/trpc_agent_sdk/sessions/compact/advanced/_formats.py @@ -9,125 +9,9 @@ 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 - -_FRONTMATTER_PATTERN = re.compile(r"\A---\n(?P.*?)\n---(?:\n|\Z)", re.DOTALL) -_UPDATED_AT_PATTERN = re.compile(r"^updated_at:\s*(?P\S+)\s*$", re.MULTILINE) - - -def _as_utc(value: datetime) -> datetime: - """Normalize an aware or naive datetime to UTC.""" - if value.tzinfo is None: - value = value.replace(tzinfo=timezone.utc) - return value.astimezone(timezone.utc) - - -class MemoryType(str, Enum): - """Semantic types allowed for long-term memory documents.""" - - USER = "user" - FEEDBACK = "feedback" - PROJECT = "project" - REFERENCE = "reference" - - -@dataclass(frozen=True) -class MemoryIndexEntry: - """Represent one standard entry in MEMORY.md.""" - - name: str - filename: str - summary: str - - def __post_init__(self) -> None: - """Validate that index fields are non-empty single-line strings.""" - for field_name, value in ( - ("name", self.name), - ("filename", self.filename), - ("summary", self.summary), - ): - if not value.strip() or "\n" in value or "\r" in value: - raise ValueError(f"{field_name} must be non-empty single-line text") - - def to_markdown(self) -> str: - """Render one long-term memory index entry.""" - return f"- [{self.name.strip()}]({self.filename.strip()}):{self.summary.strip()}" - - -@dataclass(frozen=True) -class MemoryDocument: - """Represent a long-term memory document with frontmatter.""" - - name: str - description: str - memory_type: MemoryType - content: str - updated_at: datetime | None = None - - def __post_init__(self) -> None: - """Validate frontmatter and reject unsafe multiline values.""" - for field_name, value in ( - ("name", self.name), - ("description", self.description), - ): - if not value.strip() or "\n" in value or "\r" in value: - raise ValueError(f"{field_name} must be non-empty single-line text") - - def to_markdown(self) -> str: - """Render standard frontmatter and document content.""" - body = self.content.strip() - updated_at = (_as_utc(self.updated_at).isoformat() if self.updated_at is not None else None) - updated_at_line = f"updated_at: {updated_at}\n" if updated_at else "" - return ("---\n" - f"name: {self.name.strip()}\n" - f"description: {self.description.strip()}\n" - f"type: {self.memory_type.value}\n" - f"{updated_at_line}" - "---\n" - f"{body}\n") - - -def parse_memory_updated_at(content: str) -> datetime | None: - """Extract the UTC update timestamp from a memory document.""" - frontmatter_match = _FRONTMATTER_PATTERN.match(content) - if frontmatter_match is None: - return None - match = _UPDATED_AT_PATTERN.search(frontmatter_match.group("frontmatter")) - if match is None: - return None - try: - value = match.group("value").replace("Z", "+00:00") - parsed = datetime.fromisoformat(value) - except ValueError: - return None - if parsed.tzinfo is None: - parsed = parsed.replace(tzinfo=timezone.utc) - return _as_utc(parsed) - - -def memory_freshness(updated_at: datetime | None, *, now: datetime | None = None) -> str: - """Return a compact freshness bucket suitable for model-facing output.""" - if updated_at is None: - return "unknown" - current = _as_utc(now or datetime.now(timezone.utc)) - timestamp = _as_utc(updated_at) - age_days = max(0, int((current - timestamp).total_seconds()) // 86_400) - if age_days == 0: - return "today" - if age_days == 1: - return "yesterday" - if age_days <= 7: - return "within 7 days" - if age_days <= 30: - return "within 30 days" - return "over 30 days" - SESSION_MEMORY_SECTIONS = ( "Session Title", diff --git a/trpc_agent_sdk/sessions/compact/_history_snip.py b/trpc_agent_sdk/sessions/compact/advanced/_history_snip.py similarity index 81% rename from trpc_agent_sdk/sessions/compact/_history_snip.py rename to trpc_agent_sdk/sessions/compact/advanced/_history_snip.py index 5f8867551..214f17bbb 100644 --- a/trpc_agent_sdk/sessions/compact/_history_snip.py +++ b/trpc_agent_sdk/sessions/compact/advanced/_history_snip.py @@ -12,21 +12,19 @@ import json from dataclasses import dataclass from typing import Any -from typing import TYPE_CHECKING +from typing_extensions import override -from ._callbacks import install_staged_callback -from ._runtime import SessionCompactRuntime +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.models import LlmRequest + +from ._base import BaseCompactSummarizerHandler +from ._runtime import AdvancedAutoCompactSummarizerRuntime from ._tool_result_budget import is_budget_replacement_response from ._tool_result_budget import serialize_tool_response from ._tool_result_budget import stable_tool_result_id from ._tool_result_budget import tool_result_sha256 from ._token_budget import TokenContextTracker -if TYPE_CHECKING: - from trpc_agent_sdk.agents import LlmAgent - from trpc_agent_sdk.context import InvocationContext - from trpc_agent_sdk.models import LlmRequest - HISTORY_SNIP_CLEARED_MESSAGE = "[Older tool result removed by history snip]" @@ -64,7 +62,7 @@ class HistorySnipResult: token_source: str | None = None -def estimate_request_chars(request: "LlmRequest") -> int: +def estimate_request_chars(request: LlmRequest) -> int: """Estimate the full model request using stable JSON serialization.""" payload = request.model_dump( mode="python", @@ -83,15 +81,16 @@ 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: SessionCompactRuntime) -> None: + def __init__(self, runtime: AdvancedAutoCompactSummarizerRuntime) -> None: """Initialize history-snip state and per-session async locks.""" - self._runtime = memory_runtime + self._runtime = runtime + self._config = runtime.config.history_snip self._states: dict[str, HistorySnipState] = {} self._session_locks: dict[str, asyncio.Lock] = {} self._scoped_processors: dict[object, "HistorySnip"] = {} @property - def runtime(self) -> SessionCompactRuntime: + def runtime(self) -> AdvancedAutoCompactSummarizerRuntime: """Return the runtime bound to this history snipper.""" return self._runtime @@ -119,7 +118,7 @@ async def _load_state(self, session_id: str) -> HistorySnipState: def _collect_candidates(self, request: "LlmRequest") -> list[HistorySnipCandidate]: """Collect eligible function results in request order.""" - allowed_tools = set(self._runtime.config.history_snip_tool_names) + allowed_tools = set(self._config.tool_names) candidates: list[HistorySnipCandidate] = [] serialized_by_result_id: dict[str, str] = {} for content in request.contents: @@ -152,18 +151,17 @@ def _snipped_response(self) -> dict[str, str]: async def apply( self, - request: "LlmRequest", + request: LlmRequest, *, - session_id: str, - ctx: "InvocationContext | None" = None, + ctx: InvocationContext, force: bool = False, ) -> HistorySnipResult: """Clean old tool results when over budget or explicitly forced.""" - config = self._runtime.config - if not config.enabled or not config.history_snip_enabled: + if not self._config.enabled: request_chars = estimate_request_chars(request) return HistorySnipResult(None, 0, 0, 0, request_chars, request_chars) - if ctx is None or hasattr(self._runtime, "scope"): + session_id = ctx.session_id + if self._runtime.scope: return await self._apply_scoped(request, session_id=session_id, ctx=ctx, force=force) runtime = self._runtime.for_session(ctx.session) processor = self._scoped_processors.get(runtime.scope) @@ -173,19 +171,20 @@ async def apply( processor._states = {} processor._session_locks = {} self._scoped_processors[runtime.scope] = processor - return await processor.apply(request, session_id=session_id, ctx=ctx, force=force) + return await processor.apply(request, ctx=ctx, force=force) async def _apply_scoped( self, - request: "LlmRequest", + request: LlmRequest, *, session_id: str, - ctx: "InvocationContext", + ctx: InvocationContext, force: bool, ) -> HistorySnipResult: """Apply one tenant-bound history-snipping operation.""" - config = self._runtime.config - tracker = TokenContextTracker(config) + token_context_tracker_config = self._runtime.config.token_context_tracker + history_snip_config = self._config + tracker = TokenContextTracker(token_context_tracker_config) async with self._session_lock(session_id): request.contents = [content.model_copy(deep=True) for content in request.contents] state = await self._load_state(session_id) @@ -206,7 +205,7 @@ async def _apply_scoped( token_mode = token_budget_before.token_mode_enabled current_tokens = token_budget_before.estimate.tokens if not force and (current_tokens <= token_budget_before.warning_threshold_tokens - if token_mode else request_chars_before <= config.history_snip_trigger_chars): + if token_mode else request_chars_before <= history_snip_config.trigger_chars): return HistorySnipResult( None, 0, @@ -220,7 +219,7 @@ async def _apply_scoped( ) trigger = "force" if force else "pressure" - protected_ids = {candidate.result_id for candidate in candidates[-config.history_snip_keep_recent:]} + protected_ids = {candidate.result_id for candidate in candidates[-history_snip_config.keep_recent:]} eligible = [ candidate for candidate in candidates if candidate.result_id not in state.snipped_ids and candidate.result_id not in protected_ids @@ -230,7 +229,7 @@ async def _apply_scoped( snipped_count = 0 for candidate in eligible: if not force and (current_tokens <= token_budget_before.warning_threshold_tokens - if token_mode else current_chars <= config.history_snip_target_chars): + if token_mode else current_chars <= history_snip_config.target_chars): break candidate_saving = max(0, candidate.original_size - replacement_size) if candidate_saving == 0: @@ -258,39 +257,20 @@ async def _apply_scoped( ) -class HistorySnipCallback: +class HistorySnipHandler(BaseCompactSummarizerHandler): """Adapt history snip to before_model_callback.""" - advanced_memory_stage = 20 - - def __init__(self, history_snip: HistorySnip) -> None: + def __init__(self) -> None: """Store the history-snip processor run before model requests.""" - self._history_snip = history_snip + self._history_snip: HistorySnip | None = None - @property - def history_snip(self) -> HistorySnip: - """Return the history-snip processor used by this callback.""" - return self._history_snip - - async def __call__(self, ctx: "InvocationContext", request: "LlmRequest") -> None: + @override + async def handle(self, ctx: InvocationContext, request: LlmRequest) -> None: """Run history snip before a request based on request size.""" - await self._history_snip.apply(request, session_id=ctx.session_id, ctx=ctx) - return None - - -def setup_history_snip( - agent: "LlmAgent", - memory_runtime: SessionCompactRuntime, -) -> HistorySnip: - """Install history snip while preserving context stage order.""" - history_snip = HistorySnip(memory_runtime) - callback = HistorySnipCallback(history_snip) - existing_snip = install_staged_callback( - agent, - callback, - callback_type=HistorySnipCallback, - component_attribute="history_snip", - memory_runtime=memory_runtime, - conflict_message="History snip is already configured with another runtime", - ) - return existing_snip or history_snip + summarizer = self.get_summarizer(ctx) + from ._auto_compact import AdvancedAutoCompactSummarizer + if not isinstance(summarizer, AdvancedAutoCompactSummarizer): + raise ValueError("Summarizer is not an AdvancedAutoCompactSummarizer") + if self._history_snip is None: + self._history_snip = HistorySnip(summarizer.runtime) + await self._history_snip.apply(request, ctx=ctx) diff --git a/trpc_agent_sdk/sessions/compact/advanced/_manager.py b/trpc_agent_sdk/sessions/compact/advanced/_manager.py new file mode 100644 index 000000000..8759a5e95 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/_manager.py @@ -0,0 +1,101 @@ +# Tencent is pleased to support the open source community by making +# contributions to the open source ecosystem. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Integrate Session Compact with the native SessionService lifecycle.""" + +from __future__ import annotations + +from typing_extensions import override + +from trpc_agent_sdk.abc import CompactSummarizerManagerABC +from trpc_agent_sdk.abc import CompactTrigger +from trpc_agent_sdk.abc import RequestABC +from trpc_agent_sdk.abc import ResponseABC +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.models import LlmRequest + +from ..._session import Session +from ._auto_compact import AdvancedAutoCompactSummarizer +from ._formats import parse_session_memory_state +from ._formats import SESSION_MEMORY_STATE_KEY +from ._compaction_memory_extractor import SessionMemoryExtractor +from ._history_snip import HistorySnipHandler +from ._micro_compact import MicroCompactHandler +from ._tool_result_budget import ToolResultBudgetHandler +from ._auto_compact import AdvancedAutoCompactSummarizerHandler + + +class AdvancedAutoCompactSummarizerManager(CompactSummarizerManagerABC): + """Coordinate Advanced Compact state without wrapping a SessionService.""" + + def __init__( + self, + summarizer: AdvancedAutoCompactSummarizer, + compact_trigger: CompactTrigger = CompactTrigger.BEFORE_MODEL, + ) -> None: + """Store configuration until Runner supplies the Agent.""" + super().__init__(summarizer, compact_trigger=compact_trigger) + self._tool_result_budget_handler = ToolResultBudgetHandler() + self._history_snip_handler = HistorySnipHandler() + self._micro_compact_handler = MicroCompactHandler() + self._advanced_auto_compact_summarizer_handler = AdvancedAutoCompactSummarizerHandler() + + def get_session_memory_extractor(self) -> SessionMemoryExtractor: + """Get the session memory extractor.""" + return self.summarizer.session_memory_extractor + + @override + async def create_session_summary( + self, + session: Session, + force: bool = False, + ctx: InvocationContext | None = None, + ) -> None: + """Compact persisted Events when configured for end-of-turn execution.""" + if self.compact_trigger != CompactTrigger.AFTER_TURN: + return + if ctx is None: + raise ValueError("Invocation context is required for advanced compaction") + session_memory_extractor = self.get_session_memory_extractor() + if session_memory_extractor is not None: + await session_memory_extractor.extract_if_needed( + ctx, + force=False, + ) + if force or await self.summarizer.should_summarize(session): + await self.summarizer.create_session_summary( + session, + ctx=ctx, + store_historical_events=True, + ) + + @override + async def create_session_summary_before_model( + self, + request: RequestABC, + ctx: InvocationContext, + force: bool = False, + ) -> ResponseABC | None: + """Run the advanced request pipeline immediately before model generation.""" + if self.compact_trigger != CompactTrigger.BEFORE_MODEL: + return None + if not isinstance(request, LlmRequest): + raise TypeError("Advanced compaction requires an LlmRequest") + await self._tool_result_budget_handler.handle(ctx, request) + await self._history_snip_handler.handle(ctx, request) + await self._micro_compact_handler.handle(ctx, request) + return await self._advanced_auto_compact_summarizer_handler.handle( + ctx, + request, + force=force, + ) + + async def get_session_summary(self, session: Session) -> str | None: + """Read compact Session Memory through the existing summary API.""" + parsed = parse_session_memory_state(session.state.get(SESSION_MEMORY_STATE_KEY)) + if parsed is not None: + return parsed[0].to_markdown() + return None diff --git a/trpc_agent_sdk/sessions/compact/_microcompact.py b/trpc_agent_sdk/sessions/compact/advanced/_micro_compact.py similarity index 67% rename from trpc_agent_sdk/sessions/compact/_microcompact.py rename to trpc_agent_sdk/sessions/compact/advanced/_micro_compact.py index 896c45212..da26d9295 100644 --- a/trpc_agent_sdk/sessions/compact/_microcompact.py +++ b/trpc_agent_sdk/sessions/compact/advanced/_micro_compact.py @@ -11,26 +11,26 @@ import copy import time from dataclasses import dataclass -from typing import Any -from typing import TYPE_CHECKING +from dataclasses import field +from typing import Optional +from typing_extensions import override -from ._callbacks import install_staged_callback -from ._runtime import SessionCompactRuntime +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.types import Part + +from ._base import BaseCompactSummarizerHandler +from ._runtime import AdvancedAutoCompactSummarizerRuntime from ._tool_result_budget import is_budget_replacement_response from ._tool_result_budget import serialize_tool_response from ._tool_result_budget import stable_tool_result_id from ._tool_result_budget import tool_result_sha256 -if TYPE_CHECKING: - from trpc_agent_sdk.agents import LlmAgent - from trpc_agent_sdk.context import InvocationContext - from trpc_agent_sdk.models import LlmRequest - MICROCOMPACT_CLEARED_MESSAGE = "[Old tool result content cleared]" @dataclass -class MicrocompactState: +class MicroCompactState: """Store identifiers for mechanically cleaned tool results.""" cleared_ids: set[str] @@ -38,24 +38,24 @@ class MicrocompactState: @dataclass(frozen=True) -class MicrocompactCandidate: +class MicroCompactCandidate: """Describe a function response eligible for mechanical cleanup.""" result_id: str tool_name: str original_size: int original_sha256: str - part: Any + part: Part @dataclass(frozen=True) -class MicrocompactResult: +class MicroCompactResult: """Summarize new and repeated mechanical cleanup operations.""" - trigger: str | None - cleared_count: int - reapplied_count: int - chars_saved: int + trigger: Optional[str] = field(default=None) + cleared_count: int = field(default=0) + reapplied_count: int = field(default=0) + chars_saved: int = field(default=0) def find_last_assistant_timestamp(ctx: "InvocationContext") -> float | None: @@ -71,20 +71,16 @@ def find_last_assistant_timestamp(ctx: "InvocationContext") -> float | None: return None -class Microcompact: +class MicroCompact: """Local compressor that cleans old tool results by time or count.""" - def __init__(self, memory_runtime: SessionCompactRuntime) -> None: + def __init__(self, runtime: AdvancedAutoCompactSummarizerRuntime) -> None: """Initialize mechanical-compaction state and per-session locks.""" - self._runtime = memory_runtime - self._states: dict[str, MicrocompactState] = {} + self._runtime = runtime + self._config = runtime.config.micro_compact + self._states: dict[str, MicroCompactState] = {} self._session_locks: dict[str, asyncio.Lock] = {} - self._scoped_processors: dict[object, "Microcompact"] = {} - - @property - def runtime(self) -> SessionCompactRuntime: - """Return the runtime bound to this mechanical compressor.""" - return self._runtime + self._scoped_processors: dict[object, MicroCompact] = {} def _session_lock(self, session_id: str) -> asyncio.Lock: """Return the unique async compaction lock for a session.""" @@ -95,23 +91,23 @@ def _session_lock(self, session_id: str) -> asyncio.Lock: self._session_locks[key] = lock return lock - async def _load_state(self, session_id: str) -> MicrocompactState: - """Return process-local microcompact state.""" + async def _load_state(self, session_id: str) -> MicroCompactState: + """Return process-local micro-compact state.""" state_key = self._runtime.session_key(session_id) if hasattr(self._runtime, "session_key") else session_id state = self._states.get(state_key) if state is not None: return state - state = MicrocompactState( + state = MicroCompactState( cleared_ids=set(), result_hashes={}, ) self._states[state_key] = state return state - def _collect_candidates(self, request: "LlmRequest") -> list[MicrocompactCandidate]: + def _collect_candidates(self, request: LlmRequest) -> list[MicroCompactCandidate]: """Collect eligible tool results in request order.""" - allowed_tools = set(self._runtime.config.microcompact_tool_names) - candidates: list[MicrocompactCandidate] = [] + allowed_tools = set(self._config.tool_names) + candidates: list[MicroCompactCandidate] = [] serialized_by_result_id: dict[str, str] = {} for content in request.contents: for part in content.parts or []: @@ -127,7 +123,7 @@ def _collect_candidates(self, request: "LlmRequest") -> list[MicrocompactCandida raise ValueError(f"Tool result id {result_id!r} is reused with different content") serialized_by_result_id[result_id] = serialized candidates.append( - MicrocompactCandidate( + MicroCompactCandidate( result_id=result_id, tool_name=function_response.name, original_size=len(serialized), @@ -142,21 +138,19 @@ def _cleared_response(self) -> dict[str, str]: async def apply( self, - request: "LlmRequest", + request: LlmRequest, *, - session_id: str, + ctx: InvocationContext, last_assistant_timestamp: float | None, - ctx: "InvocationContext | None" = None, now: float | None = None, - ) -> MicrocompactResult: + ) -> 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) - if ctx is None or hasattr(self._runtime, "scope"): + if not self._config.enabled: + return MicroCompactResult() + if self._runtime.scope: return await self._apply_scoped( request, - session_id=session_id, + session_id=ctx.session_id, last_assistant_timestamp=last_assistant_timestamp, now=now, ) @@ -170,7 +164,6 @@ async def apply( 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, @@ -178,14 +171,13 @@ async def apply( async def _apply_scoped( self, - request: "LlmRequest", + request: LlmRequest, *, session_id: str, last_assistant_timestamp: float | None, now: float | None, - ) -> MicrocompactResult: + ) -> 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) @@ -196,7 +188,7 @@ async def _apply_scoped( raise ValueError(f"Tool result id {candidate.result_id!r} is reused with different content") reapplied_count = 0 - active_candidates: list[MicrocompactCandidate] = [] + active_candidates: list[MicroCompactCandidate] = [] for candidate in candidates: if candidate.result_id in state.cleared_ids: candidate.part.function_response.response = self._cleared_response() @@ -206,19 +198,19 @@ async def _apply_scoped( current_time = time.time() if now is None else now gap_seconds = current_time - last_assistant_timestamp if last_assistant_timestamp is not None else None - if gap_seconds is not None and gap_seconds >= config.microcompact_gap_seconds: + if gap_seconds is not None and gap_seconds >= self._config.gap_seconds: trigger = "time" - elif len(active_candidates) > config.microcompact_trigger_count: + elif len(active_candidates) > self._config.trigger_count: trigger = "count" else: trigger = None if trigger is None: - return MicrocompactResult(None, 0, reapplied_count, 0) + return MicroCompactResult(reapplied_count=reapplied_count) - clear_candidates = active_candidates[:-config.microcompact_keep_recent] + clear_candidates = active_candidates[:-self._config.keep_recent] if not clear_candidates: - return MicrocompactResult(None, 0, reapplied_count, 0) + return MicroCompactResult(reapplied_count=reapplied_count) cleared_size = len(serialize_tool_response(self._cleared_response())) chars_saved = 0 @@ -228,7 +220,7 @@ async def _apply_scoped( state.result_hashes[candidate.result_id] = candidate.original_sha256 chars_saved += max(0, candidate.original_size - cleared_size) - return MicrocompactResult( + return MicroCompactResult( trigger=trigger, cleared_count=len(clear_candidates), reapplied_count=reapplied_count, @@ -236,44 +228,20 @@ async def _apply_scoped( ) -class MicrocompactCallback: +class MicroCompactHandler(BaseCompactSummarizerHandler): """Adapt the mechanical compressor to before_model_callback.""" - advanced_memory_stage = 30 - - def __init__(self, microcompact: Microcompact) -> None: + def __init__(self) -> None: """Store the compressor executed before model requests.""" - self._microcompact = microcompact + self._micro_compact: MicroCompact | None = None - @property - def microcompact(self) -> Microcompact: - """Return the compressor used by this callback.""" - return self._microcompact - - async def __call__(self, ctx: "InvocationContext", request: "LlmRequest") -> None: + @override + async def handle(self, ctx: InvocationContext, request: LlmRequest) -> None: """Calculate the time gap and run mechanical cleanup before a request.""" - await self._microcompact.apply( - request, - session_id=ctx.session_id, - last_assistant_timestamp=find_last_assistant_timestamp(ctx), - ctx=ctx, - ) - return None - - -def setup_microcompact( - agent: "LlmAgent", - memory_runtime: SessionCompactRuntime, -) -> Microcompact: - """Install the mechanical callback while preserving existing order.""" - microcompact = Microcompact(memory_runtime) - callback = MicrocompactCallback(microcompact) - existing_microcompact = install_staged_callback( - agent, - callback, - callback_type=MicrocompactCallback, - component_attribute="microcompact", - memory_runtime=memory_runtime, - conflict_message="Microcompact is already configured with another runtime", - ) - return existing_microcompact or microcompact + summarizer = self.get_summarizer(ctx) + from ._auto_compact import AdvancedAutoCompactSummarizer + if not isinstance(summarizer, AdvancedAutoCompactSummarizer): + raise ValueError("Summarizer is not an AdvancedAutoCompactSummarizer") + if self._micro_compact is None: + self._micro_compact = MicroCompact(summarizer.runtime) + await self._micro_compact.apply(request, ctx=ctx, last_assistant_timestamp=find_last_assistant_timestamp(ctx)) diff --git a/trpc_agent_sdk/sessions/compact/advanced/_runtime.py b/trpc_agent_sdk/sessions/compact/advanced/_runtime.py new file mode 100644 index 000000000..1369e88e4 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/_runtime.py @@ -0,0 +1,30 @@ +# Tencent is pleased to support the open source ecosystem. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# Licensed under Apache-2.0. +"""Runtime coordination for Session Compact.""" + +from __future__ import annotations + +from dataclasses import dataclass, field + +from ..._session import Session +from ._config import AdvancedAutoCompactSummarizerConfig +from ._coordination import SessionOperationCoordinator + + +@dataclass +class AdvancedAutoCompactSummarizerRuntime: + """Hold advanced auto compact summarizer configuration and per-session coordination only.""" + + config: AdvancedAutoCompactSummarizerConfig = field(default_factory=AdvancedAutoCompactSummarizerConfig) + coordination: SessionOperationCoordinator = field(default_factory=SessionOperationCoordinator) + scope: str = field(default="") + + def for_session(self, session: Session) -> AdvancedAutoCompactSummarizerRuntime: + return AdvancedAutoCompactSummarizerRuntime(config=self.config, + coordination=self.coordination, + scope=f"{session.app_name}\0{session.user_id}") + + def session_key(self, session_id: str) -> str: + return f"{self.scope}\0{session_id}" diff --git a/trpc_agent_sdk/sessions/compact/_token_budget.py b/trpc_agent_sdk/sessions/compact/advanced/_token_budget.py similarity index 68% rename from trpc_agent_sdk/sessions/compact/_token_budget.py rename to trpc_agent_sdk/sessions/compact/advanced/_token_budget.py index aad9af666..22a77cd36 100644 --- a/trpc_agent_sdk/sessions/compact/_token_budget.py +++ b/trpc_agent_sdk/sessions/compact/advanced/_token_budget.py @@ -11,13 +11,17 @@ import math import re from dataclasses import dataclass +from dataclasses import field from typing import Any -from typing import Protocol -from typing import TYPE_CHECKING +from typing import Optional +from typing_extensions import override -if TYPE_CHECKING: - from trpc_agent_sdk.context import InvocationContext - from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.models import LlmRequest +from trpc_agent_sdk.types import Content + +from ._config import TokenContextTrackerConfig +from ._base import BaseTokenEstimator _CJK_CHARACTER = re.compile(r"[\u3400-\u4dbf\u4e00-\u9fff\uf900-\ufaff]") @@ -28,7 +32,7 @@ class ContextTokenEstimate: tokens: int source: str - usage_event_id: str | None = None + usage_event_id: Optional[str] = field(default=None) @dataclass(frozen=True) @@ -36,11 +40,11 @@ class ContextBudget: """Describe the model window and the request's position within it.""" estimate: ContextTokenEstimate - context_window_tokens: int | None - effective_window_tokens: int | None - warning_threshold_tokens: int | None - autocompact_threshold_tokens: int | None - blocking_threshold_tokens: int | None + context_window_tokens: Optional[int] = field(default=None) + effective_window_tokens: Optional[int] = field(default=None) + warning_threshold_tokens: Optional[int] = field(default=None) + auto_compact_threshold_tokens: Optional[int] = field(default=None) + blocking_threshold_tokens: Optional[int] = field(default=None) @property def token_mode_enabled(self) -> bool: @@ -48,23 +52,10 @@ def token_mode_enabled(self) -> bool: return self.effective_window_tokens is not None -class TokenEstimator(Protocol): - """Define the replaceable token estimator interface.""" - - def estimate_payload_tokens(self, payload: Any) -> int: - """Estimate tokens for any JSON-compatible payload.""" - - -class ModelContextWindowResolver(Protocol): - """Define the model-identifier context-window resolver interface.""" - - def resolve_context_window_tokens(self, model: Any) -> int | None: - """Return the model context window, or None when unknown.""" - - -class HeuristicTokenEstimator: +class HeuristicTokenEstimator(BaseTokenEstimator): """Estimate JSON request tokens with a mixed-language heuristic.""" + @override def estimate_payload_tokens(self, payload: Any) -> int: """Estimate CJK at one token per character and other text at four characters per token.""" rendered = json.dumps( @@ -111,17 +102,17 @@ def _usage_context_tokens(usage: Any) -> int | None: return sum(value for value in values if isinstance(value, int)) -def _request_static_fingerprint(request: "LlmRequest") -> str: +def _request_static_fingerprint(request: LlmRequest) -> str: """Extract fingerprints for model, instructions, and tool configuration.""" - config = getattr(request, "config", None) - if hasattr(config, "model_dump"): + config = request.config + if config is not None: config_payload = config.model_dump( mode="python", by_alias=True, exclude_none=True, ) else: - config_payload = config + config_payload = {} return json.dumps( { "model": request.model, @@ -134,45 +125,46 @@ def _request_static_fingerprint(request: "LlmRequest") -> str: ) -class TokenContextTracker: +class TokenContextTracker(BaseTokenEstimator): """Estimate request context tokens from usage and new content.""" - def __init__(self, config: Any) -> None: + def __init__(self, config: TokenContextTrackerConfig): """Store configuration and choose the default or injected estimator.""" self._config = config - estimator = getattr(config, "token_estimator", None) - self._estimator: TokenEstimator = estimator or HeuristicTokenEstimator() + self._estimator = config.estimator or HeuristicTokenEstimator() - def _resolve_window_tokens(self, ctx: "InvocationContext | None") -> int | None: + def _resolve_window_tokens(self, ctx: InvocationContext) -> Optional[int]: """Resolve the model context window from config or an application resolver.""" - explicit = getattr(self._config, "model_context_window_tokens", None) - if isinstance(explicit, int) and explicit > 0: + if not self._config.enabled: + return None + explicit = self._config.model_context_window_tokens + if explicit is not None and explicit > 0: return explicit - resolver = getattr(self._config, "context_window_resolver", None) + resolver = self._config.context_window_resolver if resolver is None or ctx is None: return None - model = getattr(getattr(ctx, "agent", None), "model", None) + model = ctx.agent.model if ctx.agent else None resolved = resolver.resolve_context_window_tokens(model) - return resolved if isinstance(resolved, int) and resolved > 0 else None + return resolved if resolved is not None and resolved > 0 else None - def _estimate_request(self, request: "LlmRequest") -> int: + def _estimate_request(self, request: LlmRequest) -> int: """Estimate tokens for a complete LlmRequest.""" payload = request.model_dump(mode="python", by_alias=True, exclude_none=True) return self._estimator.estimate_payload_tokens(payload) - def _estimate_new_contents(self, contents: list[Any]) -> int: + def _estimate_new_contents(self, contents: list[Content]) -> int: """Estimate content tokens added after a usage baseline.""" payload = [content.model_dump(mode="python", by_alias=True, exclude_none=True) for content in contents] return self._estimator.estimate_payload_tokens(payload) if payload else 0 def _latest_usage_baseline( self, - request: "LlmRequest", - ctx: "InvocationContext | None", + request: LlmRequest, + ctx: InvocationContext, ) -> ContextTokenEstimate | None: """Match the latest usage event and estimate subsequent context.""" - session = getattr(ctx, "session", None) if ctx is not None else None - events = getattr(session, "events", None) + session = ctx.session + events = session.events if not isinstance(events, list): return None fingerprints = [_content_fingerprint(content) for content in request.contents] @@ -200,10 +192,10 @@ def _latest_usage_baseline( ) return None - def estimate( + def _estimate( self, - request: "LlmRequest", - ctx: "InvocationContext | None" = None, + request: LlmRequest, + ctx: InvocationContext, ) -> ContextTokenEstimate: """Prefer recent model usage, falling back to a full request estimate.""" baseline = self._latest_usage_baseline(request, ctx) @@ -214,21 +206,30 @@ def estimate( source="estimated", ) + def estimate( + self, + request: LlmRequest, + ctx: InvocationContext, + ) -> ContextTokenEstimate: + """Return the best available token estimate for a request.""" + return self._estimate(request, ctx) + + @override def estimate_payload_tokens(self, payload: Any) -> int: """Reuse the same estimator for non-request inputs such as session memory.""" return self._estimator.estimate_payload_tokens(payload) - def estimate_request_tokens(self, request: "LlmRequest") -> int: + 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: + def token_mode_enabled(self, ctx: InvocationContext) -> bool: """Return whether the configuration resolves a model context window.""" return self._resolve_window_tokens(ctx) is not None def effective_context_window_tokens( self, - ctx: "InvocationContext | None" = None, + ctx: InvocationContext, ) -> int | None: """Return the input window after max output, or None when unknown.""" context_window = self._resolve_window_tokens(ctx) @@ -237,39 +238,40 @@ def effective_context_window_tokens( effective = context_window - getattr(self._config, "max_output_tokens", 0) return effective if effective > 0 else None + @classmethod def record_request_context( - self, - request: "LlmRequest", - ctx: "InvocationContext | None", + cls, + request: LlmRequest, + ctx: InvocationContext, ) -> None: """Stage the final request fingerprint for persistence on the response Event.""" - session = getattr(ctx, "session", None) if ctx is not None else None - state = getattr(session, "state", None) + session = ctx.session + state = session.state if isinstance(state, dict): state["advanced_memory_pending_request_context_fingerprint"] = _request_static_fingerprint(request) def budget( self, - request: "LlmRequest", - ctx: "InvocationContext | None" = None, + request: LlmRequest, + ctx: InvocationContext, ) -> ContextBudget: """Calculate the effective window, thresholds, and token estimate.""" - estimate = self.estimate(request, ctx) + estimate = self._estimate(request, ctx) context_window = self._resolve_window_tokens(ctx) if context_window is None: - return ContextBudget(estimate, None, None, None, None, None) - max_output_tokens = getattr(self._config, "max_output_tokens", 0) + return ContextBudget(estimate=estimate) + max_output_tokens = self._config.max_output_tokens effective = context_window - max_output_tokens if effective <= 0: - return ContextBudget(estimate, context_window, None, None, None, None) - warning = math.floor(effective * getattr(self._config, "token_warning_ratio", 0.85)) - autocompact = math.floor(effective * getattr(self._config, "token_autocompact_ratio", 0.90)) - blocking = math.floor(effective * getattr(self._config, "token_blocking_ratio", 0.95)) + return ContextBudget(estimate=estimate, context_window_tokens=context_window) + warning = math.floor(effective * self._config.warning_ratio) + auto_compact = math.floor(effective * self._config.auto_compact_ratio) + blocking = math.floor(effective * self._config.blocking_ratio) return ContextBudget( - estimate, - context_window, - effective, - warning, - autocompact, - blocking, + estimate=estimate, + context_window_tokens=context_window, + effective_window_tokens=effective, + warning_threshold_tokens=warning, + auto_compact_threshold_tokens=auto_compact, + blocking_threshold_tokens=blocking, ) diff --git a/trpc_agent_sdk/sessions/compact/_tool_result_budget.py b/trpc_agent_sdk/sessions/compact/advanced/_tool_result_budget.py similarity index 77% rename from trpc_agent_sdk/sessions/compact/_tool_result_budget.py rename to trpc_agent_sdk/sessions/compact/advanced/_tool_result_budget.py index 3ad6cda17..0f9fd51ec 100644 --- a/trpc_agent_sdk/sessions/compact/_tool_result_budget.py +++ b/trpc_agent_sdk/sessions/compact/advanced/_tool_result_budget.py @@ -11,17 +11,16 @@ import copy import hashlib import json -from dataclasses import dataclass -from typing import Any -from typing import TYPE_CHECKING +from dataclasses import dataclass, field +from typing import Any, Optional +from typing_extensions import override -from ._callbacks import install_staged_callback -from ._runtime import SessionCompactRuntime +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.models import LlmRequest -if TYPE_CHECKING: - from trpc_agent_sdk.agents import LlmAgent - from trpc_agent_sdk.context import InvocationContext - from trpc_agent_sdk.models import LlmRequest +from ..._session import Session +from ._runtime import AdvancedAutoCompactSummarizerRuntime +from ._base import BaseCompactSummarizerHandler TOOL_RESULT_REPLACEMENT_SCHEMA_VERSION = 1 @@ -40,11 +39,11 @@ 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 part: Any + event_id: Optional[str] = field(default=None) @dataclass(frozen=True) @@ -60,9 +59,9 @@ class ToolResultReplacement: class ToolResultBudgetResult: """Summarize replacements and character savings from budget processing.""" - replaced_count: int - original_chars: int - replacement_chars: int + replaced_count: int = field(default=0) + original_chars: int = field(default=0) + replacement_chars: int = field(default=0) def serialize_tool_response(response: Any) -> str: @@ -112,21 +111,17 @@ 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: SessionCompactRuntime) -> None: + def __init__(self, runtime: AdvancedAutoCompactSummarizerRuntime) -> None: """Initialize the budget processor and per-session state locks.""" - self._runtime = memory_runtime + self._runtime: AdvancedAutoCompactSummarizerRuntime = runtime + self._tool_result_budget_config = runtime.config.tool_result_budget self._states: dict[str, ToolResultBudgetState] = {} self._session_locks: dict[str, asyncio.Lock] = {} - self._scoped_processors: dict[object, "ToolResultBudget"] = {} - - @property - def runtime(self) -> SessionCompactRuntime: - """Return the runtime bound to this budget processor.""" - return self._runtime + self._scoped_processors: dict[object, ToolResultBudget] = {} def _session_lock(self, session_id: str) -> asyncio.Lock: """Return the unique async budget lock for a session.""" - key = self._runtime.session_key(session_id) if hasattr(self._runtime, "session_key") else session_id + key = self._runtime.session_key(session_id) lock = self._session_locks.get(key) if lock is None: lock = asyncio.Lock() @@ -135,7 +130,7 @@ def _session_lock(self, session_id: str) -> asyncio.Lock: async def _load_state(self, session_id: str) -> ToolResultBudgetState: """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_key = self._runtime.session_key(session_id) state = self._states.get(state_key) if state is not None: return state @@ -149,20 +144,22 @@ async def _load_state(self, session_id: str) -> ToolResultBudgetState: def _collect_candidates( self, - request: "LlmRequest", - session: Any | None = None, + request: LlmRequest, + session: Session, ) -> list[list[ToolResultCandidate]]: """Group function responses from consecutive user contents.""" event_ids: dict[str, str] = {} - for event in getattr(session, "events", []) or []: + for event in session.events: event_id = getattr(event, "id", None) - if not isinstance(event_id, str): + if event_id is None: + continue + event_content = event.content + if event_content is None: continue - event_content = getattr(event, "content", None) - for event_part in getattr(event_content, "parts", []) or []: + for event_part in event_content.parts: response = getattr(event_part, "function_response", None) - response_id = getattr(response, "id", None) - if isinstance(response_id, str): + response_id = response.id if response is not None else None + if response_id is not None: event_ids[response_id] = event_id candidate_groups: list[list[ToolResultCandidate]] = [] serialized_by_result_id: dict[str, str] = {} @@ -204,7 +201,7 @@ def _build_replacement( """Build an event reference and model-visible preview.""" preview, truncated = _preview_text( candidate.serialized_result, - self._runtime.config.tool_result_preview_chars, + self._tool_result_budget_config.preview_chars, ) replacement_response = { "_advanced_memory": { @@ -231,7 +228,6 @@ def _select_replacements( ) -> list[ToolResultReplacement]: """Apply per-result limits, then select results under the aggregate limit.""" selected: dict[str, ToolResultReplacement] = {} - config = self._runtime.config for group in groups: fresh = [ candidate for candidate in group @@ -239,7 +235,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: + if candidate.original_size > self._tool_result_budget_config.max_chars: selected[candidate.result_id] = self._build_replacement(candidate) visible_size = 0 @@ -257,7 +253,7 @@ def _select_replacements( remaining_fresh.append(candidate) for candidate in sorted(remaining_fresh, key=lambda item: item.original_size, reverse=True): - if visible_size <= config.tool_results_per_message_max_chars: + if visible_size <= self._tool_result_budget_config.per_message_max_chars: break replacement = self._build_replacement(candidate) if replacement.replacement_size >= candidate.original_size: @@ -266,18 +262,12 @@ def _select_replacements( visible_size -= candidate.original_size - replacement.replacement_size return list(selected.values()) - async def apply( - self, - request: "LlmRequest", - *, - session_id: str, - ctx: "InvocationContext | None" = None, - ) -> ToolResultBudgetResult: + async def apply(self, request: LlmRequest, ctx: InvocationContext) -> ToolResultBudgetResult: """Process a model request without mutating session Events.""" - if not self._runtime.config.enabled: - return ToolResultBudgetResult(0, 0, 0) - if ctx is None or hasattr(self._runtime, "scope"): - return await self._apply_scoped(request, session_id, getattr(ctx, "session", None)) + if not self._tool_result_budget_config.enabled: + return ToolResultBudgetResult() + if self._runtime.scope: + return await self._apply_scoped(request, ctx.session) runtime = self._runtime.for_session(ctx.session) processor = self._scoped_processors.get(runtime.scope) if processor is None: @@ -286,18 +276,13 @@ async def apply( processor._states = {} processor._session_locks = {} self._scoped_processors[runtime.scope] = processor - return await processor.apply(request, session_id=session_id, ctx=ctx) + return await processor.apply(request, ctx=ctx) - async def _apply_scoped( - self, - request: "LlmRequest", - session_id: str, - session: Any | None, - ) -> ToolResultBudgetResult: + async def _apply_scoped(self, request: LlmRequest, session: Session) -> ToolResultBudgetResult: """Apply budgeting while ``_runtime`` is bound to the current tenant.""" - async with self._session_lock(session_id): + 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) + state = await self._load_state(session.id) groups = self._collect_candidates(request, session) for group in groups: for candidate in group: @@ -340,39 +325,21 @@ async def _apply_scoped( ) -class ToolResultBudgetCallback: +class ToolResultBudgetHandler(BaseCompactSummarizerHandler): """Adapt the tool-result budget processor to before_model_callback.""" - advanced_memory_stage = 10 - - def __init__(self, budget: ToolResultBudget) -> None: + def __init__(self) -> None: """Store the budget processor run before model requests.""" - self._budget = budget + self._budget: ToolResultBudget | None = None - @property - def budget(self) -> ToolResultBudget: - """Return the budget processor used by this callback.""" - return self._budget - - async def __call__(self, ctx: "InvocationContext", request: "LlmRequest") -> None: + @override + async def handle(self, ctx: InvocationContext, request: LlmRequest) -> None: """Apply tool-result budgeting without truncating model calls.""" - await self._budget.apply(request, session_id=ctx.session_id, ctx=ctx) + summarizer = self.get_summarizer(ctx) + from ._auto_compact import AdvancedAutoCompactSummarizer + if not isinstance(summarizer, AdvancedAutoCompactSummarizer): + raise ValueError("Summarizer is not an AdvancedAutoCompactSummarizer") + if self._budget is None: + self._budget = ToolResultBudget(summarizer.runtime) + await self._budget.apply(request, ctx=ctx) return None - - -def setup_tool_result_budget( - agent: "LlmAgent", - memory_runtime: SessionCompactRuntime, -) -> ToolResultBudget: - """Install the budget callback while preserving existing callbacks.""" - budget = ToolResultBudget(memory_runtime) - callback = ToolResultBudgetCallback(budget) - existing_budget = install_staged_callback( - agent, - callback, - callback_type=ToolResultBudgetCallback, - component_attribute="budget", - memory_runtime=memory_runtime, - conflict_message="Tool result budget is already configured with another runtime", - ) - return existing_budget or budget diff --git a/trpc_agent_sdk/sessions/compact/advanced/_utils.py b/trpc_agent_sdk/sessions/compact/advanced/_utils.py new file mode 100644 index 000000000..40dea0d9a --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/advanced/_utils.py @@ -0,0 +1,74 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Advanced utils for compact session manager.""" + +import hashlib +import json +from contextlib import contextmanager +from typing import Any +from typing import Iterator + +from trpc_agent_sdk.context import AgentContext +from trpc_agent_sdk.types import Content + +INTERNAL_COMPACTION_METADATA_KEY = "_trpc_agent_internal_compaction" +_MISSING = object() + + +@contextmanager +def internal_compaction_call(agent_context: AgentContext | None) -> Iterator[None]: + """Prevent the compact filter from recursively handling its own model call.""" + if agent_context is None: + yield + return + previous = agent_context.get_metadata(INTERNAL_COMPACTION_METADATA_KEY, _MISSING) + agent_context.with_metadata(INTERNAL_COMPACTION_METADATA_KEY, True) + try: + yield + finally: + if previous is _MISSING: + agent_context.metadata.pop(INTERNAL_COMPACTION_METADATA_KEY, None) + else: + agent_context.with_metadata(INTERNAL_COMPACTION_METADATA_KEY, previous) + + +def content_signature(content: Content) -> str: + """Generate a stable signature that preserves message identity.""" + parts: list[dict[str, Any]] = [] + for part in content.parts or []: + if part.text is not None: + parts.append({ + "type": "text", + "sha256": hashlib.sha256(part.text.encode("utf-8")).hexdigest(), + }) + elif part.function_call is not None: + parts.append({ + "type": "function_call", + "id": getattr(part.function_call, "id", None), + "name": part.function_call.name, + }) + elif part.function_response is not None: + parts.append({ + "type": "function_response", + "id": getattr(part.function_response, "id", None), + "name": part.function_response.name, + }) + elif part.executable_code is not None: + parts.append({"type": "executable_code"}) + elif part.code_execution_result is not None: + parts.append({"type": "code_execution_result"}) + else: + parts.append({"type": "other"}) + serialized = json.dumps( + { + "role": content.role, + "parts": parts + }, + ensure_ascii=False, + sort_keys=True, + separators=(",", ":"), + ) + return hashlib.sha256(serialized.encode("utf-8")).hexdigest() diff --git a/trpc_agent_sdk/sessions/compact/default/__init__.py b/trpc_agent_sdk/sessions/compact/default/__init__.py new file mode 100644 index 000000000..092b2dda5 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/default/__init__.py @@ -0,0 +1,34 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Default compact session manager.""" + +from ._checker import CheckSummarizerFunction +from ._checker import set_summarizer_token_threshold +from ._checker import set_summarizer_events_count_threshold +from ._checker import set_summarizer_time_interval_threshold +from ._checker import set_summarizer_important_content_threshold +from ._checker import set_summarizer_conversation_threshold +from ._checker import set_summarizer_check_functions_by_and +from ._checker import set_summarizer_check_functions_by_or +from ._summarizer import DEFAULT_SUMMARIZER_PROMPT +from ._summarizer import DefaultSessionSummary +from ._summarizer import DefaultSessionSummarizer +from ._summarizer_manager import DefaultSessionSummarizerManager + +__all__ = [ + "DEFAULT_SUMMARIZER_PROMPT", + "DefaultSessionSummarizer", + "DefaultSessionSummarizerManager", + "DefaultSessionSummary", + "CheckSummarizerFunction", + "set_summarizer_token_threshold", + "set_summarizer_events_count_threshold", + "set_summarizer_time_interval_threshold", + "set_summarizer_important_content_threshold", + "set_summarizer_conversation_threshold", + "set_summarizer_check_functions_by_and", + "set_summarizer_check_functions_by_or", +] diff --git a/trpc_agent_sdk/sessions/_summarizer_checker.py b/trpc_agent_sdk/sessions/compact/default/_checker.py similarity index 96% rename from trpc_agent_sdk/sessions/_summarizer_checker.py rename to trpc_agent_sdk/sessions/compact/default/_checker.py index 37863cc9b..50d97d38d 100644 --- a/trpc_agent_sdk/sessions/_summarizer_checker.py +++ b/trpc_agent_sdk/sessions/compact/default/_checker.py @@ -14,7 +14,7 @@ from trpc_agent_sdk.events import Event from trpc_agent_sdk.log import logger -from ._session import Session +from ..._session import Session CheckSummarizerFunction = Callable[[Session], bool] @@ -137,10 +137,7 @@ def set_summarizer_conversation_threshold(conversation_count: int = 100) -> Chec """ def _decorator(session: Session) -> bool: - if session.conversation_count > conversation_count: - session.conversation_count = 0 - return True - return False + return session.conversation_count > conversation_count return _decorator diff --git a/trpc_agent_sdk/sessions/_session_summarizer.py b/trpc_agent_sdk/sessions/compact/default/_summarizer.py similarity index 96% rename from trpc_agent_sdk/sessions/_session_summarizer.py rename to trpc_agent_sdk/sessions/compact/default/_summarizer.py index da7171107..0c6d826b0 100644 --- a/trpc_agent_sdk/sessions/_session_summarizer.py +++ b/trpc_agent_sdk/sessions/compact/default/_summarizer.py @@ -34,11 +34,13 @@ from typing import Dict from typing import List from typing import Optional +from typing_extensions import override from pydantic import BaseModel from pydantic import ConfigDict from pydantic import Field +from trpc_agent_sdk.abc import CompactSummarizerABC from trpc_agent_sdk.context import InvocationContext from trpc_agent_sdk.events import Event from trpc_agent_sdk.log import logger @@ -47,10 +49,10 @@ from trpc_agent_sdk.types import Content from trpc_agent_sdk.types import Part -from ._session import Session -from ._summarizer_checker import CheckSummarizerFunction -from ._summarizer_checker import set_summarizer_conversation_threshold -from ._utils import find_events_for_summary +from ..._session import Session +from ..._utils import find_events_for_summary +from ._checker import CheckSummarizerFunction +from ._checker import set_summarizer_conversation_threshold DEFAULT_SUMMARIZER_PROMPT = dedent("""\ Please summarize the following conversation, focusing on: @@ -68,7 +70,7 @@ Summary:""") -class SessionSummary(BaseModel): +class DefaultSessionSummary(BaseModel): """Represents a summary of a session's conversation history. This class encapsulates the summary information including the summary text, @@ -88,6 +90,8 @@ class SessionSummary(BaseModel): """The timestamp when the summary was created.""" metadata: Dict[str, Any] = Field(default_factory=dict) """Additional metadata about the summarization.""" + model_name: str = "" + """The name of the model used for summarization.""" def get_compression_ratio(self) -> float: """Get the compression ratio achieved by summarization. @@ -111,13 +115,13 @@ def to_dict(self) -> Dict[str, Any]: "original_event_count": self.original_event_count, "compressed_event_count": self.compressed_event_count, "summary_timestamp": self.summary_timestamp, - "model_name": self.model.name, + "model_name": self.model_name, "compression_ratio": self.get_compression_ratio(), "metadata": self.metadata, } -class SessionSummarizer: +class DefaultSessionSummarizer(CompactSummarizerABC): """Summarizes conversation history to reduce memory usage. This class provides functionality to compress long conversation histories @@ -155,6 +159,7 @@ def model(self) -> LLMModel: """Get the LLM model for summarization.""" return self._model + @override async def should_summarize(self, session: Session) -> bool: """Check if the session should be summarized. @@ -291,6 +296,7 @@ def _extract_conversation_text(self, events: List[Event]) -> str: current_author = author current_branch = branch current_text = "" + continue if is_partial and current_author == author and current_text and current_branch == branch: # Merge with current accumulated text current_text += event_text @@ -355,6 +361,7 @@ def _create_summarization_prompt(self, conversation_text: str) -> str: """ return self._summarizer_prompt.format(conversation_text=conversation_text) + @override async def create_session_summary_by_events( self, events: List[Event], @@ -413,6 +420,7 @@ async def create_session_summary_by_events( logger.error("Failed to compress session %s: %s", session_id, ex, exc_info=True) return None, events + @override async def create_session_summary(self, session: Session, ctx: InvocationContext | None = None, @@ -436,6 +444,7 @@ async def create_session_summary(self, store_historical_events=store_historical_events) return summary_text + @override def get_summary_metadata(self) -> Dict[str, Any]: """Get metadata about the summarizer configuration. diff --git a/trpc_agent_sdk/sessions/_summarizer_manager.py b/trpc_agent_sdk/sessions/compact/default/_summarizer_manager.py similarity index 80% rename from trpc_agent_sdk/sessions/_summarizer_manager.py rename to trpc_agent_sdk/sessions/compact/default/_summarizer_manager.py index 2a07f58c0..65fd5af0b 100644 --- a/trpc_agent_sdk/sessions/_summarizer_manager.py +++ b/trpc_agent_sdk/sessions/compact/default/_summarizer_manager.py @@ -31,18 +31,20 @@ from typing import Any from typing import Dict from typing import Optional +from typing_extensions import override -from trpc_agent_sdk.abc import SessionServiceABC +from trpc_agent_sdk.abc import CompactSummarizerManagerABC +from trpc_agent_sdk.abc import CompactTrigger from trpc_agent_sdk.context import InvocationContext from trpc_agent_sdk.log import logger from trpc_agent_sdk.models import LLMModel -from ._session import Session -from ._session_summarizer import SessionSummarizer -from ._session_summarizer import SessionSummary +from ..._session import Session +from ._summarizer import DefaultSessionSummarizer +from ._summarizer import DefaultSessionSummary -class SummarizerSessionManager: +class DefaultSessionSummarizerManager(CompactSummarizerManagerABC): """Session service with automatic summarization capabilities. This service extends the basic session service with automatic @@ -53,8 +55,9 @@ class SummarizerSessionManager: def __init__( self, model: LLMModel, - summarizer: Optional[SessionSummarizer] = None, + summarizer: Optional[DefaultSessionSummarizer] = None, auto_summarize: bool = True, + compact_trigger: CompactTrigger = CompactTrigger.AFTER_TURN, ): """Initialize the summarizer session service. @@ -64,33 +67,13 @@ def __init__( summarizer: The session summarizer to use auto_summarize: Whether to automatically summarize sessions """ - self._base_service = None if not summarizer and model: - summarizer = SessionSummarizer(model=model) - self._summarizer: SessionSummarizer = summarizer + summarizer = DefaultSessionSummarizer(model=model) + super().__init__(summarizer, compact_trigger=compact_trigger) self._auto_summarize = auto_summarize - self._summarizer_cache: Dict[str, Dict[str, Dict[str, SessionSummary]]] = {} - - def set_session_service(self, session_service: SessionServiceABC, force: bool = False) -> None: - """Set the session service to use. - - Args: - session_service: The session service to use - force: Whether to force update even if already set - """ - if not self._base_service or force: - self._base_service = session_service - - def set_summarizer(self, summarizer: SessionSummarizer, force: bool = False) -> None: - """Set the summarizer to use. - - Args: - summarizer: The summarizer to use - force: Whether to force update even if already set - """ - if not self._summarizer or force: - self._summarizer = summarizer + self._summarizer_cache: Dict[str, Dict[str, Dict[str, DefaultSessionSummary]]] = {} + @override async def create_session_summary(self, session: Session, force: bool = False, @@ -100,6 +83,8 @@ async def create_session_summary(self, Args: session: The session to summarize """ + if self.compact_trigger != CompactTrigger.AFTER_TURN: + return is_should_summarize = await self.should_summarize_session(session) or force # Check if session should be summarized if is_should_summarize: @@ -122,25 +107,28 @@ async def create_session_summary(self, self._summarizer_cache[app_name] = {} if user_id not in self._summarizer_cache[app_name]: self._summarizer_cache[app_name][user_id] = {} - self._summarizer_cache[app_name][user_id][session.id] = SessionSummary( + self._summarizer_cache[app_name][user_id][session.id] = DefaultSessionSummary( session_id=session.id, summary_text=summary_text, original_event_count=original_event_count, compressed_event_count=len(session.events), summary_timestamp=time.time(), + model_name=self._summarizer.model.name, ) + session.conversation_count = 0 # Update the stored session if self._base_service: await self._base_service.update_session(session) - async def get_session_summary(self, session: Session) -> Optional[SessionSummary]: + @override + async def get_session_summary(self, session: Session) -> Optional[DefaultSessionSummary]: """Get a summary of a session. Args: session: The session to summarize Returns: - SessionSummary if successful, None otherwise + DefaultSessionSummary if successful, None otherwise """ if not self._summarizer or not self._summarizer_cache: return None diff --git a/trpc_agent_sdk/tools/__init__.py b/trpc_agent_sdk/tools/__init__.py index 20a937149..071b5365a 100644 --- a/trpc_agent_sdk/tools/__init__.py +++ b/trpc_agent_sdk/tools/__init__.py @@ -14,9 +14,9 @@ # Lazy re-export — see ``_LAZY_REEXPORTS`` below. from trpc_agent_sdk.agents.sub_agent import DynamicSubAgentTool as DynamicSubAgentTool # noqa: F401 from trpc_agent_sdk.agents.sub_agent import SpawnSubAgentTool as SpawnSubAgentTool # noqa: F401 - from trpc_agent_sdk.tools._advanced_memory_tool import AdvancedMemoryTools as AdvancedMemoryTools # noqa: F401 - from trpc_agent_sdk.tools._advanced_memory_tool import ( # noqa: F401 - create_advanced_memory_tools as create_advanced_memory_tools, ) + # from trpc_agent_sdk.tools._advanced_memory_tool import AdvancedMemoryTools as AdvancedMemoryTools # noqa: F401 + # from trpc_agent_sdk.tools._advanced_memory_tool import ( # noqa: F401 + # create_advanced_memory_tools as create_advanced_memory_tools, ) from ._agent_tool import AGENT_TOOL_APP_NAME_SUFFIX from ._agent_tool import AgentTool @@ -202,14 +202,14 @@ # the tools package free of optional file/web tool dependencies) but exposed # here for discoverability. Not in ``__all__`` so ``import *`` stays lazy. _LAZY_REEXPORTS = { - "AdvancedMemoryTools": ( - "trpc_agent_sdk.tools._advanced_memory_tool", - "AdvancedMemoryTools", - ), - "create_advanced_memory_tools": ( - "trpc_agent_sdk.tools._advanced_memory_tool", - "create_advanced_memory_tools", - ), + # "AdvancedMemoryTools": ( + # "trpc_agent_sdk.tools._advanced_memory_tool", + # "AdvancedMemoryTools", + # ), + # "create_advanced_memory_tools": ( + # "trpc_agent_sdk.tools._advanced_memory_tool", + # "create_advanced_memory_tools", + # ), "DynamicSubAgentTool": ("trpc_agent_sdk.agents.sub_agent", "DynamicSubAgentTool"), "SpawnSubAgentTool": ("trpc_agent_sdk.agents.sub_agent", "SpawnSubAgentTool"), } diff --git a/trpc_agent_sdk/tools/_advanced_memory_tool.py b/trpc_agent_sdk/tools/_advanced_memory_tool.py index 76bce5249..e69de29bb 100644 --- a/trpc_agent_sdk/tools/_advanced_memory_tool.py +++ b/trpc_agent_sdk/tools/_advanced_memory_tool.py @@ -1,171 +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 long-term memory read/write tools for standalone Advanced Memory.""" - -from __future__ import annotations - -import asyncio -import re -from typing import Any - -from trpc_agent_sdk.sessions.compact._formats import MemoryDocument -from trpc_agent_sdk.sessions.compact._formats import MemoryIndexEntry -from trpc_agent_sdk.sessions.compact._formats import MemoryType -from trpc_agent_sdk.sessions.compact._formats import memory_freshness -from trpc_agent_sdk.sessions.compact._formats import parse_memory_updated_at -from trpc_agent_sdk.advanced_memory._runtime import AdvancedMemoryRuntime - -from ._function_tool import FunctionTool - -ADVANCED_MEMORY_TOOL_NAMES = frozenset({ - "save_memory", - "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 - - -class AdvancedMemoryTools: - """Wrap long-term memory storage as three official Agent-callable tools.""" - - def __init__(self, runtime: AdvancedMemoryRuntime) -> None: - """Store the runtime and create tenant-scoped index update locks.""" - self._runtime = runtime - self._index_locks: dict[str, asyncio.Lock] = {} - self._tools = ( - FunctionTool(self.save_memory), - FunctionTool(self.read_memory), - FunctionTool(self.list_memory_index), - ) - - @property - def runtime(self) -> AdvancedMemoryRuntime: - """Return the Advanced Memory runtime bound to these tools.""" - return self._runtime - - 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, - filename: str, - name: str, - description: str, - 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: - resolved_type = MemoryType(memory_type) - except ValueError as exc: - allowed = ", ".join(item.value for item in MemoryType) - raise ValueError(f"memory_type must be one of: {allowed}") from exc - document = MemoryDocument( - name=name, - description=description, - memory_type=resolved_type, - content=content, - ) - 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 runtime.long_term_memory.read_index()) - new_entry = MemoryIndexEntry( - name=name, - filename=path.name, - summary=summary, - ) - entries = [entry for entry in entries if entry.filename != new_entry.filename] - entries.insert(0, new_entry) - 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": 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, tool_context: Any | None = None) -> dict: - """Read a complete long-term memory by its filename in MEMORY.md.""" - 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) - freshness = memory_freshness(updated_at) - return { - "found": - True, - "filename": - filename, - "content": - content, - "updated_at": - updated_at.isoformat() if updated_at is not None else None, - "freshness": - freshness, - "freshness_notice": (f"This memory was last updated {freshness}. 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."), - } - - 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": _memory_index_reference(runtime), - "index": await runtime.long_term_memory.read_index(), - } - - -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() From 11b3f0c947bb4131336491f265815a9430dcfd86 Mon Sep 17 00:00:00 2001 From: congkechen Date: Fri, 11 Sep 2026 17:31:08 +0800 Subject: [PATCH 6/6] =?UTF-8?q?feature:=20=E4=BC=98=E5=8C=96=20advanced=20?= =?UTF-8?q?memory=20=E9=95=BF=E6=9C=9F=E8=AE=B0=E5=BF=86=E5=AE=9E=E7=8E=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitignore | 4 + .../memory_service_with_advanced_memory/.env | 10 +- .../README.md | 150 +++++--- .../agent/agent.py | 2 + .../run_agent.py | 129 +++++++ .../.env | 4 +- .../README.md | 328 +++++------------- .../run_agent.py | 9 +- .../.env | 9 +- .../README.md | 189 ++++------ .../run_agent.py | 9 +- .../.env | 10 - .../README.md | 103 ------ .../agent/__init__.py | 5 - .../agent/agent.py | 39 --- .../agent/config.py | 19 - .../agent/prompts.py | 10 - .../agent/tools.py | 11 - .../run_agent.py | 130 ------- .../.env | 11 - .../README.md | 100 ------ .../agent/__init__.py | 5 - .../agent/agent.py | 39 --- .../agent/config.py | 19 - .../agent/prompts.py | 10 - .../agent/tools.py | 11 - .../run_agent.py | 126 ------- .../test_advanced_memory_tools.py | 6 +- tests/advanced_memory/test_memory_context.py | 36 +- tests/advanced_memory/test_preload_memory.py | 52 ++- trpc_agent_sdk/memory/__init__.py | 5 +- .../memory/_advanced_memory_service.py | 126 +++++++ .../memory/advanced_memory/__init__.py | 53 +++ .../memory/advanced_memory/_config.py | 84 +++++ .../memory/advanced_memory/_formats.py | 133 +++++++ .../memory/advanced_memory/_integration.py | 106 ++++++ .../memory/advanced_memory/_memory_context.py | 126 +++++++ .../memory/advanced_memory/_paths.py | 110 ++++++ .../memory/advanced_memory/_preload_memory.py | 286 +++++++++++++++ .../memory/advanced_memory/_redis_stores.py | 197 +++++++++++ .../memory/advanced_memory/_runtime.py | 202 +++++++++++ .../memory/advanced_memory/_sql_stores.py | 315 +++++++++++++++++ .../memory/advanced_memory/_storage.py | 221 ++++++++++++ trpc_agent_sdk/runners.py | 5 + trpc_agent_sdk/sessions/compact/_callbacks.py | 49 +++ .../advanced/_compaction_memory_extractor.py | 13 +- .../compact/default/_summarizer_manager.py | 2 +- trpc_agent_sdk/tools/__init__.py | 22 +- trpc_agent_sdk/tools/_advanced_memory_tool.py | 159 +++++++++ 49 files changed, 2682 insertions(+), 1117 deletions(-) delete mode 100644 examples/session_service_with_advanced_memory_redis/.env delete mode 100644 examples/session_service_with_advanced_memory_redis/README.md delete mode 100644 examples/session_service_with_advanced_memory_redis/agent/__init__.py delete mode 100644 examples/session_service_with_advanced_memory_redis/agent/agent.py delete mode 100644 examples/session_service_with_advanced_memory_redis/agent/config.py delete mode 100644 examples/session_service_with_advanced_memory_redis/agent/prompts.py delete mode 100644 examples/session_service_with_advanced_memory_redis/agent/tools.py delete mode 100644 examples/session_service_with_advanced_memory_redis/run_agent.py delete mode 100644 examples/session_service_with_advanced_memory_sql/.env delete mode 100644 examples/session_service_with_advanced_memory_sql/README.md delete mode 100644 examples/session_service_with_advanced_memory_sql/agent/__init__.py delete mode 100644 examples/session_service_with_advanced_memory_sql/agent/agent.py delete mode 100644 examples/session_service_with_advanced_memory_sql/agent/config.py delete mode 100644 examples/session_service_with_advanced_memory_sql/agent/prompts.py delete mode 100644 examples/session_service_with_advanced_memory_sql/agent/tools.py delete mode 100644 examples/session_service_with_advanced_memory_sql/run_agent.py create mode 100644 trpc_agent_sdk/memory/advanced_memory/__init__.py create mode 100644 trpc_agent_sdk/memory/advanced_memory/_config.py create mode 100644 trpc_agent_sdk/memory/advanced_memory/_formats.py create mode 100644 trpc_agent_sdk/memory/advanced_memory/_integration.py create mode 100644 trpc_agent_sdk/memory/advanced_memory/_memory_context.py create mode 100644 trpc_agent_sdk/memory/advanced_memory/_paths.py create mode 100644 trpc_agent_sdk/memory/advanced_memory/_preload_memory.py create mode 100644 trpc_agent_sdk/memory/advanced_memory/_redis_stores.py create mode 100644 trpc_agent_sdk/memory/advanced_memory/_runtime.py create mode 100644 trpc_agent_sdk/memory/advanced_memory/_sql_stores.py create mode 100644 trpc_agent_sdk/memory/advanced_memory/_storage.py create mode 100644 trpc_agent_sdk/sessions/compact/_callbacks.py diff --git a/.gitignore b/.gitignore index 91426e394..60da44eaa 100644 --- a/.gitignore +++ b/.gitignore @@ -28,3 +28,7 @@ pyrightconfig.json # spec-workflow tool artifacts .spec-workflow + +# Local-only examples +examples/session_service_with_advanced_memory_sql/ +examples/session_service_with_advanced_memory_redis/ diff --git a/examples/memory_service_with_advanced_memory/.env b/examples/memory_service_with_advanced_memory/.env index e4183ff5b..8061a2bc8 100644 --- a/examples/memory_service_with_advanced_memory/.env +++ b/examples/memory_service_with_advanced_memory/.env @@ -1,12 +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= - -# Optional TTL settings. Leave empty to disable automatic expiration. -M_TTL=120 -SESSION_TTL=60 +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 1e23757f7..e81a4aea4 100644 --- a/examples/memory_service_with_advanced_memory/README.md +++ b/examples/memory_service_with_advanced_memory/README.md @@ -1,61 +1,111 @@ -# Standard SessionService + Advanced Compact + Advanced Memory - -本示例使用统一后的组合方式: - -```text -InMemorySessionService -└── AdvancedSessionCompactManager - ├── Session Memory - ├── Tool Result Budget - ├── History Snip - ├── Microcompact - └── AutoCompact - -AdvancedMemoryService -├── save_memory -├── read_memory -├── list_memory_index -└── long-term memory injection +# 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 ``` -不再使用独立的 Advanced SessionService。Session 的创建、Event 保存和状态管理始终 -由标准 `InMemorySessionService`、`RedisSessionService` 或 `SqlSessionService` -负责;Advanced Compact 通过 `BaseSessionCompactManager` 生命周期接入。 - -## 核心组装 - -```python -config = AdvancedMemoryServiceConfig( - root_dir=Path(__file__).resolve().parent, -) - -session_service = InMemorySessionService( - session_config=SessionServiceConfig( - store_historical_events=True, - ), - session_compact_manager=AdvancedSessionCompactManager( - config=AdvancedCompactConfig(), - ), -) - -memory_service = AdvancedMemoryService(config=config) -runner = Runner( - app_name="advanced_memory_demo", - agent=agent, - session_service=session_service, - memory_service=memory_service, -) +## 代码构建 + +```bash +git clone https://github.com/trpc-group/trpc-agent-python.git +cd trpc-agent-python +./build.sh +source .venv/bin/activate ``` -Session Compact 与 Advanced Memory 使用独立配置和 Runtime。Compact 只使用 -SessionService 的 events、historical_events 和 state。 +如果已经在当前项目中创建了 Python 3.10+ 虚拟环境,也可以直接安装: -## 运行 +```bash +python -m pip install -e . +``` -在 `.env` 中配置模型,然后执行: +## 运行 ```bash +cd examples/memory_service_with_advanced_memory python run_agent.py ``` -示例会在两个 Session 中使用同一用户,验证用户级长期记忆可以跨 Session 使用。 +示例会使用同一用户运行多个会话,验证长期记忆可以在不同会话之间复用。 + +## 运行结果(实测) + +```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. + +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. + +Is there anything else you'd like me to remember or clarify about your language preferences? +``` diff --git a/examples/memory_service_with_advanced_memory/agent/agent.py b/examples/memory_service_with_advanced_memory/agent/agent.py index 8f93758ec..2165a5f26 100644 --- a/examples/memory_service_with_advanced_memory/agent/agent.py +++ b/examples/memory_service_with_advanced_memory/agent/agent.py @@ -7,6 +7,7 @@ from trpc_agent_sdk.agents import LlmAgent from trpc_agent_sdk.models import OpenAIModel +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerFilter from .config import get_model_config from .prompts import INSTRUCTION @@ -22,6 +23,7 @@ def create_agent() -> LlmAgent: model_name=model_name, api_key=api_key, base_url=base_url, + filters=[AdvancedAutoCompactSummarizerFilter()], ), instruction=INSTRUCTION, ) diff --git a/examples/memory_service_with_advanced_memory/run_agent.py b/examples/memory_service_with_advanced_memory/run_agent.py index e69de29bb..b1027950c 100644 --- a/examples/memory_service_with_advanced_memory/run_agent.py +++ b/examples/memory_service_with_advanced_memory/run_agent.py @@ -0,0 +1,129 @@ +#!/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 the two-session Advanced Memory demonstration.""" + +import asyncio +from pathlib import Path + +from dotenv import load_dotenv +from trpc_agent_sdk.memory.advanced_memory import AdvancedMemoryServiceConfig +from trpc_agent_sdk.memory import AdvancedMemoryService +from trpc_agent_sdk.sessions import InMemorySessionService +from trpc_agent_sdk.sessions import SessionServiceConfig +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizer +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerConfig +from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerManager +from trpc_agent_sdk.types import Content +from trpc_agent_sdk.types import Part + +from agent.agent import create_agent + +load_dotenv(Path(__file__).with_name(".env"), override=True) + + +def create_session_service() -> InMemorySessionService: + """Create the session service with the independent Compact manager.""" + compact_manager = AdvancedAutoCompactSummarizerManager( + summarizer=AdvancedAutoCompactSummarizer( + config=AdvancedAutoCompactSummarizerConfig(), + ), + ) + return InMemorySessionService( + session_config=SessionServiceConfig( + ttl=SessionServiceConfig.create_ttl_config( + enable=True, + ttl_seconds=60, + cleanup_interval_seconds=5, + ), + store_historical_events=True, + ), + summarizer_manager=compact_manager, + ) + + +def create_memory_service() -> AdvancedMemoryService: + """Create the independent long-term Advanced Memory service.""" + memory_config = AdvancedMemoryServiceConfig( + root_dir=Path(__file__).resolve().parent, + memory_ttl_seconds=120, + memory_focus_instruction=("特别关注并主动记住用户长期稳定的兴趣爱好、" + "编程语言偏好、开发习惯和测试习惯。"), + ) + return AdvancedMemoryService(config=memory_config) + + +async def run_turn(runner, *, user_id: str, session_id: str, prompt: str) -> None: + """Run one turn and print tool activity and the final response.""" + print(f"\n👤 [{session_id}] {prompt}") + content = Content(parts=[Part.from_text(text=prompt)]) + async for event in runner.run_async( + user_id=user_id, + session_id=session_id, + new_message=content, + ): + if not event.content or not event.content.parts: + continue + for part in event.content.parts: + if part.function_call: + print(f"🔧 {part.function_call.name}({part.function_call.args})") + elif part.function_response: + print(f"📊 {part.function_response.response}") + elif part.text and not part.thought and not event.partial: + print(f"🤖 {part.text}") + + +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 = [ + ("Please remember that my favorite programming language is Python. " + "Save this as a user preference."), + "I use Python mainly for backend services and data processing.", + "I prefer typed Python code with clear dataclasses and small modules.", + "For testing Python code, I usually prefer pytest and focused unit tests.", + "When documenting projects, I prefer concise examples with runnable commands.", + ] + for prompt in session_one_prompts: + await run_turn( + runner, + user_id="demo-user", + session_id="session-1", + prompt=prompt, + ) + + await run_turn( + runner, + user_id="demo-user", + session_id="session-1", + prompt="Summarize what you learned about my Python development preferences.", + ) + + await run_turn( + runner, + user_id="demo-user", + session_id="session-2", + prompt="What do you remember about my favorite programming language?", + ) + + finally: + await runner.close() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/examples/memory_service_with_advanced_memory_redis/.env b/examples/memory_service_with_advanced_memory_redis/.env index 6a46edc83..52b372762 100644 --- a/examples/memory_service_with_advanced_memory_redis/.env +++ b/examples/memory_service_with_advanced_memory_redis/.env @@ -3,6 +3,4 @@ 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= - -M_TTL=120 \ No newline at end of file +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 index d9a8273d0..dcd6a2a34 100644 --- a/examples/memory_service_with_advanced_memory_redis/README.md +++ b/examples/memory_service_with_advanced_memory_redis/README.md @@ -1,30 +1,40 @@ -# Advanced Memory Redis 示例 +# Advanced Memory Redis 持久化示例 -本示例演示如何将 Advanced Memory 的本地文件存储切换为 Redis,并验证: +本示例演示如何使用 `AdvancedMemoryService` 将长期记忆保存到 Redis,实现跨会话、跨 Python 进程的持久化记忆。 -- Redis:`AdvancedMemoryService(storage_backend="redis")` -- 长期 memory 可以跨 Python 进程持久化; -- 同一用户在不同 `session_id` 中可以读取自己的长期 memory; -- session 相关数据和长期 memory 可以分别设置 TTL; -- Redis 中的 Markdown、Stream 和索引数据如何组织。 +## 关键特性 -本示例只关注长期 Memory 的 Redis 持久化: +- **主动式记忆**:Agent 根据对话内容主动调用工具保存长期有效的信息。 +- **记忆分类**:每条记忆包含名称、描述、类型、摘要和详细内容。 +- **基于记忆索引的记忆召回**:先读取 Redis 中的 `MEMORY.md` 索引, + 再读取与问题相关的记忆内容。 +- **Redis 持久化**:多个进程或实例使用相同的 Redis、应用名和用户 ID时,可以访问同一份长期记忆。 -```text -AdvancedMemoryService -└── Redis 保存长期 memory index 和 topic +## Agent 层级结构说明 -Runner -└── InMemorySessionService(仅用于运行示例) -``` +`AdvancedMemoryService` 通过 `Runner` 绑定到 Agent,并根据配置使用 Redis 保存记忆索引和记忆主题。Agent 通过三个工具主动管理长期记忆。 + +## 关键代码解释 + +### `save_memory` + +保存或更新一条长期记忆,同时更新 Redis 中的记忆索引。 + +### `list_memory_index` + +读取当前用户的记忆索引,帮助 Agent 找到与当前问题相关的记忆文件。 + +### `read_memory` + +根据索引中的文件名读取完整记忆内容。 ## 环境要求 -- Python 3.10+,推荐 Python 3.12; -- 可访问的 Redis 服务; -- 可正常调用的模型服务。 +- Python 3.10 或更高版本 +- 可访问的 Redis 服务 +- 一个可访问的 OpenAI 兼容模型服务 -如果还没有 Redis,可以使用 Docker: +**启动本地 Redis:** ```bash docker run --name advanced-memory-redis \ @@ -32,7 +42,13 @@ docker run --name advanced-memory-redis \ -d redis:7-alpine ``` -容器已创建过时不要重复执行 `docker run`,直接启动: +然后在当前目录的 `.env` 中配置: + +```dotenv +REDIS_URL=redis://localhost:6379/0 +``` + +如果容器已经存在,执行: ```bash docker start advanced-memory-redis @@ -45,87 +61,65 @@ docker exec advanced-memory-redis redis-cli PING # PONG ``` -## Redis 配置方式 - -### 方式一:使用完整连接串 - -在当前目录的 `.env` 中配置: - -```dotenv -REDIS_URL=redis://localhost:6379/0 -``` - -带密码: +如果使用已有的**远程 Redis 服务**,不需要执行 Docker 命令,只需要在当前目录的`.env` 中配置 Redis 连接信息: ```dotenv REDIS_URL=redis://:password@redis.example.com:6379/0 ``` -Redis ACL 用户名和密码: +如果 Redis 使用 ACL 用户名和密码: ```dotenv REDIS_URL=redis://username:password@redis.example.com:6379/0 ``` -启用 TLS: +启用 TLS 时使用 `rediss` 协议: ```dotenv -REDIS_URL=rediss://:password@redis.example.com:6380/0 -``` - -密码包含 `@`、`:`、`/`、`#` 等特殊字符时,需要进行 URL 编码。 - -### 方式二:分别配置连接参数 - -也可以不设置 `REDIS_URL`,改为: - -```dotenv -REDIS_HOST=127.0.0.1 -REDIS_PORT=6379 -REDIS_DB=0 -REDIS_USER= -REDIS_PASSWORD= -REDIS_TLS=false +REDIS_URL=rediss://username:password@redis.example.com:6380/0 ``` -云 Redis 使用示例: +也可以拆分配置: ```dotenv -REDIS_HOST=your-redis.example.com +REDIS_HOST=redis.example.com REDIS_PORT=6379 REDIS_DB=0 REDIS_USER=your-user REDIS_PASSWORD=your-password -REDIS_TLS=true +REDIS_TLS=false ``` -代码会优先使用 `REDIS_URL`;未设置时才根据上述字段构造连接串。 +代码会优先使用 `REDIS_URL`;未设置时,才会根据这些字段构造连接串。密码包含 `@`、`:`、`/`、`#` 等特殊字符时,需要进行 URL 编码。 -## 模型和 TTL 配置 +## 模型配置 -`.env` 示例: +在当前目录的 `.env` 中配置: ```dotenv TRPC_AGENT_API_KEY=your-api-key -TRPC_AGENT_BASE_URL=your-model-base-url +TRPC_AGENT_BASE_URL=https://your-llm-endpoint/v1 TRPC_AGENT_MODEL_NAME=your-model-name +``` -REDIS_URL=redis://localhost:6379/0 +Redis 配置请参考上面的本地 Redis 或远程 Redis 配置方式。 -# 长期 memory 的 TTL,单位为秒 -M_TTL=120 +## 代码构建 +```bash +git clone https://github.com/trpc-group/trpc-agent-python.git +cd trpc-agent-python +./build.sh +source .venv/bin/activate ``` -TTL 规则: +如果已经在当前项目中创建了 Python 3.10+ 虚拟环境,也可以直接安装: -- `M_TTL` 管理用户级长期 memory 的全部 Redis key; -- TTL 会在访问或写入时刷新,是“最后一次活动后过期”; -- `M_TTL` 必须设置为大于 0 的整数。 - -更多 Advanced Memory 配置请参考[Advanced Memory README](../memory_service_with_advanced_memory/README.md)。 +```bash +python -m pip install -e . +``` -## 运行示例 +## 运行 ```bash cd examples/memory_service_with_advanced_memory_redis @@ -133,98 +127,46 @@ source ../../.venv/bin/activate python run_agent.py ``` -脚本会自动启动两个独立的 Python 子进程: +脚本会依次启动写入和读取两个独立进程,验证 Redis 中的记忆可以跨进程和不同会话读取。也可以单独运行某个阶段: -```text -RUNNER A PROCESS -├── 使用 7 条对话模拟记忆建立过程 -└── Alice 的姓名和 favorite color 会被保存到长期 memory - -RUNNER B PROCESS -├── 使用新的 session -├── 询问 Alice 的 name -└── 询问 Alice 的 favorite color +```bash +python run_agent.py --phase write +python run_agent.py --phase read ``` -两个进程使用相同的: +## Redis 中的存储 -```text -app_name = advanced-memory-redis-demo -user_id = redis-demo-user -``` - -但使用不同的 `session_id`。第二个进程应该能够回答: +记忆索引和主题内容会以 Redis key 保存,key 前缀为: ```text -name: Alice -favorite color: blue +advanced-memory-redis-demo:v1:* ``` -这证明了 Redis 数据可以跨进程、跨 session 持久化。 - -也可以单独运行某个阶段: +查看本示例写入的 key: ```bash -python run_agent.py --phase write # Runner A -python run_agent.py --phase read # Runner B -``` - -## 最基本的构建方式 - -Redis 版本最核心的构建过程可以简化为三步: - -```python -redis_url = "redis://:password@localhost:6379/0" - -memory_service = AdvancedMemoryService( - AdvancedMemoryServiceConfig( - storage_backend="redis", - redis_url=redis_url, - memory_ttl_seconds=120, # from M_TTL; omit to disable expiration - ) -) - -runner = Runner( - app_name="advanced-memory-redis-demo", - agent=create_agent(), - session_service=InMemorySessionService(), - memory_service=memory_service, -) +docker exec advanced-memory-redis redis-cli --scan \ + --pattern 'advanced-memory-redis-demo:v1:*' ``` -其中: - -- 用户只需要配置长期 Memory 的 `M_TTL`; -- `AdvancedMemoryService` 只负责长期 memory; -- Session Service 的 Redis 高级压缩接入请看 - [`session_service_with_advanced_memory_redis`](../session_service_with_advanced_memory_redis/); -- 运行请求时通过 `user_id` 和 `session_id` 指定用户及会话。 +示例中的记忆 TTL 在代码的 `AdvancedMemoryServiceConfig` 中配置为 `memory_ttl_seconds=120`。 ## 运行结果(实测) -```text - user: Do you remember my name? -🔧 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:index', 'index': ''} -🤖 Assistant: I checked my long-term memory, but I'm afraid I don't have anything saved yet — the memory index is currently empty, so I don't know your name. - -If you'd like, just tell me your name (and anything else you'd like me to remember about you), and I'll save it so I can recall it in future conversations! +```txt +==================== WRITE PROCESS ==================== ------ Runner A, query 2 ----- +----- Runner A, query 1 ----- -📝 user: Do you remember my favorite color? -🔧 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:index', 'index': ''} -🤖 Assistant: I checked my long-term memory, but I don't have anything saved about your favorite color yet — my memory index is currently empty. +📝 user: Do you remember my name? +🤖 Assistant: -If you'd like, tell me your favorite color and I'll remember it for future conversations. 💬 ------ Runner A, query 3 ----- +🔧 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. -📝 user: what is the weather like in paris? -🔧 tool call: get_weather_report({'city': 'Paris'}) -📊 Tool Result: {'status': 'success', 'report': 'The weather in Paris is sunny with a temperature of 25 degrees Celsius.'} -🤖 Assistant: The weather in Paris is **sunny** with a temperature of **25°C**. ☀️ +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. 😊 ...... @@ -233,110 +175,16 @@ If you'd like, tell me your favorite color and I'll remember it for future conve ----- Runner B, query 1 ----- 📝 user: Do you remember my name? -🔧 tool call: read_memory({'filename': 'alice-identity.md'}) -📊 Tool Result: {'found': True, 'filename': 'alice-identity.md', 'content': "---\nname: Alice's identity\ndescription: User's name and personal preferences.\ntype: user\nupdated_at: 2026-09-08T07:17:40.559641+00:00\n---\n- Name: Alice\n- Favorite color: blue\n", 'updated_at': '2026-09-08T07:17:40.559641+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 do — your name is Alice! 😊 And I also remember that your favorite color is blue. +🔧 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': 'alice-identity.md'}) -📊 Tool Result: {'found': True, 'filename': 'alice-identity.md', 'content': "---\nname: Alice's identity\ndescription: User's name and personal preferences.\ntype: user\nupdated_at: 2026-09-08T07:17:40.559641+00:00\n---\n- Name: Alice\n- Favorite color: blue\n", 'updated_at': '2026-09-08T07:17:40.559641+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! According to your memory profile, your favorite color is **blue**. 💙 -``` - -## 查看 Redis 中的数据 - -进入 Redis CLI: - -```bash -docker exec -it advanced-memory-redis redis-cli -``` - -查看本示例写入的全部 Redis key: - -```redis -SCAN 0 MATCH advanced-memory-redis-demo:v1:* COUNT 100 -``` - -也可以在命令行中直接查看全部 key: - -```bash -docker exec advanced-memory-redis redis-cli --scan \ - --pattern 'advanced-memory-redis-demo:v1:*' -``` - -`SCAN` 不会像 `KEYS *` 一样阻塞 Redis,适合共享或云 Redis 环境。 - -## 查看 TTL - -长期 memory: - -```redis -TTL "advanced-memory-redis-demo:v1:{advanced-memory-redis-demo:redis-demo-user}:memory:index" -TTL "advanced-memory-redis-demo:v1:{advanced-memory-redis-demo:redis-demo-user}:memory:topic:user_favorite_project_code.md" -``` - -预期接近 `120`。 - -TTL 含义: - -```text --1 永不过期 --2 key 不存在或已经过期 -大于 0 剩余秒数 -``` - -## 清理测试数据 - -只删除本示例的 Advanced Memory key: - -```bash -docker exec advanced-memory-redis redis-cli --scan \ - --pattern 'advanced-memory-redis-demo:v1:*' \ - | xargs -r docker exec -i advanced-memory-redis redis-cli DEL -``` - -测试 Redis 独占一个数据库时,也可以清空当前数据库: - -```bash -docker exec -it advanced-memory-redis redis-cli FLUSHDB -``` - -`FLUSHDB` 会删除当前 Redis DB 中的所有数据,不要在共享或生产数据库执行。 - -## Redis 中的存储形式 - -### 长期 memory - -本地文件概念: - -```text -MEMORY/MEMORY.md -MEMORY/user_favorite_project_code.md -``` - -Redis 映射: - -```text -{prefix}:{app:user}:memory:index -{prefix}:{app:user}:memory:topic:user_favorite_project_code.md -``` - -类型都是 Redis String,内容是 Markdown。 - -topic 列表的辅助索引: - -```text -{prefix}:{app:user}:memory:topics -``` - -类型是 ZSet,member 是 topic 文件名,score 是更新时间。 - -memory TTL registry: - -```text -{prefix}:{app:user}:memory:keys -``` - -它记录该用户的所有长期 memory key,用于统一刷新 `M_TTL`。 +🔧 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/run_agent.py b/examples/memory_service_with_advanced_memory_redis/run_agent.py index e8b175ce0..d33e031af 100644 --- a/examples/memory_service_with_advanced_memory_redis/run_agent.py +++ b/examples/memory_service_with_advanced_memory_redis/run_agent.py @@ -14,13 +14,13 @@ from dotenv import load_dotenv from agent.agent import create_agent -from trpc_agent_sdk.advanced_memory import AdvancedMemoryServiceConfig +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")) +load_dotenv(Path(__file__).with_name(".env"), override=True) RUNNER_A_QUERIES = [ "Do you remember my name?", @@ -62,14 +62,13 @@ def build_redis_url_from_environment() -> str: def create_advanced_memory_service(redis_url: str) -> AdvancedMemoryService: """Create the long-term Advanced Memory service backed by Redis.""" - memory_ttl = os.getenv("M_TTL") config = AdvancedMemoryServiceConfig( storage_backend="redis", redis_url=redis_url, redis_key_prefix="advanced-memory-redis-demo:v1", - memory_ttl_seconds=int(memory_ttl) if memory_ttl else None, + memory_ttl_seconds=120, ) - return AdvancedMemoryService(config) + return AdvancedMemoryService(config=config) async def ask(runner: Runner, session_id: str, prompt: str) -> None: diff --git a/examples/memory_service_with_advanced_memory_sql/.env b/examples/memory_service_with_advanced_memory_sql/.env index 81dbccf4a..5e8b0ded0 100644 --- a/examples/memory_service_with_advanced_memory_sql/.env +++ b/examples/memory_service_with_advanced_memory_sql/.env @@ -4,10 +4,9 @@ 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 +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 -M_TTL=120 +# 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 index bc7c7dcf0..86f927347 100644 --- a/examples/memory_service_with_advanced_memory_sql/README.md +++ b/examples/memory_service_with_advanced_memory_sql/README.md @@ -1,174 +1,119 @@ -# Advanced Memory SQL 示例 +# Advanced Memory SQL 持久化示例 -本示例使用 SQL 保存 Advanced Memory,并验证同一用户的长期 memory 可以跨 Python 进程和不同 session 读取。 +本示例演示如何使用 `AdvancedMemoryService` 将长期记忆保存到 SQL 数据库,实现跨会话、跨 Python 进程的持久化记忆。 -- SQL:`AdvancedMemoryService(storage_backend="sql")` +## 关键特性 -```text -AdvancedMemoryService -└── SQL 保存长期 memory index 和 topic +- **主动式记忆**:Agent 根据对话内容主动调用工具保存长期有效的信息。 +- **记忆分类**:每条记忆包含名称、描述、类型、摘要和详细内容。 +- **基于记忆索引的记忆召回**:先读取数据库中的记忆索引,再读取与问题相关的记忆内容。 +- **SQL 持久化**:多个进程或实例使用相同的数据库、应用名和用户 ID 时,可以访问同一份长期记忆。 -Runner -└── InMemorySessionService(仅用于运行示例) -``` +## Agent 层级结构说明 + +`AdvancedMemoryService` 通过 `Runner` 绑定到 Agent,并根据配置使用 SQL 保存记忆索引和记忆主题。Agent 通过三个工具主动管理长期记忆。 + +## 关键代码解释 + +### `save_memory` + +保存或更新一条长期记忆,同时更新数据库中的记忆索引。 + +### `list_memory_index` -## 配置 +读取当前用户的记忆索引,帮助 Agent 找到与当前问题相关的记忆文件。 -默认使用 SQLite,运行示例不需要额外启动数据库: +### `read_memory` + +根据索引中的文件名读取完整记忆内容。 + +## 环境要求 + +- Python 3.10 或更高版本 +- SQLite 或可访问的 MySQL 数据库 +- 一个可访问的 OpenAI 兼容模型服务 + +默认使用 **SQLite**,不需要额外启动数据库: ```dotenv SQL_URL=sqlite:///advanced-memory-sql-demo.db SQL_IS_ASYNC=false ``` -使用 MySQL 时: +使用 **MySQL** 时: ```dotenv -SQL_URL=mysql+aiomysql://user:password@host:3306/trpc_agent_advanced_memory?charset=utf8mb4 +SQL_URL=mysql+aiomysql://user:password@localhost:3306/trpc_agent_advanced_memory?charset=utf8mb4 SQL_IS_ASYNC=true ``` -也可以通过 `MYSQL_USER`、`MYSQL_PASSWORD`、`MYSQL_HOST`、`MYSQL_PORT` 和 -`MYSQL_DB` 构造 MySQL URL。模型配置需要设置: +## 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=your-base-url +TRPC_AGENT_BASE_URL=https://your-llm-endpoint/v1 TRPC_AGENT_MODEL_NAME=your-model-name ``` -`M_TTL` 控制长期 memory 的过期时间,单位为秒。 - -更多 Advanced Memory 配置请参考[Advanced Memory README](../memory_service_with_advanced_memory/README.md)。 +也可以使用 `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 -cd examples/memory_service_with_advanced_memory_sql -python run_agent.py -``` - -脚本会依次启动两个独立进程: - -```text -RUNNER A PROCESS -├── 使用 7 条对话模拟记忆建立过程 -└── Alice 的姓名和 favorite color 会被保存到长期 memory - -RUNNER B PROCESS -├── 使用新的 session -├── 询问 Alice 的 name -└── 询问 Alice 的 favorite color ``` -Runner B 应该能够回答: +如果已经在当前项目中创建了 Python 3.10+ 虚拟环境,也可以直接安装: -```text -name: Alice -favorite color: blue +```bash +python -m pip install -e . ``` -也可以单独运行: +## 运行 ```bash -python run_agent.py --phase write # Runner A -python run_agent.py --phase read # Runner B +cd examples/memory_service_with_advanced_memory_sql +source ../../.venv/bin/activate +python run_agent.py ``` -第一次运行后,SQLite 文件 `advanced-memory-sql-demo.db` 会自动创建, -Advanced Memory 的表也会自动创建。 - -## 最基本的构建方式 - -SQL 版本最核心的构建过程可以简化为三步: +脚本会依次启动写入和读取两个独立进程,验证数据库中的记忆可以跨进程和不同会话读取。也可以单独运行某个阶段: -```python -sql_url = "mysql+aiomysql://user:password@localhost:3306/trpc_agent_advanced_memory" - -memory_service = AdvancedMemoryService( - AdvancedMemoryServiceConfig( - storage_backend="sql", - sql_url=sql_url, - sql_is_async=True, - memory_ttl_seconds=120, # from M_TTL; omit to disable expiration - ) -) - -runner = Runner( - app_name="advanced-memory-sql-demo", - agent=create_agent(), - session_service=InMemorySessionService(), - memory_service=memory_service, -) +```bash +python run_agent.py --phase write +python run_agent.py --phase read ``` -其中: - -- 用户只需要配置长期 Memory 的 `M_TTL`; -- `AdvancedMemoryService` 只负责长期 memory; -- Session Service 的 SQL 高级压缩接入请看 - [`session_service_with_advanced_memory_sql`](../session_service_with_advanced_memory_sql/); -- 运行请求时通过 `user_id` 和 `session_id` 指定用户及会话; -- 多个节点只要使用相同的 SQL 数据库、`app_name` 和 `user_id`,就能访问同一份长期 memory。 +首次运行时,SQLite 数据库文件和 Advanced Memory 数据表会自动创建。示例中的记忆 TTL 在代码的 `AdvancedMemoryServiceConfig` 中配置为`memory_ttl_seconds=120`。 ## 运行结果(实测) -```text - -==================== WRITE PROCESS ==================== +```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/index', 'index': ''} -🤖 Assistant: I checked my long-term memory index, and it's currently empty — I don't have any saved memories about you yet, so I don't remember your name. - -If you'd like, tell me your name (or anything else you'd like me to remember about you), and I'll save it to my memory so I can remember it across future conversations. - ------ Runner A, query 2 ----- -📝 user: Do you remember my favorite color? +📊 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/index', 'index': ''} -🤖 Assistant: I checked my long-term memory index, and it's currently empty — I don't have any saved memories yet, so I don't remember your favorite color. - -If you tell me your favorite color (or any other preferences you'd like me to keep in mind), I can save it for future conversations! +📊 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. 😊 ------ Runner A, query 3 ----- -📝 user: what is the weather like in paris? -🔧 tool call: get_weather_report({'city': 'Paris'}) -📊 Tool Result: {'status': 'success', 'report': 'The weather in Paris is sunny with a temperature of 25 degrees Celsius.'} -🤖 Assistant: The weather in Paris is **sunny** with a temperature of **25°C**. ☀️ +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. -...... - -==================== READ PROCESS ==================== - ------ Runner B, query 1 ----- -📝 user: Do you remember my name? -🔧 tool call: list_memory_index({}) -🔧 tool call: read_memory({'filename': 'user_identity.md'}) -📊 Tool Result: {'index_path': 'advanced-memory://sql/advanced-memory-sql-demo/sql-demo-user/memory/index', 'index': '- [User identity](user_identity.md):Alice, favorite color blue.\n'} -📊 Tool Result: {'found': True, 'filename': 'user_identity.md', 'content': "---\nname: User identity\ndescription: The user's name and basic identity.\ntype: user\nupdated_at: 2026-09-08T07:26:18.283583+00:00\n---\n---\nname: User identity\ndescription: The user's name and basic identity.\ntype: user\nupdated_at: 2026-09-08T07:26:06.708029+00:00\n---\nThe user's name is Alice. She introduced herself in conversation. Her favorite color is blue.\n", 'updated_at': '2026-09-08T07:26:18.283583+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** — from my long-term memory, that's how you introduced yourself. I also have noted that your favorite color is blue. 😊 - -If any of that has changed, just let me know and I'll update my memory records. +... ----- Runner B, query 2 ----- 📝 user: Do you remember my favorite color? -🔧 tool call: read_memory({'filename': 'user_identity.md'}) -📊 Tool Result: {'found': True, 'filename': 'user_identity.md', 'content': "---\nname: User identity\ndescription: The user's name and basic identity.\ntype: user\nupdated_at: 2026-09-08T07:26:18.283583+00:00\n---\n---\nname: User identity\ndescription: The user's name and basic identity.\ntype: user\nupdated_at: 2026-09-08T07:26:06.708029+00:00\n---\nThe user's name is Alice. She introduced herself in conversation. Her favorite color is blue.\n", 'updated_at': '2026-09-08T07:26:18.283583+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**, Alice. 💙 -``` - -## SQL 表 - -Advanced Memory 使用独立的表,不复用原始 `SqlMemoryService` 的 `mem_events`: - -```text -advanced_memory_indexes -advanced_memory_topics -``` - -Markdown 内容保存在 `TEXT` 字段,`expires_at` 用于 Memory TTL。 -SQL 后端在读取时过滤过期数据,并在访问或写入时刷新同一用户的长期 Memory。 +🔧 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/run_agent.py b/examples/memory_service_with_advanced_memory_sql/run_agent.py index 0fbde7f74..d0484cb6a 100644 --- a/examples/memory_service_with_advanced_memory_sql/run_agent.py +++ b/examples/memory_service_with_advanced_memory_sql/run_agent.py @@ -14,13 +14,13 @@ from dotenv import load_dotenv from agent.agent import create_agent -from trpc_agent_sdk.advanced_memory import AdvancedMemoryServiceConfig +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")) +load_dotenv(Path(__file__).with_name(".env"), override=True) RUNNER_A_QUERIES = [ "Do you remember my name?", @@ -59,14 +59,13 @@ def sql_is_async() -> bool: def create_advanced_memory_service(sql_url: str) -> AdvancedMemoryService: """Create the long-term Advanced Memory service backed by SQL.""" - memory_ttl = os.getenv("M_TTL") config = AdvancedMemoryServiceConfig( storage_backend="sql", sql_url=sql_url, sql_is_async=sql_is_async(), - memory_ttl_seconds=int(memory_ttl) if memory_ttl else None, + memory_ttl_seconds=120, ) - return AdvancedMemoryService(config) + return AdvancedMemoryService(config=config) async def run_phase(phase: str) -> None: diff --git a/examples/session_service_with_advanced_memory_redis/.env b/examples/session_service_with_advanced_memory_redis/.env deleted file mode 100644 index 4858f369a..000000000 --- a/examples/session_service_with_advanced_memory_redis/.env +++ /dev/null @@ -1,10 +0,0 @@ -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 deleted file mode 100644 index 91886f100..000000000 --- a/examples/session_service_with_advanced_memory_redis/README.md +++ /dev/null @@ -1,103 +0,0 @@ -# Redis SessionService + Session Compact - -本示例只演示如何在已有 `RedisSessionService` 上增加: - -- Tool Result Budget -- History Snip -- Microcompact -- AutoCompact -- AutoCompact 触发时生成的 Session Memory - -压缩能力来自 `trpc_agent_sdk.sessions.compact`,长期记忆仍属于独立的 -`trpc_agent_sdk.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 deleted file mode 100644 index bc6e483f9..000000000 --- a/examples/session_service_with_advanced_memory_redis/agent/__init__.py +++ /dev/null @@ -1,5 +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. diff --git a/examples/session_service_with_advanced_memory_redis/agent/agent.py b/examples/session_service_with_advanced_memory_redis/agent/agent.py deleted file mode 100644 index 57093a8c1..000000000 --- a/examples/session_service_with_advanced_memory_redis/agent/agent.py +++ /dev/null @@ -1,39 +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. -"""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 deleted file mode 100644 index 9ff843472..000000000 --- a/examples/session_service_with_advanced_memory_redis/agent/config.py +++ /dev/null @@ -1,19 +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. -"""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 deleted file mode 100644 index 8913fbef9..000000000 --- a/examples/session_service_with_advanced_memory_redis/agent/prompts.py +++ /dev/null @@ -1,10 +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. -"""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 deleted file mode 100644 index 9a109117a..000000000 --- a/examples/session_service_with_advanced_memory_redis/agent/tools.py +++ /dev/null @@ -1,11 +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. -"""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 deleted file mode 100644 index 828eab31e..000000000 --- a/examples/session_service_with_advanced_memory_redis/run_agent.py +++ /dev/null @@ -1,130 +0,0 @@ -#!/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.runners import Runner -from trpc_agent_sdk.sessions import RedisSessionService -from trpc_agent_sdk.sessions import SessionServiceConfig -from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizer -from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerConfig -from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerManager -from trpc_agent_sdk.sessions.compact import AutoCompactSummarizerConfig -from trpc_agent_sdk.sessions.compact import SessionMemoryExtractorConfig -from trpc_agent_sdk.sessions.compact import TokenContextTrackerConfig -from trpc_agent_sdk.types import Content -from trpc_agent_sdk.types import Part - -load_dotenv() - - -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() -> AdvancedAutoCompactSummarizerConfig: - """Configure only the settings needed to demonstrate one compaction.""" - return AdvancedAutoCompactSummarizerConfig( - token_context_tracker=TokenContextTrackerConfig( - model_context_window_tokens=4096, - max_output_tokens=256, - warning_ratio=0.25, - auto_compact_ratio=0.30, - blocking_ratio=0.95, - ), - session_memory=SessionMemoryExtractorConfig( - initial_tokens=500, - update_tokens=500, - ), - auto_compact=AutoCompactSummarizerConfig(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 = AdvancedAutoCompactSummarizerManager( - AdvancedAutoCompactSummarizer(compact_config), - ) - session_config = SessionServiceConfig(store_historical_events=True) - session_service = RedisSessionService( - db_url=redis_url(), - is_async=True, - session_config=session_config, - summarizer_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 deleted file mode 100644 index 693f8ecb1..000000000 --- a/examples/session_service_with_advanced_memory_sql/.env +++ /dev/null @@ -1,11 +0,0 @@ -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 deleted file mode 100644 index 522ce8555..000000000 --- a/examples/session_service_with_advanced_memory_sql/README.md +++ /dev/null @@ -1,100 +0,0 @@ -# SQL SessionService + Session Compact - -本示例只演示如何在已有 `SqlSessionService` 上增加: - -- Tool Result Budget -- History Snip -- Microcompact -- AutoCompact -- AutoCompact 触发时生成的 Session Memory - -压缩能力来自 `trpc_agent_sdk.sessions.compact`,长期记忆仍属于独立的 -`trpc_agent_sdk.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 deleted file mode 100644 index bc6e483f9..000000000 --- a/examples/session_service_with_advanced_memory_sql/agent/__init__.py +++ /dev/null @@ -1,5 +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. diff --git a/examples/session_service_with_advanced_memory_sql/agent/agent.py b/examples/session_service_with_advanced_memory_sql/agent/agent.py deleted file mode 100644 index 5501a1b0b..000000000 --- a/examples/session_service_with_advanced_memory_sql/agent/agent.py +++ /dev/null @@ -1,39 +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. -"""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 deleted file mode 100644 index 91236eaf9..000000000 --- a/examples/session_service_with_advanced_memory_sql/agent/config.py +++ /dev/null @@ -1,19 +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. -"""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 deleted file mode 100644 index d7213fa0e..000000000 --- a/examples/session_service_with_advanced_memory_sql/agent/prompts.py +++ /dev/null @@ -1,10 +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. -"""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 deleted file mode 100644 index cf472e2b6..000000000 --- a/examples/session_service_with_advanced_memory_sql/agent/tools.py +++ /dev/null @@ -1,11 +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. -"""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 deleted file mode 100644 index 5c1bb99b7..000000000 --- a/examples/session_service_with_advanced_memory_sql/run_agent.py +++ /dev/null @@ -1,126 +0,0 @@ -#!/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.runners import Runner -from trpc_agent_sdk.sessions import SessionServiceConfig -from trpc_agent_sdk.sessions import SqlSessionService -from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizer -from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerConfig -from trpc_agent_sdk.sessions.compact import AdvancedAutoCompactSummarizerManager -from trpc_agent_sdk.sessions.compact import AutoCompactSummarizerConfig -from trpc_agent_sdk.sessions.compact import SessionMemoryExtractorConfig -from trpc_agent_sdk.sessions.compact import TokenContextTrackerConfig -from trpc_agent_sdk.types import Content -from trpc_agent_sdk.types import Part - -load_dotenv() - - -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() -> AdvancedAutoCompactSummarizerConfig: - """Configure only the settings needed to demonstrate one compaction.""" - return AdvancedAutoCompactSummarizerConfig( - token_context_tracker=TokenContextTrackerConfig( - model_context_window_tokens=4096, - max_output_tokens=256, - warning_ratio=0.25, - auto_compact_ratio=0.30, - blocking_ratio=0.95, - ), - session_memory=SessionMemoryExtractorConfig( - initial_tokens=500, - update_tokens=500, - ), - auto_compact=AutoCompactSummarizerConfig(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 = AdvancedAutoCompactSummarizerManager( - AdvancedAutoCompactSummarizer(compact_config), - ) - session_config = SessionServiceConfig(store_historical_events=True) - session_service = SqlSessionService( - db_url=sql_url(), - is_async=False, - session_config=session_config, - summarizer_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_tools.py b/tests/advanced_memory/test_advanced_memory_tools.py index 59124377e..dc2e181b2 100644 --- a/tests/advanced_memory/test_advanced_memory_tools.py +++ b/tests/advanced_memory/test_advanced_memory_tools.py @@ -8,9 +8,9 @@ import pytest -from trpc_agent_sdk.advanced_memory import AdvancedMemoryServiceConfig -from trpc_agent_sdk.advanced_memory import AdvancedMemoryPaths -from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime +from trpc_agent_sdk.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 diff --git a/tests/advanced_memory/test_memory_context.py b/tests/advanced_memory/test_memory_context.py index a2385d9d3..4a4b319b7 100644 --- a/tests/advanced_memory/test_memory_context.py +++ b/tests/advanced_memory/test_memory_context.py @@ -7,11 +7,14 @@ import pytest -from trpc_agent_sdk.advanced_memory import AdvancedMemoryServiceConfig -from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime -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.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 @@ -26,6 +29,20 @@ def _runtime(tmp_path: Path) -> AdvancedMemoryRuntime: )) +@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) @@ -66,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="项目约定", diff --git a/tests/advanced_memory/test_preload_memory.py b/tests/advanced_memory/test_preload_memory.py index a2c61deb8..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 AdvancedMemoryServiceConfig -from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime -from trpc_agent_sdk.advanced_memory import MemoryDocument -from trpc_agent_sdk.advanced_memory import MemoryPreloader -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,6 +33,44 @@ 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( diff --git a/trpc_agent_sdk/memory/__init__.py b/trpc_agent_sdk/memory/__init__.py index db9e1c012..e0a69ed5e 100644 --- a/trpc_agent_sdk/memory/__init__.py +++ b/trpc_agent_sdk/memory/__init__.py @@ -7,12 +7,13 @@ 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 from trpc_agent_sdk.abc import MemoryServiceConfig +from ._advanced_memory_service import AdvancedMemoryService from ._in_memory_memory_service import EventTtl from ._in_memory_memory_service import InMemoryMemoryService from ._redis_memory_service import RedisMemoryService @@ -26,6 +27,7 @@ __all__ = [ "BaseMemoryService", "MemoryServiceConfig", + "AdvancedMemoryService", "EventTtl", "InMemoryMemoryService", "RedisMemoryService", @@ -36,4 +38,3 @@ "extract_words_lower", "format_timestamp", ] - diff --git a/trpc_agent_sdk/memory/_advanced_memory_service.py b/trpc_agent_sdk/memory/_advanced_memory_service.py index e69de29bb..a95406f5e 100644 --- a/trpc_agent_sdk/memory/_advanced_memory_service.py +++ b/trpc_agent_sdk/memory/_advanced_memory_service.py @@ -0,0 +1,126 @@ +# 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. +"""Runner-compatible facade for the Advanced Memory mechanism.""" + +from __future__ import annotations + +from typing import Any +from typing import Optional +from typing import TYPE_CHECKING + +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 + +if TYPE_CHECKING: + 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(MemoryServiceABC): + """Expose tool-driven long-term Memory through the Runner memory API. + + ``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: AdvancedMemoryServiceConfig | None = None, + *, + runtime: AdvancedMemoryRuntime | 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.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 AdvancedMemoryServiceConfig()) + super().__init__(MemoryServiceConfig(enabled=resolved_config.enabled)) + self._runtime = runtime or AdvancedMemoryRuntime.create(resolved_config) + self._preload_memory_model = preload_memory_model + self._install_long_term_memory_tools = install_long_term_memory_tools + self._integration: LongTermMemoryIntegration | None = None + self._bound_agent: Any | None = None + + @property + def config(self) -> AdvancedMemoryServiceConfig: + """Return the Advanced Memory configuration.""" + return self._runtime.config + + @property + def runtime(self) -> AdvancedMemoryRuntime: + """Return the Advanced Memory runtime.""" + return self._runtime + + @property + 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 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") + return session_service + + self._integration = setup_long_term_memory( + agent, + self._runtime, + preload_memory_model=self._preload_memory_model, + install_tools=self._install_long_term_memory_tools, + ) + self._bound_agent = agent + return session_service + + @override + async def store_session( + self, + session: SessionABC, + agent_context: Optional[AgentContext] = None, + ) -> None: + """Keep the standard hook side-effect free. + + 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, + query: str, + limit: int = 10, + agent_context: Optional[AgentContext] = None, + ) -> SearchMemoryResponse: + """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 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..9bcc09654 --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_formats.py @@ -0,0 +1,133 @@ +# 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/memory/advanced_memory/_memory_context.py b/trpc_agent_sdk/memory/advanced_memory/_memory_context.py new file mode 100644 index 000000000..b3e62a9df --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_memory_context.py @@ -0,0 +1,126 @@ +# 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. +"""Inject the long-term memory index into model system instructions.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +from trpc_agent_sdk.sessions.compact._callbacks import install_staged_callback +from ._runtime import AdvancedMemoryRuntime + +if TYPE_CHECKING: + from trpc_agent_sdk.agents import LlmAgent + from trpc_agent_sdk.context import InvocationContext + from trpc_agent_sdk.models import LlmRequest + +LONG_TERM_MEMORY_MARKER = "" + + +class LongTermMemoryContext: + """Load a bounded MEMORY.md index for each model request.""" + + def __init__(self, memory_runtime: AdvancedMemoryRuntime) -> None: + """Store the runtime bound to this long-term memory context.""" + self._runtime = memory_runtime + + @property + def runtime(self) -> AdvancedMemoryRuntime: + """Return the runtime bound to this long-term memory context.""" + return self._runtime + + async def apply(self, request: "LlmRequest", ctx: "InvocationContext | None" = None) -> bool: + """Append the MEMORY.md index and on-demand read guidance.""" + 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 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 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 " + "lead, not a complete fact. Use it only when relevant to the current task. For exact details, prefer " + "read_memory on the referenced file; do not infer details from one index line.\n\n" + "Memory records are point-in-time observations and may become stale. Before relying on a memory for " + "current code, configuration, or external state, verify it against the current source or resource. " + "If a memory is incorrect or outdated, update the existing memory instead of creating a duplicate.\n\n" + "## Proactively maintain long-term memory\n" + "If the save_memory tool is available, proactively save information that is sufficiently certain and " + "useful across sessions; do not wait for the user to say \"remember this\". Prefer saving:\n" + "- user: stable identity, role, preferences, skill level, work habits, or explicit personal constraints;\n" + "- feedback: corrections, confirmations, or preferences about collaboration, format, and quality;\n" + "- project: goals, confirmed technical decisions, architecture/process conventions, important state, or " + "deadlines that cannot be reliably inferred from code alone;\n" + "- reference: locations, purposes, and usage constraints for external systems, docs, APIs, repositories, " + "or resources.\n" + "For corrections, replacements, or important additions, update the existing memory with the same " + "filename instead of creating a duplicate. Keep each memory focused on one stable, concrete, actionable " + "topic; inspect the index first and reuse an existing topic when possible.\n\n" + "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 " + 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: " + 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]) + return True + + +class LongTermMemoryContextCallback: + """Adapt the long-term memory index injector to before_model_callback.""" + + advanced_memory_stage = 5 + + def __init__(self, memory_context: LongTermMemoryContext) -> None: + """Store the injector executed before each model request.""" + self._memory_context = memory_context + + @property + def memory_context(self) -> LongTermMemoryContext: + """Return the memory context used by this callback.""" + return self._memory_context + + async def __call__(self, ctx: "InvocationContext", request: "LlmRequest") -> None: + """Inject the long-term memory index before a model request.""" + await self._memory_context.apply(request, ctx) + return None + + +def setup_long_term_memory_context( + agent: "LlmAgent", + memory_runtime: AdvancedMemoryRuntime, +) -> LongTermMemoryContext: + """Install the index callback while preserving pipeline stage order.""" + memory_context = LongTermMemoryContext(memory_runtime) + callback = LongTermMemoryContextCallback(memory_context) + existing_context = install_staged_callback( + agent, + callback, + callback_type=LongTermMemoryContextCallback, + component_attribute="memory_context", + memory_runtime=memory_runtime, + conflict_message="Long-term memory context is already configured with another runtime", + ) + return existing_context or memory_context 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/memory/advanced_memory/_preload_memory.py b/trpc_agent_sdk/memory/advanced_memory/_preload_memory.py new file mode 100644 index 000000000..f2bcfce8b --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_preload_memory.py @@ -0,0 +1,286 @@ +# 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. +"""Select and preload long-term memories for the current model request.""" + +from __future__ import annotations + +import json +import re +from dataclasses import dataclass +from datetime import datetime +from datetime import timezone +from html import escape +from typing import Protocol +from typing import TYPE_CHECKING + +from trpc_agent_sdk.log import logger +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 + +if TYPE_CHECKING: + from trpc_agent_sdk.context import InvocationContext + +_FRONTMATTER_FIELD = re.compile(r"^(?P[A-Za-z_]+):\s*(?P.*)$", re.MULTILINE) + + +@dataclass(frozen=True) +class MemoryCandidate: + """Describe one topic file using only its frontmatter metadata.""" + + filename: str + name: str + description: str + memory_type: str + updated_at: datetime | None + + def to_selector_dict(self) -> dict[str, str]: + """Render the metadata passed to the relevance selector.""" + return { + "filename": self.filename, + "name": self.name, + "description": self.description, + "type": self.memory_type, + "freshness": memory_freshness(self.updated_at), + } + + +class MemoryRelevanceSelector(Protocol): + """Select relevant topic filenames for a user query.""" + + async def select( + self, + query: str, + candidates: list[MemoryCandidate], + ctx: "InvocationContext", + *, + limit: int, + ) -> list[str]: + """Return at most ``limit`` filenames from the candidate list.""" + + +def _frontmatter(content: str) -> dict[str, str]: + """Parse the simple single-line frontmatter used by MemoryDocument.""" + if not content.startswith("---\n"): + return {} + end = content.find("\n---", 4) + if end < 0: + return {} + return {match.group("field"): match.group("value").strip() for match in _FRONTMATTER_FIELD.finditer(content[4:end])} + + +def _candidate_from_content(filename: str, content: str) -> MemoryCandidate: + """Build candidate metadata from a topic document.""" + metadata = _frontmatter(content) + return MemoryCandidate( + filename=filename, + name=metadata.get("name", filename), + description=metadata.get("description", ""), + memory_type=metadata.get("type", ""), + updated_at=parse_memory_updated_at(content), + ) + + +class ModelMemoryRelevanceSelector: + """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 + + 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 + + @staticmethod + def _build_prompt(query: str, candidates: list[MemoryCandidate], limit: int) -> str: + """Build the strict JSON selection prompt.""" + candidate_payload = [candidate.to_selector_dict() for candidate in candidates] + return ("Select the long-term memory files that are clearly relevant to the user's query.\n" + f"Return at most {limit} filenames. If none are clearly relevant, return an empty list.\n" + "Use freshness as one relevance signal, but do not discard an older memory solely because it is old.\n" + "Only return filenames from the candidate list. Do not explain your choices.\n" + 'Return exactly one JSON object: {"selected_memories": ["filename.md"]}\n\n' + f"User query:\n{query}\n\n" + f"Candidate memories:\n{json.dumps(candidate_payload, ensure_ascii=False, indent=2)}") + + @staticmethod + def _parse_selection( + text: str, + candidates: list[MemoryCandidate], + limit: int, + ) -> list[str]: + """Parse and validate the selector's JSON response.""" + payload = None + decoder = json.JSONDecoder() + for index, character in enumerate(text): + if character != "{": + continue + try: + value, _ = decoder.raw_decode(text[index:]) + except json.JSONDecodeError: + continue + if isinstance(value, dict): + payload = value + break + if payload is None: + raise ValueError("Memory relevance selector returned no JSON object") + selected = payload.get("selected_memories") + if not isinstance(selected, list): + raise ValueError("Memory relevance selector returned an invalid selected_memories list") + valid_filenames = {candidate.filename for candidate in candidates} + result: list[str] = [] + for filename in selected: + if isinstance(filename, str) and filename in valid_filenames and filename not in result: + result.append(filename) + if len(result) >= limit: + break + return result + + async def select( + self, + query: str, + candidates: list[MemoryCandidate], + ctx: "InvocationContext", + *, + limit: int, + ) -> list[str]: + """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))], + ) + ], + ) + 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( + query: str, + candidates: list[MemoryCandidate], + ctx: "InvocationContext", + *, + selector: MemoryRelevanceSelector, + limit: int, +) -> list[str]: + """Select relevant memory filenames behind a replaceable screening boundary.""" + selected = await selector.select(query, candidates, ctx, limit=limit) + valid_filenames = {candidate.filename for candidate in candidates} + return list(dict.fromkeys(filename for filename in selected if filename in valid_filenames))[:limit] + + +class MemoryPreloader: + """Find and render relevant topic files for automatic prompt injection.""" + + def __init__( + self, + runtime: AdvancedMemoryRuntime, + selector: MemoryRelevanceSelector | None = None, + ) -> None: + """Store the runtime and replaceable relevance selector.""" + self._runtime = runtime + self._selector = selector or ModelMemoryRelevanceSelector() + + 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 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[: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.""" + config = self._runtime.config + if not config.enabled or not config.preload_memory_enabled or not query.strip(): + return None + try: + candidates = await self._candidates(ctx) + except Exception as exc: # noqa: BLE001 + logger.warning("Advanced Memory preload candidate loading failed: %s", exc) + return None + if not candidates: + return None + try: + selected = await select_relevant_memory_filenames( + query, + candidates, + ctx, + selector=self._selector, + limit=config.preload_memory_max_topics, + ) + except Exception as exc: # noqa: BLE001 + logger.warning("Advanced Memory preload selection failed: %s", exc) + return None + by_filename = {candidate.filename: candidate for candidate in candidates} + sections: list[str] = [] + used_chars = 0 + for filename in selected: + candidate = by_filename.get(filename) + if candidate is None: + continue + try: + 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 + if full_content is None: + continue + remaining = config.preload_memory_max_chars - used_chars + if remaining <= 0: + break + truncated = len(full_content) > remaining + content = full_content[:remaining] + safe_filename = escape(filename, quote=True) + sections.append(f'\n' + f"{content}\n" + "") + used_chars += len(content) + if not sections: + return None + return ( + "\n" + "The following memories were automatically selected for the current request. " + "They are historical observations, not guaranteed current facts. Verify them when necessary " + "and update them if they are outdated or incorrect. Each memory includes its source filename " + "and has already been read for this request. You may read or update these files again when needed.\n\n" + + "\n\n".join(sections) + "\n") diff --git a/trpc_agent_sdk/memory/advanced_memory/_redis_stores.py b/trpc_agent_sdk/memory/advanced_memory/_redis_stores.py new file mode 100644 index 000000000..8e0dba64a --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_redis_stores.py @@ -0,0 +1,197 @@ +"""Redis implementations of the Advanced Memory storage contracts.""" + +from __future__ import annotations + +import asyncio +from contextlib import asynccontextmanager +from dataclasses import replace +from datetime import datetime, timezone +from pathlib import Path +from typing import Any +from uuid import uuid4 + +from trpc_agent_sdk.storage import RedisCommand, RedisExpire, RedisStorage +from trpc_agent_sdk.types import Ttl + +from ._config import AdvancedMemoryServiceConfig +from ._formats import MemoryDocument, MemoryIndexEntry +from ._formats import limit_memory_index +from ._paths import AdvancedMemoryPaths +from ._storage import parse_memory_index, prune_memory_index + +_RELEASE_LOCK_SCRIPT = """ +if redis.call('GET', KEYS[1]) == ARGV[1] then + return redis.call('DEL', KEYS[1]) +end +return 0 +""" + + +class _RedisStore: + + def __init__( + self, + config: AdvancedMemoryServiceConfig, + paths: AdvancedMemoryPaths, + storage: RedisStorage, + ) -> None: + if paths.scope is None: + raise ValueError("Redis Advanced Memory storage requires a tenant scope") + self._config, self._paths, self._storage = config, paths, storage + app_component = paths.tenant_root_dir.parent.name + user_component = paths.tenant_root_dir.name + self._user_base = f"{config.redis_key_prefix}:{{{app_component}:{user_component}}}" + + async def _command(self, method: str, *args: Any, **kwargs: Any) -> Any: + command_expire = kwargs.pop("_command_expire", None) + async with self._storage.create_db_session() as connection: + return await self._storage.execute_command( + connection, + RedisCommand(method=method, args=args, kwargs=kwargs, expire=command_expire or RedisExpire()), + ) + + def _memory_registry(self) -> str: + return f"{self._user_base}:memory:keys" + + def _memory_lock_key(self) -> str: + """Return the distributed lock key for this app/user memory scope.""" + return f"{self._user_base}:memory:lock" + + @asynccontextmanager + async def _memory_write_lock(self): + """Serialize long-term memory writes across processes and nodes.""" + token = uuid4().hex + key = self._memory_lock_key() + deadline = asyncio.get_running_loop().time() + self._config.memory_lock_acquire_timeout_seconds + acquired = False + while asyncio.get_running_loop().time() < deadline: + result = await self._command( + "set", + key, + token, + nx=True, + ex=self._config.memory_lock_ttl_seconds, + _command_expire=RedisExpire( + key=key, + ttl=Ttl(ttl_seconds=self._config.memory_lock_ttl_seconds), + ), + ) + if result is True or result in (b"OK", "OK"): + acquired = True + break + await asyncio.sleep(min(0.05, max(0.0, deadline - asyncio.get_running_loop().time()))) + if not acquired: + raise TimeoutError(f"Timed out acquiring Advanced Memory lock for {self._paths.scope.storage_key}") + try: + yield + finally: + await self._command( + "eval", + _RELEASE_LOCK_SCRIPT, + 1, + key, + token, + ) + + async def _refresh_ttl_group( + self, + registry: str, + keys: list[str], + ttl: int | None, + skip_prefixes: tuple[str, ...] = (), + ) -> None: + """Track and refresh every key in one logical memory group.""" + if ttl is None: + return + if keys: + await self._command("sadd", registry, *keys) + tracked = await self._command("smembers", registry) or [] + tracked_keys = {self._text(value) for value in tracked} + tracked_keys.update(keys) + for key in tracked_keys: + if key and not key.startswith(skip_prefixes): + await self._command("expire", key, ttl) + await self._command("expire", registry, ttl) + + async def _refresh_memory_ttl(self, *keys: str) -> None: + await self._refresh_ttl_group( + self._memory_registry(), + list(keys), + self._config.memory_ttl_seconds, + ) + + @staticmethod + def _text(value: Any) -> str | None: + if value is None: + return None + return value.decode("utf-8") if isinstance(value, bytes) else str(value) + + +class RedisLongTermMemoryStore(_RedisStore): + + async def initialize(self) -> None: + key = f"{self._user_base}:memory:index" + await self._command("setnx", key, "") + await self._refresh_memory_ttl(key) + + async def read_index(self) -> str: + key = f"{self._user_base}:memory:index" + value = self._text(await self._command("get", key)) or "" + await self._refresh_memory_ttl() + index = value + valid_filenames = set() + for entry in parse_memory_index(index): + topic_key = f"{self._user_base}:memory:topic:{self._topic_name(entry.filename)}" + if await self._command("exists", topic_key): + valid_filenames.add(entry.filename) + pruned_index = prune_memory_index(index, valid_filenames) + if pruned_index != index: + async with self._memory_write_lock(): + await self._command("set", key, pruned_index) + await self._refresh_memory_ttl(key) + return limit_memory_index( + pruned_index, + max_lines=self._config.memory_index_max_lines, + max_bytes=self._config.memory_index_max_bytes, + encoding=self._config.encoding, + ) + + async def write_index(self, entries: list[MemoryIndexEntry]) -> None: + content = "\n".join(entry.to_markdown() for entry in entries) + key = f"{self._user_base}:memory:index" + async with self._memory_write_lock(): + await self._command("set", key, f"{content}\n" if content else "") + await self._refresh_memory_ttl(key) + + def _topic_name(self, topic_name: str) -> str: + return self._paths.memory_topic_path(topic_name).name + + async def read_topic(self, topic_name: str) -> str | None: + key = f"{self._user_base}:memory:topic:{self._topic_name(topic_name)}" + value = await self._command("get", key) + await self._refresh_memory_ttl() + return self._text(value) + + async def read_topic_frontmatter(self, topic_name: str) -> str | None: + content = await self.read_topic(topic_name) + if content is None: + return None + end = content.find("\n---", 4) if content.startswith("---\n") else -1 + return content[:end + 4] if end >= 0 else content + + async def write_topic(self, topic_name: str, document: MemoryDocument) -> Path: + name = self._topic_name(topic_name) + document = replace(document, updated_at=datetime.now(timezone.utc)) + topic_key = f"{self._user_base}:memory:topic:{name}" + topics_key = f"{self._user_base}:memory:topics" + async with self._memory_write_lock(): + await self._command("set", topic_key, document.to_markdown()) + await self._command("zadd", topics_key, {name: document.updated_at.timestamp()}) + await self._refresh_memory_ttl(topic_key, topics_key) + return Path(name) + + async def list_topics(self) -> list[Path]: + key = f"{self._user_base}:memory:topics" + values = await self._command("zrange", key, 0, -1) + await self._refresh_memory_ttl() + return [Path(self._text(value) or "") for value in values] diff --git a/trpc_agent_sdk/memory/advanced_memory/_runtime.py b/trpc_agent_sdk/memory/advanced_memory/_runtime.py new file mode 100644 index 000000000..c3eb60d1e --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_runtime.py @@ -0,0 +1,202 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Unified runtime entry point for the independent memory mechanism.""" + +from __future__ import annotations + +from dataclasses import dataclass +from dataclasses import field +import shutil +import threading +from typing import Any + +from ._config import AdvancedMemoryServiceConfig +from trpc_agent_sdk.sessions.compact.advanced._coordination import CrossLoopLock +from ._paths import AdvancedMemoryPaths +from ._paths import MemoryScope +from ._storage import LocalAdvancedMemoryCleanup +from ._storage import LongTermMemoryStore + + +@dataclass(frozen=True) +class AdvancedMemoryRuntime: + """Aggregate configuration, paths, and long-term memory storage.""" + + config: AdvancedMemoryServiceConfig + paths: AdvancedMemoryPaths + long_term_memory: LongTermMemoryStore + _scoped_runtimes: dict[MemoryScope, "ScopedAdvancedMemoryRuntime"] = field( + default_factory=dict, + repr=False, + compare=False, + ) + _scoped_runtimes_lock: threading.Lock = field( + default_factory=threading.Lock, + repr=False, + compare=False, + ) + _redis_storage: Any | None = field(default=None, repr=False, compare=False) + _sql_storage: Any | None = field(default=None, repr=False, compare=False) + _sql_cleanup: Any | None = field(default=None, repr=False, compare=False) + _local_cleanup: LocalAdvancedMemoryCleanup | None = field(default=None, repr=False, compare=False) + _close_lock: CrossLoopLock = field( + default_factory=CrossLoopLock, + repr=False, + compare=False, + ) + _closed: bool = field(default=False, repr=False, compare=False) + + @classmethod + def create(cls, config: AdvancedMemoryServiceConfig | None = None) -> "AdvancedMemoryRuntime": + """Create a runtime isolated from the legacy mechanism.""" + resolved_config = config or AdvancedMemoryServiceConfig() + paths = AdvancedMemoryPaths(resolved_config) + redis_storage = None + sql_storage = None + sql_cleanup = None + local_cleanup = None + if resolved_config.storage_backend == "redis": + from trpc_agent_sdk.storage import RedisStorage + redis_storage = RedisStorage(redis_url=resolved_config.redis_url, is_async=resolved_config.redis_is_async) + elif resolved_config.storage_backend == "sql": + from trpc_agent_sdk.storage import SqlStorage + from ._sql_stores import AdvancedMemorySqlBase, SqlAdvancedMemoryCleanup + sql_storage = SqlStorage( + is_async=resolved_config.sql_is_async, + db_url=resolved_config.sql_url, + metadata=AdvancedMemorySqlBase.metadata, + expire_on_commit=False, + ) + sql_cleanup = SqlAdvancedMemoryCleanup(resolved_config, sql_storage) + else: + local_cleanup = LocalAdvancedMemoryCleanup(resolved_config) + return cls( + config=resolved_config, + paths=paths, + long_term_memory=LongTermMemoryStore(resolved_config, paths), + _redis_storage=redis_storage, + _sql_storage=sql_storage, + _sql_cleanup=sql_cleanup, + _local_cleanup=local_cleanup, + ) + + def for_scope(self, app_name: str, user_id: str) -> "ScopedAdvancedMemoryRuntime": + """Return the stores isolated to one application user.""" + scope = MemoryScope(app_name, user_id) + with self._scoped_runtimes_lock: + runtime = self._scoped_runtimes.get(scope) + if runtime is None: + paths = self.paths.for_scope(app_name, user_id) + if self.config.storage_backend == "redis": + from trpc_agent_sdk.storage import RedisStorage + from ._redis_stores import RedisLongTermMemoryStore + + storage = self._redis_storage or RedisStorage( + redis_url=self.config.redis_url, + is_async=self.config.redis_is_async, + ) + long_term_memory = RedisLongTermMemoryStore(self.config, paths, storage) + elif self.config.storage_backend == "sql": + from ._sql_stores import SqlLongTermMemoryStore + storage = self._sql_storage + if storage is None: + raise RuntimeError("SQL Advanced Memory storage is not initialized") + long_term_memory = SqlLongTermMemoryStore(self.config, paths, storage) + else: + long_term_memory = LongTermMemoryStore(self.config, paths) + runtime = ScopedAdvancedMemoryRuntime( + root=self, + scope=scope, + paths=paths, + long_term_memory=long_term_memory, + ) + self._scoped_runtimes[scope] = runtime + return runtime + + def for_session(self, session: object) -> "ScopedAdvancedMemoryRuntime": + """Return the scoped runtime for a SessionABC-compatible object.""" + app_name = getattr(session, "app_name", None) + user_id = getattr(session, "user_id", None) + if not isinstance(app_name, str) or not isinstance(user_id, str): + raise ValueError("Advanced Memory requires session app_name and user_id") + return self.for_scope(app_name, user_id) + + def migrate_legacy(self, app_name: str, user_id: str) -> "ScopedAdvancedMemoryRuntime": + """Move an old flat Advanced Memory layout into one explicit tenant. + + Refuses to overwrite a tenant that already contains data. + """ + scoped = self.for_scope(app_name, user_id) + legacy_paths = self.paths + target_root = scoped.paths.tenant_root_dir + if target_root.exists(): + raise FileExistsError(f"Target Advanced Memory tenant already exists: {target_root}") + if not legacy_paths.memory_dir.exists(): + raise FileNotFoundError("No legacy Advanced Memory directories exist") + target_root.mkdir(parents=True) + if legacy_paths.memory_dir.exists(): + shutil.move(str(legacy_paths.memory_dir), str(scoped.paths.memory_dir)) + return scoped + + async def initialize(self) -> bool: + """Create memory directories only when the mechanism is enabled.""" + if not self.config.enabled: + return False + if self.config.storage_backend == "sql": + if self._sql_storage is None: + raise RuntimeError("SQL Advanced Memory storage is not initialized") + if self._sql_cleanup is not None: + await self._sql_cleanup.start() + async with self._sql_storage.create_db_session(): + pass + return True + if self.config.storage_backend == "redis": + return True + if self._local_cleanup is not None: + await self._local_cleanup.start() + await self.long_term_memory.initialize() + return True + + async def close(self) -> None: + """Release shared external backend resources.""" + async with self._close_lock: + if self._closed: + return + if self._local_cleanup is not None: + await self._local_cleanup.close() + if self._redis_storage is not None: + await self._redis_storage.close() + if self._sql_cleanup is not None: + await self._sql_cleanup.close() + if self._sql_storage is not None: + await self._sql_storage.close() + object.__setattr__(self, "_closed", True) + + +@dataclass(frozen=True) +class ScopedAdvancedMemoryRuntime: + """A tenant-bound view of an :class:`AdvancedMemoryRuntime`.""" + + root: AdvancedMemoryRuntime + scope: MemoryScope + paths: AdvancedMemoryPaths + long_term_memory: LongTermMemoryStore + + @property + def config(self) -> AdvancedMemoryServiceConfig: + """Return the root runtime configuration.""" + return self.root.config + + async def initialize(self) -> bool: + """Initialize only this tenant's local directories.""" + if not self.config.enabled: + return False + if self.config.storage_backend == "local" and self.root._local_cleanup is not None: + await self.root._local_cleanup.start() + if self.config.storage_backend == "sql" and self.root._sql_cleanup is not None: + await self.root._sql_cleanup.start() + await self.long_term_memory.initialize() + return True diff --git a/trpc_agent_sdk/memory/advanced_memory/_sql_stores.py b/trpc_agent_sdk/memory/advanced_memory/_sql_stores.py new file mode 100644 index 000000000..05874e849 --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_sql_stores.py @@ -0,0 +1,315 @@ +"""SQL implementations of the Advanced Memory storage contracts.""" + +from __future__ import annotations + +import asyncio +from datetime import datetime, timedelta, timezone +from dataclasses import replace +from pathlib import Path +from typing import Any + +from sqlalchemy import DateTime, String, Text, func +from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column + +from trpc_agent_sdk.storage import ( + DEFAULT_MAX_KEY_LENGTH, + DEFAULT_MAX_VARCHAR_LENGTH, + PreciseTimestamp, + SqlCondition, + SqlKey, + SqlStorage, +) + +from ._config import AdvancedMemoryServiceConfig +from ._formats import MemoryDocument, MemoryIndexEntry +from ._formats import limit_memory_index +from ._paths import AdvancedMemoryPaths +from ._storage import prune_memory_index + + +class AdvancedMemorySqlBase(DeclarativeBase): + """Metadata owned exclusively by Advanced Memory SQL stores.""" + + +class SqlMemoryIndex(AdvancedMemorySqlBase): + __tablename__ = "advanced_memory_indexes" + + app_name: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + user_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + content: Mapped[str] = mapped_column(Text, default="") + updated_at: Mapped[datetime] = mapped_column(PreciseTimestamp, default=func.now(), onupdate=func.now()) + expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) + + +class SqlMemoryTopic(AdvancedMemorySqlBase): + __tablename__ = "advanced_memory_topics" + + app_name: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + user_id: Mapped[str] = mapped_column(String(DEFAULT_MAX_KEY_LENGTH), primary_key=True) + topic_name: Mapped[str] = mapped_column(String(DEFAULT_MAX_VARCHAR_LENGTH), primary_key=True) + content: Mapped[str] = mapped_column(Text) + updated_at: Mapped[datetime] = mapped_column(PreciseTimestamp, default=func.now(), onupdate=func.now()) + expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) + + +class _SqlStore: + + def __init__( + self, + config: AdvancedMemoryServiceConfig, + paths: AdvancedMemoryPaths, + storage: SqlStorage, + ) -> None: + if paths.scope is None: + raise ValueError("SQL Advanced Memory storage requires a tenant scope") + self._config = config + self._paths = paths + self._storage = storage + self._app_name = paths.scope.app_name + self._user_id = paths.scope.user_id + + @staticmethod + def _now() -> datetime: + return datetime.now(timezone.utc).replace(tzinfo=None) + + def _expiry(self, ttl: int | None) -> datetime | None: + return self._now() + timedelta(seconds=ttl) if ttl is not None else None + + @staticmethod + def _expired(value: datetime | None) -> bool: + if value is None: + return False + return value.replace(tzinfo=None) <= datetime.now(timezone.utc).replace(tzinfo=None) + + async def initialize(self) -> None: + async with self._storage.create_db_session(): + pass + + async def _refresh_memory_scope(self, db: Any) -> None: + expiry = self._expiry(self._config.memory_ttl_seconds) + if expiry is None: + return + index = await self._storage.get(db, SqlKey( + key=(self._app_name, self._user_id), + storage_cls=SqlMemoryIndex, + )) + if index is not None: + index.expires_at = expiry + topics = await self._storage.query( + db, + SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryTopic), + SqlCondition(filters=[ + SqlMemoryTopic.app_name == self._app_name, + SqlMemoryTopic.user_id == self._user_id, + SqlMemoryTopic.expires_at.is_(None) | (SqlMemoryTopic.expires_at > self._now()), + ]), + ) + for topic in topics: + topic.expires_at = expiry + + +class SqlLongTermMemoryStore(_SqlStore): + + async def initialize(self) -> None: + await super().initialize() + async with self._storage.create_db_session() as db: + key = SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryIndex) + row = await self._storage.get(db, key) + if row is None: + await self._storage.add( + db, + SqlMemoryIndex( + app_name=self._app_name, + user_id=self._user_id, + content="", + expires_at=self._expiry(self._config.memory_ttl_seconds), + )) + await self._storage.commit(db) + + async def read_index(self) -> str: + async with self._storage.create_db_session() as db: + row = await self._storage.get(db, SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryIndex)) + if row is None or self._expired(row.expires_at): + return "" + await self._refresh_memory_scope(db) + content = row.content + valid_topics = await self._storage.query( + db, + SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryTopic), + SqlCondition(filters=[ + SqlMemoryTopic.app_name == self._app_name, + SqlMemoryTopic.user_id == self._user_id, + SqlMemoryTopic.expires_at.is_(None) | (SqlMemoryTopic.expires_at > self._now()), + ]), + ) + valid_filenames = {topic.topic_name for topic in valid_topics} + pruned_content = prune_memory_index(content, valid_filenames) + if pruned_content != content: + row.content = pruned_content + content = pruned_content + await self._storage.commit(db) + return limit_memory_index( + content, + max_lines=self._config.memory_index_max_lines, + max_bytes=self._config.memory_index_max_bytes, + encoding=self._config.encoding, + ) + + async def write_index(self, entries: list[MemoryIndexEntry]) -> None: + content = "\n".join(entry.to_markdown() for entry in entries) + if content: + content += "\n" + async with self._storage.create_db_session() as db: + # Keep the tenant's lock row locked until this transaction commits. + await self._storage.get_for_update( + db, + SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryIndex), + ) + key = SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryIndex) + row = await self._storage.get(db, key) + if row is None: + row = SqlMemoryIndex(app_name=self._app_name, user_id=self._user_id) + await self._storage.add(db, row) + row.content = content + row.expires_at = self._expiry(self._config.memory_ttl_seconds) + await self._refresh_memory_scope(db) + await self._storage.commit(db) + + def _topic_key(self, topic_name: str) -> tuple[str, str, str]: + return self._app_name, self._user_id, self._paths.memory_topic_path(topic_name).name + + async def read_topic(self, topic_name: str) -> str | None: + async with self._storage.create_db_session() as db: + row = await self._storage.get(db, SqlKey(key=self._topic_key(topic_name), storage_cls=SqlMemoryTopic)) + if row is None or self._expired(row.expires_at): + return None + await self._refresh_memory_scope(db) + await self._storage.commit(db) + return row.content + + async def read_topic_frontmatter(self, topic_name: str) -> str | None: + content = await self.read_topic(topic_name) + if content is None: + return None + end = content.find("\n---", 4) if content.startswith("---\n") else -1 + return content[:end + 4] if end >= 0 else content + + async def write_topic(self, topic_name: str, document: MemoryDocument) -> Path: + name = self._paths.memory_topic_path(topic_name).name + async with self._storage.create_db_session() as db: + # Serialize all long-term writes for this app/user scope. + await self._storage.get_for_update( + db, + SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryIndex), + ) + key = self._topic_key(name) + row = await self._storage.get(db, SqlKey(key=key, storage_cls=SqlMemoryTopic)) + if row is None: + row = SqlMemoryTopic(app_name=key[0], user_id=key[1], topic_name=key[2]) + await self._storage.add(db, row) + row.content = replace(document, updated_at=self._now().replace(tzinfo=timezone.utc)).to_markdown() + row.updated_at = self._now() + row.expires_at = self._expiry(self._config.memory_ttl_seconds) + await self._refresh_memory_scope(db) + await self._storage.commit(db) + return Path(name) + + async def list_topics(self) -> list[Path]: + async with self._storage.create_db_session() as db: + rows = await self._storage.query( + db, + SqlKey(key=(self._app_name, self._user_id), storage_cls=SqlMemoryTopic), + SqlCondition(filters=[ + SqlMemoryTopic.app_name == self._app_name, + SqlMemoryTopic.user_id == self._user_id, + ]), + ) + rows = [row for row in rows if not self._expired(row.expires_at)] + await self._refresh_memory_scope(db) + await self._storage.commit(db) + return [Path(row.topic_name) for row in sorted(rows, key=lambda item: item.topic_name)] + + +class SqlAdvancedMemoryCleanup: + """Periodically remove expired Advanced Memory SQL rows.""" + + _models = ( + SqlMemoryIndex, + SqlMemoryTopic, + ) + + def __init__(self, config: AdvancedMemoryServiceConfig, storage: SqlStorage) -> None: + self._config = config + self._storage = storage + self._task: asyncio.Task[None] | None = None + self._stop_event: asyncio.Event | None = None + + async def start(self) -> None: + if self._task is not None or self._config.memory_ttl_seconds is None: + return + self._stop_event = asyncio.Event() + self._task = asyncio.create_task(self._run()) + + async def cleanup_once(self) -> None: + now = datetime.now(timezone.utc).replace(tzinfo=None) + async with self._storage.create_db_session() as db: + for model in self._models: + await self._storage.delete( + db, + SqlKey(key=tuple(), storage_cls=model), + SqlCondition(filters=[model.expires_at.is_not(None), model.expires_at <= now]), + ) + indexes = await self._storage.query( + db, + SqlKey(key=tuple(), storage_cls=SqlMemoryIndex), + ) + for index in indexes: + topics = await self._storage.query( + db, + SqlKey( + key=(index.app_name, index.user_id), + storage_cls=SqlMemoryTopic, + ), + SqlCondition(filters=[ + SqlMemoryTopic.app_name == index.app_name, + SqlMemoryTopic.user_id == index.user_id, + SqlMemoryTopic.expires_at.is_(None) | (SqlMemoryTopic.expires_at > now), + ]), + ) + valid_filenames = {topic.topic_name for topic in topics} + index.content = prune_memory_index(index.content, valid_filenames) + await self._storage.commit(db) + + async def _run(self) -> None: + if self._stop_event is None: + return + try: + while not self._stop_event.is_set(): + try: + await asyncio.wait_for( + self._stop_event.wait(), + timeout=self._config.sql_cleanup_interval_seconds, + ) + except asyncio.TimeoutError: + await self.cleanup_once() + except asyncio.CancelledError: + raise + + async def close(self) -> None: + if self._stop_event is not None: + self._stop_event.set() + if self._task is not None and not self._task.done(): + self._task.cancel() + try: + await self._task + except asyncio.CancelledError: + pass + self._task = None + self._stop_event = None + + +__all__ = [ + "AdvancedMemorySqlBase", + "SqlAdvancedMemoryCleanup", + "SqlLongTermMemoryStore", +] diff --git a/trpc_agent_sdk/memory/advanced_memory/_storage.py b/trpc_agent_sdk/memory/advanced_memory/_storage.py new file mode 100644 index 000000000..d5a0e9f08 --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_storage.py @@ -0,0 +1,221 @@ +# Tencent is pleased to support the open source ecosystem. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# Licensed under Apache-2.0. +"""Long-term memory storage owned by AdvancedMemoryService.""" + +from __future__ import annotations + +import asyncio +import os +import re +import tempfile +import time +from dataclasses import replace +from datetime import datetime +from datetime import timezone +from pathlib import Path + +from ._formats import MemoryDocument +from ._formats import MemoryIndexEntry +from ._formats import limit_memory_index + +from ._config import AdvancedMemoryServiceConfig +from ._paths import AdvancedMemoryPaths + +_MEMORY_INDEX_PATTERN = re.compile(r"^- \[(?P.+?)\]((?P.+?)):(?P.+)$") + + +def parse_memory_index(index: str) -> list[MemoryIndexEntry]: + """Parse standard entries from a MEMORY.md index.""" + entries: list[MemoryIndexEntry] = [] + for line in index.splitlines(): + match = _MEMORY_INDEX_PATTERN.match(line.strip()) + if match is not None: + entries.append(MemoryIndexEntry(**match.groupdict())) + return entries + + +def prune_memory_index(index: str, valid_filenames: set[str]) -> str: + """Remove index entries whose topic files no longer exist.""" + lines = [ + line for line in index.splitlines() + if (match := _MEMORY_INDEX_PATTERN.match(line.strip())) is None or match.group("filename") in valid_filenames + ] + if lines == index.splitlines(): + return index + return "\n".join(lines) + ("\n" if lines else "") + + +def _atomic_write_text(path: Path, content: str, *, encoding: str) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + descriptor, temporary_name = tempfile.mkstemp(prefix=f".{path.name}.", dir=path.parent) + try: + with os.fdopen(descriptor, "w", encoding=encoding) as output: + output.write(content) + output.flush() + os.fsync(output.fileno()) + os.replace(temporary_name, path) + except BaseException: + try: + os.unlink(temporary_name) + except FileNotFoundError: + pass + raise + + +def _is_expired(path: Path, ttl: int | None) -> bool: + return ttl is not None and path.exists() and time.time() - path.stat().st_mtime >= ttl + + +class LongTermMemoryStore: + """Read and write MEMORY.md and its topic files.""" + + def __init__( + self, + config: AdvancedMemoryServiceConfig, + paths: AdvancedMemoryPaths | None = None, + ) -> None: + self._config = config + self._paths = paths or AdvancedMemoryPaths(config) + + @property + def index_path(self) -> Path: + return self._paths.memory_index_path + + async def initialize(self) -> None: + await asyncio.to_thread(self._initialize_sync) + + def _initialize_sync(self) -> None: + self._paths.ensure_base_directories() + if not self.index_path.exists(): + _atomic_write_text(self.index_path, "", encoding=self._config.encoding) + + async def read_index(self) -> str: + return await asyncio.to_thread(self._read_index_sync) + + def _read_index_sync(self) -> str: + if _is_expired(self.index_path, self._config.memory_ttl_seconds): + for path in self._paths.memory_dir.glob("*.md"): + path.unlink(missing_ok=True) + return "" + if not self.index_path.exists(): + return "" + with self.index_path.open(encoding=self._config.encoding) as source: + index = source.read() + valid_filenames = { + path.name + for path in self._paths.memory_dir.glob("*.md") + if path.name != self._config.memory_index_name and not _is_expired(path, self._config.memory_ttl_seconds) + } + pruned_index = prune_memory_index(index, valid_filenames) + if pruned_index != index: + _atomic_write_text( + self.index_path, + pruned_index, + encoding=self._config.encoding, + ) + return limit_memory_index( + pruned_index, + max_lines=self._config.memory_index_max_lines, + max_bytes=self._config.memory_index_max_bytes, + encoding=self._config.encoding, + ) + + async def write_index(self, entries: list[MemoryIndexEntry]) -> None: + content = "\n".join(entry.to_markdown() for entry in entries) + await asyncio.to_thread( + _atomic_write_text, + self.index_path, + f"{content}\n" if content else "", + encoding=self._config.encoding, + ) + + async def read_topic(self, topic_name: str) -> str | None: + path = self._paths.memory_topic_path(topic_name) + return await asyncio.to_thread(lambda: path.read_text(encoding=self._config.encoding) + if path.exists() else None) + + async def read_topic_frontmatter(self, topic_name: str) -> str | None: + content = await self.read_topic(topic_name) + if content is None: + return None + lines: list[str] = [] + for line in content.splitlines(keepends=True): + lines.append(line) + if len(lines) > 1 and line.rstrip("\r\n") == "---": + break + return "".join(lines) + + async def write_topic(self, topic_name: str, document: MemoryDocument) -> Path: + path = self._paths.memory_topic_path(topic_name) + updated = replace(document, updated_at=datetime.now(timezone.utc)) + await asyncio.to_thread( + _atomic_write_text, + path, + updated.to_markdown(), + encoding=self._config.encoding, + ) + return path + + async def list_topics(self) -> list[Path]: + return await asyncio.to_thread(lambda: sorted(path for path in self._paths.memory_dir.glob("*.md") + if path.name != self._config.memory_index_name)) + + +class LocalAdvancedMemoryCleanup: + """Remove expired long-term memory files for the local backend.""" + + def __init__(self, config: AdvancedMemoryServiceConfig) -> None: + self._config = config + self._task: asyncio.Task[None] | None = None + self._stop_event: asyncio.Event | None = None + + async def start(self) -> None: + if self._task is not None or self._config.memory_ttl_seconds is None: + return + self._stop_event = asyncio.Event() + await self.cleanup_once() + self._task = asyncio.create_task(self._run()) + + async def cleanup_once(self) -> None: + await asyncio.to_thread(self._cleanup_sync) + + def _cleanup_sync(self) -> None: + root = self._config.root_dir + memory_dirs = [root / self._config.memory_dir_name] + tenants_root = root / "tenants" + if tenants_root.exists(): + for app_dir in tenants_root.iterdir(): + if app_dir.is_dir(): + memory_dirs.extend(user_dir / self._config.memory_dir_name for user_dir in app_dir.iterdir() + if user_dir.is_dir()) + for memory_dir in memory_dirs: + index_path = memory_dir / self._config.memory_index_name + if _is_expired(index_path, self._config.memory_ttl_seconds): + for path in memory_dir.glob("*.md"): + path.unlink(missing_ok=True) + + async def _run(self) -> None: + if self._stop_event is None: + return + try: + while not self._stop_event.is_set(): + try: + await asyncio.wait_for( + self._stop_event.wait(), + timeout=self._config.memory_ttl_seconds or 60, + ) + except asyncio.TimeoutError: + await self.cleanup_once() + except asyncio.CancelledError: + raise + + async def close(self) -> None: + if self._stop_event is not None: + self._stop_event.set() + if self._task is not None and not self._task.done(): + self._task.cancel() + await asyncio.gather(self._task, return_exceptions=True) + self._task = None + self._stop_event = None diff --git a/trpc_agent_sdk/runners.py b/trpc_agent_sdk/runners.py index 21afe7dae..16fa9a353 100644 --- a/trpc_agent_sdk/runners.py +++ b/trpc_agent_sdk/runners.py @@ -223,6 +223,11 @@ def __init__( the memory service. Set to False when the service is managed outside the runner. """ + if memory_service is not None: + from trpc_agent_sdk.memory import AdvancedMemoryService + + if isinstance(memory_service, AdvancedMemoryService): + session_service = memory_service.bind(agent, session_service) self.app_name = app_name self.agent = agent self.artifact_service = artifact_service diff --git a/trpc_agent_sdk/sessions/compact/_callbacks.py b/trpc_agent_sdk/sessions/compact/_callbacks.py new file mode 100644 index 000000000..44249496a --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/_callbacks.py @@ -0,0 +1,49 @@ +# Tencent is pleased to support the open source community by making +# contributions to the open source ecosystem. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Shared callback installation and stage ordering for Advanced Memory.""" + +from __future__ import annotations + +from typing import Any + + +def install_staged_callback( + agent: Any, + callback: Any, + *, + callback_type: type, + component_attribute: str, + memory_runtime: Any, + conflict_message: str, +) -> Any | None: + """Install a staged callback idempotently and validate runtime ownership.""" + existing = agent.before_model_callback + callbacks = existing if isinstance(existing, list) else ([existing] if existing else []) + for item in callbacks: + if not isinstance(item, callback_type): + continue + component = getattr(item, component_attribute) + if component.runtime is not memory_runtime: + raise ValueError(conflict_message) + return component + stage = getattr(callback, "advanced_memory_stage", None) + if not isinstance(stage, int): + raise TypeError("advanced_memory_stage must be an integer") + + def get_stage(item: Any) -> int: + item_stage = getattr(item, "advanced_memory_stage", 0) + return item_stage if isinstance(item_stage, int) else 0 + + insertion_index = next( + (index for index, item in enumerate(callbacks) if get_stage(item) > stage), + len(callbacks), + ) + agent.before_model_callback = [ + *callbacks[:insertion_index], + callback, + *callbacks[insertion_index:], + ] + return None diff --git a/trpc_agent_sdk/sessions/compact/advanced/_compaction_memory_extractor.py b/trpc_agent_sdk/sessions/compact/advanced/_compaction_memory_extractor.py index 349c6ce0e..1a9668336 100644 --- a/trpc_agent_sdk/sessions/compact/advanced/_compaction_memory_extractor.py +++ b/trpc_agent_sdk/sessions/compact/advanced/_compaction_memory_extractor.py @@ -473,6 +473,13 @@ def _event_has_tool_call(self, record: dict[str, Any]) -> bool: parts = record.get("event", {}).get("content", {}).get("parts", []) return any(isinstance(part, dict) and (part.get("function_call") or part.get("functionCall")) for part in parts) + def _event_has_tool_response(self, record: dict[str, Any]) -> bool: + """Return whether one Session Event contains a function response.""" + parts = record.get("event", {}).get("content", {}).get("parts", []) + return any( + isinstance(part, dict) and (part.get("function_response") or part.get("functionResponse")) + for part in parts) + def _fits_prompt_budget( self, extraction_input: SessionMemoryExtractionInput, @@ -716,9 +723,9 @@ async def extract_if_needed( (config.initial_tokens if token_mode else (config.update_chars if checkpoint_event_id is not None else config.initial_chars))) tool_calls = self._count_tool_calls(pending) - natural_break = not self._last_event_has_tool_call(pending) - if not natural_break: - return SessionMemoryExtractionResult(reason="unsafe-boundary") + if self._last_event_has_tool_call(pending): + return SessionMemoryExtractionResult(False, "unsafe-boundary") + natural_break = not self._event_has_tool_response(pending[-1]) threshold_met = ((context_tokens >= threshold if checkpoint_context_tokens is None else (context_tokens < checkpoint_context_tokens or context_tokens - checkpoint_context_tokens >= threshold)) if token_mode else pending_chars >= threshold) diff --git a/trpc_agent_sdk/sessions/compact/default/_summarizer_manager.py b/trpc_agent_sdk/sessions/compact/default/_summarizer_manager.py index 65fd5af0b..3e0e2ddc4 100644 --- a/trpc_agent_sdk/sessions/compact/default/_summarizer_manager.py +++ b/trpc_agent_sdk/sessions/compact/default/_summarizer_manager.py @@ -113,7 +113,7 @@ async def create_session_summary(self, original_event_count=original_event_count, compressed_event_count=len(session.events), summary_timestamp=time.time(), - model_name=self._summarizer.model.name, + model_name=getattr(self._summarizer.model, "name", ""), ) session.conversation_count = 0 # Update the stored session diff --git a/trpc_agent_sdk/tools/__init__.py b/trpc_agent_sdk/tools/__init__.py index 071b5365a..20a937149 100644 --- a/trpc_agent_sdk/tools/__init__.py +++ b/trpc_agent_sdk/tools/__init__.py @@ -14,9 +14,9 @@ # Lazy re-export — see ``_LAZY_REEXPORTS`` below. from trpc_agent_sdk.agents.sub_agent import DynamicSubAgentTool as DynamicSubAgentTool # noqa: F401 from trpc_agent_sdk.agents.sub_agent import SpawnSubAgentTool as SpawnSubAgentTool # noqa: F401 - # from trpc_agent_sdk.tools._advanced_memory_tool import AdvancedMemoryTools as AdvancedMemoryTools # noqa: F401 - # from trpc_agent_sdk.tools._advanced_memory_tool import ( # noqa: F401 - # create_advanced_memory_tools as create_advanced_memory_tools, ) + from trpc_agent_sdk.tools._advanced_memory_tool import AdvancedMemoryTools as AdvancedMemoryTools # noqa: F401 + from trpc_agent_sdk.tools._advanced_memory_tool import ( # noqa: F401 + create_advanced_memory_tools as create_advanced_memory_tools, ) from ._agent_tool import AGENT_TOOL_APP_NAME_SUFFIX from ._agent_tool import AgentTool @@ -202,14 +202,14 @@ # the tools package free of optional file/web tool dependencies) but exposed # here for discoverability. Not in ``__all__`` so ``import *`` stays lazy. _LAZY_REEXPORTS = { - # "AdvancedMemoryTools": ( - # "trpc_agent_sdk.tools._advanced_memory_tool", - # "AdvancedMemoryTools", - # ), - # "create_advanced_memory_tools": ( - # "trpc_agent_sdk.tools._advanced_memory_tool", - # "create_advanced_memory_tools", - # ), + "AdvancedMemoryTools": ( + "trpc_agent_sdk.tools._advanced_memory_tool", + "AdvancedMemoryTools", + ), + "create_advanced_memory_tools": ( + "trpc_agent_sdk.tools._advanced_memory_tool", + "create_advanced_memory_tools", + ), "DynamicSubAgentTool": ("trpc_agent_sdk.agents.sub_agent", "DynamicSubAgentTool"), "SpawnSubAgentTool": ("trpc_agent_sdk.agents.sub_agent", "SpawnSubAgentTool"), } diff --git a/trpc_agent_sdk/tools/_advanced_memory_tool.py b/trpc_agent_sdk/tools/_advanced_memory_tool.py index e69de29bb..933fd77d2 100644 --- a/trpc_agent_sdk/tools/_advanced_memory_tool.py +++ b/trpc_agent_sdk/tools/_advanced_memory_tool.py @@ -0,0 +1,159 @@ +# 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 long-term memory read/write tools for standalone Advanced Memory.""" + +from __future__ import annotations + +import asyncio +from typing import Any + +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 + +ADVANCED_MEMORY_TOOL_NAMES = frozenset({ + "save_memory", + "read_memory", + "list_memory_index", +}) + + +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.""" + 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 tenant-scoped index update locks.""" + self._runtime = runtime + self._index_locks: dict[str, asyncio.Lock] = {} + self._tools = ( + FunctionTool(self.save_memory), + FunctionTool(self.read_memory), + FunctionTool(self.list_memory_index), + ) + + @property + def runtime(self) -> AdvancedMemoryRuntime: + """Return the Advanced Memory runtime bound to these tools.""" + return self._runtime + + def as_tools(self) -> list[FunctionTool]: + """Return tools that can be appended directly to LlmAgent.tools.""" + return list(self._tools) + + 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, + filename: str, + name: str, + description: str, + 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: + resolved_type = MemoryType(memory_type) + except ValueError as exc: + allowed = ", ".join(item.value for item in MemoryType) + raise ValueError(f"memory_type must be one of: {allowed}") from exc + document = MemoryDocument( + name=name, + description=description, + memory_type=resolved_type, + content=content, + ) + 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 runtime.long_term_memory.read_index()) + new_entry = MemoryIndexEntry( + name=name, + filename=path.name, + summary=summary, + ) + entries = [entry for entry in entries if entry.filename != new_entry.filename] + entries.insert(0, new_entry) + 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": 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, tool_context: Any | None = None) -> dict: + """Read a complete long-term memory by its filename in MEMORY.md.""" + 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) + freshness = memory_freshness(updated_at) + return { + "found": + True, + "filename": + filename, + "content": + content, + "updated_at": + updated_at.isoformat() if updated_at is not None else None, + "freshness": + freshness, + "freshness_notice": (f"This memory was last updated {freshness}. 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."), + } + + 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": _memory_index_reference(runtime), + "index": await runtime.long_term_memory.read_index(), + } + + +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()