From 36ee472f49de6fdad96b933ff2323b03fa6c4b3e Mon Sep 17 00:00:00 2001 From: congkechen Date: Wed, 9 Sep 2026 10:32:21 +0800 Subject: [PATCH 1/5] =?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/5] =?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/5] =?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/5] =?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 bcdb021b69e24e63b794067781fc8d47fe537bba Mon Sep 17 00:00:00 2001 From: congkechen Date: Thu, 10 Sep 2026 16:32:38 +0800 Subject: [PATCH 5/5] =?UTF-8?q?feature:=20Advanced=20memory=20=E9=95=BF?= =?UTF-8?q?=E6=9C=9F=E8=AE=B0=E5=BF=86=E5=AE=9E=E7=8E=B0=E4=BC=98=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../memory_service_with_advanced_memory/.env | 10 +- .../README.md | 150 +++-- .../run_agent.py | 52 +- .../.env | 4 +- .../README.md | 328 +++-------- .../run_agent.py | 9 +- .../.env | 9 +- .../README.md | 189 +++---- .../run_agent.py | 9 +- .../README.md | 2 +- .../README.md | 2 +- .../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/advanced_memory/_sql_stores.py | 533 ------------------ .../advanced_memory/_storage_backend.py | 30 - trpc_agent_sdk/memory/__init__.py | 12 +- .../memory/_advanced_memory_service.py | 41 +- .../{ => memory}/advanced_memory/__init__.py | 14 +- .../{ => memory}/advanced_memory/_config.py | 1 + .../memory/advanced_memory/_formats.py | 136 +++++ .../advanced_memory/_integration.py | 0 .../advanced_memory/_memory_context.py | 0 .../{ => memory}/advanced_memory/_paths.py | 0 .../advanced_memory/_preload_memory.py | 88 ++- .../advanced_memory/_redis_stores.py | 145 +---- .../{ => memory}/advanced_memory/_runtime.py | 26 +- .../memory/advanced_memory/_sql_stores.py | 315 +++++++++++ .../{ => memory}/advanced_memory/_storage.py | 60 +- trpc_agent_sdk/sessions/_session.py | 6 +- trpc_agent_sdk/tools/_advanced_memory_tool.py | 30 +- 31 files changed, 984 insertions(+), 1311 deletions(-) delete mode 100644 trpc_agent_sdk/advanced_memory/_sql_stores.py delete mode 100644 trpc_agent_sdk/advanced_memory/_storage_backend.py rename trpc_agent_sdk/{ => memory}/advanced_memory/__init__.py (75%) rename trpc_agent_sdk/{ => memory}/advanced_memory/_config.py (99%) create mode 100644 trpc_agent_sdk/memory/advanced_memory/_formats.py rename trpc_agent_sdk/{ => memory}/advanced_memory/_integration.py (100%) rename trpc_agent_sdk/{ => memory}/advanced_memory/_memory_context.py (100%) rename trpc_agent_sdk/{ => memory}/advanced_memory/_paths.py (100%) rename trpc_agent_sdk/{ => memory}/advanced_memory/_preload_memory.py (79%) rename trpc_agent_sdk/{ => memory}/advanced_memory/_redis_stores.py (54%) rename trpc_agent_sdk/{ => memory}/advanced_memory/_runtime.py (93%) create mode 100644 trpc_agent_sdk/memory/advanced_memory/_sql_stores.py rename trpc_agent_sdk/{ => memory}/advanced_memory/_storage.py (78%) 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/run_agent.py b/examples/memory_service_with_advanced_memory/run_agent.py index efd8014e0..9d34a0dd1 100644 --- a/examples/memory_service_with_advanced_memory/run_agent.py +++ b/examples/memory_service_with_advanced_memory/run_agent.py @@ -8,11 +8,10 @@ """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.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 @@ -23,35 +22,36 @@ from agent.agent import create_agent -load_dotenv(Path(__file__).with_name(".env")) +load_dotenv(Path(__file__).with_name(".env"), override=True) -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=("特别关注并主动记住用户长期稳定的兴趣爱好、" - "编程语言偏好、开发习惯和测试习惯。"), +def create_session_service() -> InMemorySessionService: + """Create the session service with the independent Compact manager.""" + compact_manager = AdvancedSessionCompactManager( + config=AdvancedCompactConfig(), ) - compact_config = AdvancedCompactConfig() - compact_manager = AdvancedSessionCompactManager(config=compact_config) - session_service = InMemorySessionService( + return InMemorySessionService( session_config=SessionServiceConfig( ttl=SessionServiceConfig.create_ttl_config( - enable=bool(session_ttl), - ttl_seconds=session_ttl_seconds, + enable=True, + ttl_seconds=60, cleanup_interval_seconds=5, ), store_historical_events=True, ), session_compact_manager=compact_manager, ) - return session_service, AdvancedMemoryService(config=config) + + +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: @@ -77,7 +77,8 @@ 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, memory_service = create_services(agent) + session_service = create_session_service() + memory_service = create_memory_service() from trpc_agent_sdk.runners import Runner runner = Runner( @@ -86,10 +87,6 @@ async def main() -> None: 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. " @@ -121,11 +118,6 @@ async def main() -> None: 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() 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..0a21b66ba 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/README.md b/examples/session_service_with_advanced_memory_redis/README.md index 91886f100..9039db31a 100644 --- a/examples/session_service_with_advanced_memory_redis/README.md +++ b/examples/session_service_with_advanced_memory_redis/README.md @@ -9,7 +9,7 @@ - AutoCompact 触发时生成的 Session Memory 压缩能力来自 `trpc_agent_sdk.sessions.compact`,长期记忆仍属于独立的 -`trpc_agent_sdk.advanced_memory`。 +`trpc_agent_sdk.memory.advanced_memory`。 ## 组装关系 diff --git a/examples/session_service_with_advanced_memory_sql/README.md b/examples/session_service_with_advanced_memory_sql/README.md index 522ce8555..c703b700f 100644 --- a/examples/session_service_with_advanced_memory_sql/README.md +++ b/examples/session_service_with_advanced_memory_sql/README.md @@ -9,7 +9,7 @@ - AutoCompact 触发时生成的 Session Memory 压缩能力来自 `trpc_agent_sdk.sessions.compact`,长期记忆仍属于独立的 -`trpc_agent_sdk.advanced_memory`。SQL 表结构不变,但活跃/历史 Event +`trpc_agent_sdk.memory.advanced_memory`。SQL 表结构不变,但活跃/历史 Event 会按原 Session 语义重新分区。 ## 组装关系 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/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_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/memory/__init__.py b/trpc_agent_sdk/memory/__init__.py index a4b768e07..e0a69ed5e 100644 --- a/trpc_agent_sdk/memory/__init__.py +++ b/trpc_agent_sdk/memory/__init__.py @@ -7,7 +7,7 @@ This module provides memory/RAG functionality including: - Abstract memory service interfaces -- In-memory memory service implementation +- In-memory, Redis, SQL, and Advanced Memory implementations """ from trpc_agent_sdk.abc import MemoryServiceABC as BaseMemoryService @@ -27,7 +27,6 @@ __all__ = [ "BaseMemoryService", "MemoryServiceConfig", - "AdvancedMemoryServiceConfig", "AdvancedMemoryService", "EventTtl", "InMemoryMemoryService", @@ -39,12 +38,3 @@ "extract_words_lower", "format_timestamp", ] - - -def __getattr__(name: str): - """Lazily expose Advanced Memory configuration without import cycles.""" - if name == "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..a95406f5e 100644 --- a/trpc_agent_sdk/memory/_advanced_memory_service.py +++ b/trpc_agent_sdk/memory/_advanced_memory_service.py @@ -11,23 +11,28 @@ from typing import Optional from typing import TYPE_CHECKING -from trpc_agent_sdk.abc import MemoryServiceABC as BaseMemoryService +from typing_extensions import override + +from trpc_agent_sdk.abc import MemoryServiceABC from trpc_agent_sdk.abc import MemoryServiceConfig from trpc_agent_sdk.abc import SearchMemoryResponse +from trpc_agent_sdk.abc import SessionABC from trpc_agent_sdk.abc import SessionServiceABC from trpc_agent_sdk.context import AgentContext -from trpc_agent_sdk.sessions import Session if TYPE_CHECKING: - from trpc_agent_sdk.advanced_memory import AdvancedMemoryServiceConfig - from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime - from trpc_agent_sdk.advanced_memory import LongTermMemoryIntegration + from trpc_agent_sdk.memory.advanced_memory import AdvancedMemoryServiceConfig + from trpc_agent_sdk.memory.advanced_memory import AdvancedMemoryRuntime + from trpc_agent_sdk.memory.advanced_memory import LongTermMemoryIntegration -class AdvancedMemoryService(BaseMemoryService): - """Expose user-scoped long-term Memory through the Runner memory API. +class AdvancedMemoryService(MemoryServiceABC): + """Expose tool-driven long-term Memory through the Runner memory API. - ``Runner`` calls :meth:`bind` automatically. Session compression is + ``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``. """ @@ -40,8 +45,8 @@ def __init__( 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 + 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") @@ -70,7 +75,7 @@ def integration(self) -> LongTermMemoryIntegration | None: 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 + 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: @@ -86,14 +91,21 @@ def bind(self, agent: Any, session_service: SessionServiceABC) -> SessionService self._bound_agent = agent return session_service + @override async def store_session( self, - session: Session, + session: SessionABC, agent_context: Optional[AgentContext] = None, ) -> None: - """Long-term Memory is updated explicitly through its tools.""" + """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, @@ -101,13 +113,14 @@ async def search_memory( limit: int = 10, agent_context: Optional[AgentContext] = None, ) -> SearchMemoryResponse: - """Return an empty legacy-style response. + """Return the standard empty response for compatibility. Advanced long-term memory is intentionally accessed through its ``save_memory``, ``read_memory``, and ``list_memory_index`` tools. """ return SearchMemoryResponse() + @override async def close(self) -> None: """Release service-owned local or external storage resources.""" await self._runtime.close() diff --git a/trpc_agent_sdk/advanced_memory/__init__.py b/trpc_agent_sdk/memory/advanced_memory/__init__.py similarity index 75% rename from trpc_agent_sdk/advanced_memory/__init__.py rename to trpc_agent_sdk/memory/advanced_memory/__init__.py index 0f7fe7ca5..04625f61a 100644 --- a/trpc_agent_sdk/advanced_memory/__init__.py +++ b/trpc_agent_sdk/memory/advanced_memory/__init__.py @@ -6,11 +6,11 @@ """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 ._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 @@ -27,18 +27,14 @@ 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", diff --git a/trpc_agent_sdk/advanced_memory/_config.py b/trpc_agent_sdk/memory/advanced_memory/_config.py similarity index 99% rename from trpc_agent_sdk/advanced_memory/_config.py rename to trpc_agent_sdk/memory/advanced_memory/_config.py index 99588568c..9bb9fcd2f 100644 --- a/trpc_agent_sdk/advanced_memory/_config.py +++ b/trpc_agent_sdk/memory/advanced_memory/_config.py @@ -19,6 +19,7 @@ def _require_positive(**values: int | float) -> None: 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: diff --git a/trpc_agent_sdk/memory/advanced_memory/_formats.py b/trpc_agent_sdk/memory/advanced_memory/_formats.py new file mode 100644 index 000000000..bea5c3d35 --- /dev/null +++ b/trpc_agent_sdk/memory/advanced_memory/_formats.py @@ -0,0 +1,136 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Data formats used by Advanced Memory.""" + +from __future__ import annotations + +import re +from dataclasses import dataclass +from datetime import datetime +from datetime import timezone +from enum import Enum + +_FRONTMATTER_PATTERN = re.compile(r"\A---\n(?P.*?)\n---(?:\n|\Z)", re.DOTALL) +_UPDATED_AT_PATTERN = re.compile(r"^updated_at:\s*(?P\S+)\s*$", re.MULTILINE) + + +def _as_utc(value: datetime) -> datetime: + """Normalize an aware or naive datetime to UTC.""" + if value.tzinfo is None: + value = value.replace(tzinfo=timezone.utc) + return value.astimezone(timezone.utc) + + +class MemoryType(str, Enum): + """Semantic types allowed for long-term memory documents.""" + + USER = "user" + FEEDBACK = "feedback" + PROJECT = "project" + REFERENCE = "reference" + + +@dataclass(frozen=True) +class MemoryIndexEntry: + """Represent one entry in MEMORY.md.""" + + name: str + filename: str + summary: str + + def __post_init__(self) -> None: + """Validate that index fields are non-empty single-line strings.""" + for field_name, value in ( + ("name", self.name), + ("filename", self.filename), + ("summary", self.summary), + ): + if not value.strip() or "\n" in value or "\r" in value: + raise ValueError(f"{field_name} must be non-empty single-line text") + + def to_markdown(self) -> str: + """Render one standard index entry.""" + return f"- [{self.name.strip()}]({self.filename.strip()}):{self.summary.strip()}" + + +@dataclass(frozen=True) +class MemoryDocument: + """Represent one long-term memory topic.""" + + name: str + description: str + memory_type: MemoryType + content: str + updated_at: datetime | None = None + + def __post_init__(self) -> None: + """Validate frontmatter fields.""" + for field_name, value in ( + ("name", self.name), + ("description", self.description), + ): + if not value.strip() or "\n" in value or "\r" in value: + raise ValueError(f"{field_name} must be non-empty single-line text") + + def to_markdown(self) -> str: + """Render the topic as Markdown with frontmatter.""" + body = self.content.strip() + updated_at = _as_utc(self.updated_at).isoformat() if self.updated_at is not None else None + updated_at_line = f"updated_at: {updated_at}\n" if updated_at else "" + return ( + "---\n" + f"name: {self.name.strip()}\n" + f"description: {self.description.strip()}\n" + f"type: {self.memory_type.value}\n" + f"{updated_at_line}" + "---\n" + f"{body}\n" + ) + + +def parse_memory_updated_at(content: str) -> datetime | None: + """Extract the UTC update timestamp from a memory document.""" + frontmatter_match = _FRONTMATTER_PATTERN.match(content) + if frontmatter_match is None: + return None + match = _UPDATED_AT_PATTERN.search(frontmatter_match.group("frontmatter")) + if match is None: + return None + try: + parsed = datetime.fromisoformat(match.group("value").replace("Z", "+00:00")) + except ValueError: + return None + return _as_utc(parsed) + + +def memory_freshness(updated_at: datetime | None, *, now: datetime | None = None) -> str: + """Return a compact freshness bucket for model-facing output.""" + if updated_at is None: + return "unknown" + age_days = max(0, int((_as_utc(now or datetime.now(timezone.utc)) - _as_utc(updated_at)).total_seconds()) + // 86_400) + if age_days == 0: + return "today" + if age_days == 1: + return "yesterday" + if age_days <= 7: + return "within 7 days" + if age_days <= 30: + return "within 30 days" + return "over 30 days" + + +def limit_memory_index(index: str, *, max_lines: int, max_bytes: int, encoding: str) -> str: + """Return a bounded view of an index without modifying the stored index.""" + lines: list[str] = [] + used_bytes = 0 + for line in index.splitlines(keepends=True)[:max_lines]: + size = len(line.encode(encoding)) + if used_bytes + size > max_bytes: + break + lines.append(line) + used_bytes += size + return "".join(lines) diff --git a/trpc_agent_sdk/advanced_memory/_integration.py b/trpc_agent_sdk/memory/advanced_memory/_integration.py similarity index 100% rename from trpc_agent_sdk/advanced_memory/_integration.py rename to trpc_agent_sdk/memory/advanced_memory/_integration.py diff --git a/trpc_agent_sdk/advanced_memory/_memory_context.py b/trpc_agent_sdk/memory/advanced_memory/_memory_context.py similarity index 100% rename from trpc_agent_sdk/advanced_memory/_memory_context.py rename to trpc_agent_sdk/memory/advanced_memory/_memory_context.py diff --git a/trpc_agent_sdk/advanced_memory/_paths.py b/trpc_agent_sdk/memory/advanced_memory/_paths.py similarity index 100% rename from trpc_agent_sdk/advanced_memory/_paths.py rename to trpc_agent_sdk/memory/advanced_memory/_paths.py diff --git a/trpc_agent_sdk/advanced_memory/_preload_memory.py b/trpc_agent_sdk/memory/advanced_memory/_preload_memory.py similarity index 79% rename from trpc_agent_sdk/advanced_memory/_preload_memory.py rename to trpc_agent_sdk/memory/advanced_memory/_preload_memory.py index d032a8c42..f2bcfce8b 100644 --- a/trpc_agent_sdk/advanced_memory/_preload_memory.py +++ b/trpc_agent_sdk/memory/advanced_memory/_preload_memory.py @@ -16,13 +16,10 @@ 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 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 @@ -91,15 +88,20 @@ def _candidate_from_content(filename: str, content: str) -> MemoryCandidate: class ModelMemoryRelevanceSelector: - """Use an isolated lightweight Agent to select relevant topic files.""" + """Use one direct LLM call to select relevant topic files.""" def __init__(self, model: object | None = None) -> None: """Store an optional dedicated selector model.""" self._model = model - def _resolve_model(self, ctx: "InvocationContext") -> object: - """Prefer a dedicated selector model and fall back to the main model.""" - model = self._model if self._model is not None else getattr(ctx.agent, "model", None) + async def _resolve_model(self, ctx: "InvocationContext") -> object: + """Prefer a dedicated selector model and resolve the main Agent model.""" + if self._model is not None: + return self._model + resolver = getattr(ctx.agent, "_resolve_model", None) + if callable(resolver): + return await resolver(ctx) + model = getattr(ctx.agent, "model", None) if model is None: raise ValueError("Memory relevance selector cannot resolve an LLM model") return model @@ -157,48 +159,32 @@ async def select( *, limit: int, ) -> list[str]: - """Run the isolated selector Agent and validate its result.""" - app_name = f"{ctx.app_name}_advanced_memory_selector" - agent = LlmAgent( - name="advanced_memory_relevance_selector", - description="Select relevant long-term memories.", - instruction=("You are a strict long-term memory relevance selector. " - "Follow the user's query and output format exactly."), - model=self._resolve_model(ctx), - tools=[], - add_name_to_instruction=False, + """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))], + ) + ], ) - 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() + 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( diff --git a/trpc_agent_sdk/advanced_memory/_redis_stores.py b/trpc_agent_sdk/memory/advanced_memory/_redis_stores.py similarity index 54% rename from trpc_agent_sdk/advanced_memory/_redis_stores.py rename to trpc_agent_sdk/memory/advanced_memory/_redis_stores.py index 07f4f73ae..8e0dba64a 100644 --- a/trpc_agent_sdk/advanced_memory/_redis_stores.py +++ b/trpc_agent_sdk/memory/advanced_memory/_redis_stores.py @@ -3,8 +3,6 @@ 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 @@ -16,14 +14,10 @@ from trpc_agent_sdk.types import Ttl from ._config import AdvancedMemoryServiceConfig -from trpc_agent_sdk.sessions.compact._formats import MemoryDocument, MemoryIndexEntry +from ._formats import MemoryDocument, MemoryIndexEntry +from ._formats import limit_memory_index 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 -""" +from ._storage import parse_memory_index, prune_memory_index _RELEASE_LOCK_SCRIPT = """ if redis.call('GET', KEYS[1]) == ARGV[1] then @@ -47,7 +41,6 @@ def __init__( 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) @@ -57,14 +50,6 @@ async def _command(self, method: str, *args: Any, **kwargs: Any) -> Any: 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" @@ -128,17 +113,6 @@ async def _refresh_ttl_group( 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(), @@ -146,29 +120,6 @@ async def _refresh_memory_ttl(self, *keys: str) -> None: 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: @@ -187,14 +138,23 @@ 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) + 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) @@ -235,68 +195,3 @@ async def list_topics(self) -> list[Path]: 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/memory/advanced_memory/_runtime.py similarity index 93% rename from trpc_agent_sdk/advanced_memory/_runtime.py rename to trpc_agent_sdk/memory/advanced_memory/_runtime.py index 31bd28134..fa16d565a 100644 --- a/trpc_agent_sdk/advanced_memory/_runtime.py +++ b/trpc_agent_sdk/memory/advanced_memory/_runtime.py @@ -15,7 +15,6 @@ 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 @@ -28,7 +27,6 @@ class AdvancedMemoryRuntime: config: AdvancedMemoryServiceConfig paths: AdvancedMemoryPaths - coordination: SessionOperationCoordinator long_term_memory: LongTermMemoryStore _scoped_runtimes: dict[MemoryScope, "ScopedAdvancedMemoryRuntime"] = field( default_factory=dict, @@ -42,6 +40,7 @@ class AdvancedMemoryRuntime: ) _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, @@ -57,28 +56,30 @@ def create(cls, config: AdvancedMemoryServiceConfig | None = None) -> "AdvancedM 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 + 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, - coordination=SessionOperationCoordinator(), long_term_memory=LongTermMemoryStore(resolved_config, paths), _redis_storage=redis_storage, _sql_storage=sql_storage, + _sql_cleanup=sql_cleanup, _local_cleanup=local_cleanup, ) @@ -147,6 +148,8 @@ async def initialize(self) -> bool: 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 @@ -166,6 +169,8 @@ async def close(self) -> 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) @@ -185,22 +190,13 @@ 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() + 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/advanced_memory/_storage.py b/trpc_agent_sdk/memory/advanced_memory/_storage.py similarity index 78% rename from trpc_agent_sdk/advanced_memory/_storage.py rename to trpc_agent_sdk/memory/advanced_memory/_storage.py index d4d17daea..d5a0e9f08 100644 --- a/trpc_agent_sdk/advanced_memory/_storage.py +++ b/trpc_agent_sdk/memory/advanced_memory/_storage.py @@ -8,6 +8,7 @@ import asyncio import os +import re import tempfile import time from dataclasses import replace @@ -15,12 +16,36 @@ 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 ._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) @@ -76,19 +101,26 @@ def _read_index_sync(self) -> str: 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) + 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) 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/tools/_advanced_memory_tool.py b/trpc_agent_sdk/tools/_advanced_memory_tool.py index 76bce5249..933fd77d2 100644 --- a/trpc_agent_sdk/tools/_advanced_memory_tool.py +++ b/trpc_agent_sdk/tools/_advanced_memory_tool.py @@ -8,15 +8,15 @@ from __future__ import annotations import asyncio -import re from typing import Any -from trpc_agent_sdk.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 trpc_agent_sdk.memory.advanced_memory._formats import MemoryDocument +from trpc_agent_sdk.memory.advanced_memory._formats import MemoryIndexEntry +from trpc_agent_sdk.memory.advanced_memory._formats import MemoryType +from trpc_agent_sdk.memory.advanced_memory._formats import memory_freshness +from trpc_agent_sdk.memory.advanced_memory._formats import parse_memory_updated_at +from trpc_agent_sdk.memory.advanced_memory._storage import parse_memory_index +from trpc_agent_sdk.memory.advanced_memory._runtime import AdvancedMemoryRuntime from ._function_tool import FunctionTool @@ -25,7 +25,6 @@ "read_memory", "list_memory_index", }) -_INDEX_PATTERN = re.compile(r"^- \[(?P.+?)\]((?P.+?)):(?P.+)$") def _memory_index_reference(runtime: Any) -> str: @@ -35,13 +34,7 @@ def _memory_index_reference(runtime: Any) -> str: def _parse_index(index: str) -> list[MemoryIndexEntry]: """Parse standard Advanced Memory index entries from MEMORY.md.""" - entries: list[MemoryIndexEntry] = [] - for line in index.splitlines(): - match = _INDEX_PATTERN.match(line.strip()) - if match is None: - continue - entries.append(MemoryIndexEntry(**match.groupdict())) - return entries + return parse_memory_index(index) class AdvancedMemoryTools: @@ -66,11 +59,6 @@ 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: @@ -166,6 +154,6 @@ async def list_memory_index(self, tool_context: Any | None = None) -> dict: } -def create_advanced_memory_tools(runtime: AdvancedMemoryRuntime, ) -> list[FunctionTool]: +def create_advanced_memory_tools(runtime: AdvancedMemoryRuntime) -> list[FunctionTool]: """Create the official Advanced Memory tools bound to the given runtime.""" return AdvancedMemoryTools(runtime).as_tools()