diff --git a/README.zh_CN.md b/README.zh_CN.md index 29f3ea24b..5a8e7fd0e 100644 --- a/README.zh_CN.md +++ b/README.zh_CN.md @@ -497,7 +497,7 @@ skill_tool_set = SkillToolSet(repository=repository, run_tool_kwargs=tool_kwargs 建议先看: -- Session:[examples/session_service_with_in_memory](./examples/session_service_with_in_memory/README.md) / [examples/session_service_with_redis](./examples/session_service_with_redis/README.md) / [examples/session_service_with_sql](./examples/session_service_with_sql/README.md) / [examples/session_summarizer](./examples/session_summarizer/README.md) / [examples/session_state](./examples/session_state/README.md) +- Session:[examples/session_service_with_in_memory](./examples/session_service_with_in_memory/README.md) / [examples/session_service_with_redis](./examples/session_service_with_redis/README.md) / [examples/session_service_with_sql](./examples/session_service_with_sql/README.md) / [Advanced Memory Redis 压缩](./examples/session_service_with_advanced_memory_redis/README.md) / [Advanced Memory SQL 压缩](./examples/session_service_with_advanced_memory_sql/README.md) / [examples/session_summarizer](./examples/session_summarizer/README.md) / [examples/session_state](./examples/session_state/README.md) - Memory: [examples/memory_service_with_in_memory](./examples/memory_service_with_in_memory/README.md) / [examples/memory_service_with_redis](./examples/memory_service_with_redis/README.md) / [examples/memory_service_with_sql](./examples/memory_service_with_sql/README.md) / [examples/memory_service_with_mem0](./examples/memory_service_with_mem0/README.md) / [examples/memory_service_with_mempalace](./examples/memory_service_with_mempalace/README.md) - Knowledge:[examples/knowledge_with_documentloader](./examples/knowledge_with_documentloader/README.md) / [examples/knowledge_with_vectorstore](./examples/knowledge_with_vectorstore/README.md) / [examples/knowledge_with_rag_agent](./examples/knowledge_with_rag_agent/README.md) / [examples/knowledge_with_searchtool_rag_agent](./examples/knowledge_with_searchtool_rag_agent/README.md) / [examples/knowledge_with_prompt_template](./examples/knowledge_with_prompt_template/README.md) / [examples/knowledge_with_custom_components](./examples/knowledge_with_custom_components/README.md) diff --git a/examples/memory_service_with_advanced_memory/.env b/examples/memory_service_with_advanced_memory/.env index 2da17e1ce..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..dd58e491a 100644 --- a/examples/memory_service_with_advanced_memory/README.md +++ b/examples/memory_service_with_advanced_memory/README.md @@ -1,223 +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`。 - -## 示例流程 - -脚本使用同一个 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, - ) +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, + memory_service=memory_service, ) ``` -`Runner` 检测到 `AdvancedMemorySessionService` 后会自动完成 Advanced Memory -绑定,包括: - -- transcript 持久化 -- session memory 提取 -- 长期记忆 tools:`save_memory`、`read_memory`、`list_memory_index` -- `HistorySnip` -- `Microcompact` -- `AutoCompact` -- `ToolResultBudget` +Session Compact 与 Advanced Memory 可以共享一个 Runtime;Runtime 的 `close()` +支持幂等调用,因此两个 Service 的正常关闭流程不会造成重复释放错误。 -`AdvancedMemoryConfig` 默认已经启用这些能力,本示例直接使用默认配置。 - -## 数据目录 - -运行后,数据默认写入当前示例目录: - -```text -MEMORY/ -├── MEMORY.md -└── *.md # 长期记忆详情 - -SESSION/ -├── _state.json # app/user 级 state -├── session-1/ -│ ├── session.json # Session 元数据和 session state -│ ├── transcript.jsonl # 原始 Events 和 checkpoint -│ ├── session_memory.md # 结构化 Session 记忆 -│ └── tool-results/ # 超大工具结果 -└── session-2/ - ├── session.json - ├── transcript.jsonl - └── session_memory.md -``` - -其中: - -- `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) - -`.env` 中留空的变量不会覆盖默认值;如果同时在 Python 中传入 -`model_context_window_tokens` 或 `max_output_tokens`,Python 显式配置优先。 - -如果配置了模型上下文窗口,Advanced Memory 会用 -`TRPC_AGENT_MODEL_CONTEXT_WINDOW_TOKENS - TRPC_AGENT_MAX_OUTPUT_TOKENS` -作为可用于输入内容的窗口;两个变量都留空时使用字符数阈值。 - -## `AdvancedMemoryConfig` 配置项 - -下面列出当前所有可直接传入 `AdvancedMemoryConfig` 的配置项。**没有特殊需求时, -只设置 `root_dir` 即可**;示例中的值均为默认值。 - -```python -session_service = AdvancedMemorySessionService( - config=AdvancedMemoryConfig( - root_dir=Path(__file__).resolve().parent, # 当前示例目录 - # Optional - enabled=True, # 总开关和存储路径 - memory_dir_name="MEMORY", # 长期记忆目录 - session_dir_name="SESSION", # Session 数据目录 - memory_index_name="MEMORY.md", # 长期记忆索引文件 - transcript_name="transcript.jsonl", # transcript 文件 - session_memory_name="session_memory.md", # Session 摘要文件 - encoding="utf-8", # 文件编码 - transcript_fsync=False, # transcript 写入后是否 fsync - - # 长期记忆 - memory_index_max_lines=200, # 注入 prompt 的索引最大行数 - memory_index_max_bytes=25_000, # 注入 prompt 的索引最大字节数 - long_term_memory_injection_enabled=True, # 是否注入 MEMORY.md - - # 工具结果 - tool_result_max_chars=50_000, # 单个工具结果最大字符数 - tool_results_per_message_max_chars=200_000, # 单条消息工具结果总上限 - tool_result_preview_chars=2_000, # 超限结果的预览字符数 - - # HistorySnip - history_snip_enabled=True, # 是否压缩过长历史 - history_snip_trigger_chars=600_000, # 触发阈值 - history_snip_target_chars=400_000, # 压缩目标 - history_snip_keep_recent=5, # 保留最近的完整消息数 - history_snip_tool_names=( # 可处理的工具名称 - "Read", "Bash", "Grep", "Glob", - "WebSearch", "WebFetch", "Edit", "Write", - ), - - # Token 上下文预算 - # 这两个值也可以通过 .env 配置;显式传参优先于环境变量。 - # model_context_window_tokens=131072, # 显式设置后覆盖环境变量 - # max_output_tokens=8192, # 显式设置后覆盖环境变量 - # 如果省略这两行,则分别读取 .env;未配置时默认 None 和 0。 - token_warning_ratio=0.85, # 告警比例 - token_autocompact_ratio=0.90, # 自动压缩比例 - token_blocking_ratio=0.95, # 阻止继续增加上下文的比例 - token_estimator=None, # 可选:自定义 token 估算器 - context_window_resolver=None, # 可选:自定义窗口解析器 - - # Session Memory - session_memory_enabled=True, # 是否启用 Session 摘要 - session_memory_initial_chars=40_000, # 首次提取字符阈值 - session_memory_update_chars=20_000, # 后续更新字符阈值 - session_memory_initial_tokens=10_000, # 首次提取 token 阈值 - session_memory_update_tokens=5_000, # 后续更新 token 阈值 - session_memory_tool_calls_between_updates=3, # 两次更新间的工具调用数 - session_memory_prompt_max_chars=200_000, # 摘要请求最大字符数 - session_memory_request_overhead_tokens=2_048, # 请求预留 token - session_memory_section_max_chars=8_000, # 单个摘要 section 最大字符数 - session_memory_total_max_chars=54_000, # 摘要总最大字符数 - session_memory_wait_timeout_seconds=15.0, # 等待摘要 Agent 的超时时间 - - # AutoCompact - autocompact_enabled=True, # 是否启用自动压缩 - autocompact_trigger_chars=700_000, # 触发阈值 - autocompact_target_chars=350_000, # 压缩目标 - autocompact_blocking_chars=780_000, # 阻止继续增加上下文的阈值 - autocompact_keep_recent_contents=8, # 保留最近内容数 - autocompact_max_failures=3, # 最大连续失败次数 - autocompact_summary_input_max_chars=600_000, # 摘要 Agent 输入上限 - autocompact_summary_retries=3, # 摘要 Agent 重试次数 - - # Microcompact - microcompact_enabled=True, # 是否启用工具结果微压缩 - microcompact_gap_seconds=3_600.0, # 工具结果时间间隔阈值 - microcompact_trigger_count=20, # 触发工具结果数量 - microcompact_keep_recent=5, # 保留最近工具结果数 - microcompact_tool_names=( # 可处理的工具名称 - "Read", "Bash", "Grep", "Glob", - "WebSearch", "WebFetch", "Edit", "Write", - ), - - # Advanced Memory preload - preload_memory_enabled=False, # 是否自动预加载相关 topic - preload_memory_max_topics=5, # 一次最多加载的 topic 数 - preload_memory_max_chars=50_000, # 预加载内容总字符上限 - preload_memory_candidate_limit=200, # 筛选模型的候选 topic 数 - ), -) +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 097a3271d..17e8298cb 100644 --- a/examples/memory_service_with_advanced_memory/run_agent.py +++ b/examples/memory_service_with_advanced_memory/run_agent.py @@ -8,29 +8,51 @@ """Run the two-session Advanced Memory demonstration.""" import asyncio +import os 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 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.""" - return AdvancedMemorySessionService( - config=AdvancedMemoryConfig(root_dir=Path(__file__).resolve().parent), - session_config=SessionServiceConfig(ttl=SessionServiceConfig.create_ttl_config( - ttl_seconds=60, - cleanup_interval_seconds=5, - )), +def create_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 = 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, + ), + ) + 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: @@ -56,14 +78,19 @@ 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 + 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 +122,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..6a46edc83 --- /dev/null +++ b/examples/memory_service_with_advanced_memory_redis/.env @@ -0,0 +1,8 @@ +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 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..c97328e13 --- /dev/null +++ b/examples/memory_service_with_advanced_memory_redis/README.md @@ -0,0 +1,342 @@ +# Advanced Memory Redis 示例 + +本示例演示如何将 Advanced Memory 的本地文件存储切换为 Redis,并验证: + +- Redis:`AdvancedMemoryService(storage_backend="redis")` +- 长期 memory 可以跨 Python 进程持久化; +- 同一用户在不同 `session_id` 中可以读取自己的长期 memory; +- session 相关数据和长期 memory 可以分别设置 TTL; +- Redis 中的 Markdown、Stream 和索引数据如何组织。 + +本示例只关注长期 Memory 的 Redis 持久化: + +```text +AdvancedMemoryService +└── Redis 保存长期 memory index 和 topic + +Runner +└── InMemorySessionService(仅用于运行示例) +``` + +## 环境要求 + +- 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 + +``` + +TTL 规则: + +- `M_TTL` 管理用户级长期 memory 的全部 Redis key; +- TTL 会在访问或写入时刷新,是“最后一次活动后过期”; +- `M_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( + AdvancedCompactConfig( + 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, +) +``` + +其中: + +- 用户只需要配置长期 Memory 的 `M_TTL`; +- `AdvancedMemoryService` 只负责长期 memory; +- Session Service 的 Redis 高级压缩接入请看 + [`session_service_with_advanced_memory_redis`](../session_service_with_advanced_memory_redis/); +- 运行请求时通过 `user_id` 和 `session_id` 指定用户及会话。 + +## 运行结果(实测) + +```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! + +----- Runner A, query 2 ----- + +📝 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. + +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`。 + +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`。 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..93dce8789 --- /dev/null +++ b/examples/memory_service_with_advanced_memory_redis/run_agent.py @@ -0,0 +1,131 @@ +#!/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 AdvancedCompactConfig +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")) + +RUNNER_A_QUERIES = [ + "Do you remember my name?", + "Do you remember my favorite color?", + "what is the weather like in paris?", + "Hello! My name is Alice. What's your name?", + "Do you remember my name?", + "Hello! My favorite color is blue. What's your favorite color?", + "Do you remember my favorite color?", +] + +RUNNER_B_QUERIES = [ + "Do you remember my name?", + "Do you remember my favorite color?", +] + + +def build_redis_url_from_environment() -> str: + """Use REDIS_URL directly, or construct it from standard Redis variables.""" + redis_url = os.getenv("REDIS_URL") + if redis_url: + return redis_url + + host = os.getenv("REDIS_HOST", "127.0.0.1") + port = os.getenv("REDIS_PORT", "6379") + database = os.getenv("REDIS_DB", "0") + username = os.getenv("REDIS_USER", "") + password = os.getenv("REDIS_PASSWORD", "") + scheme = "rediss" if os.getenv("REDIS_TLS", "").lower() in {"1", "true", "yes"} else "redis" + + if username and password: + auth = f"{quote(username, safe='')}:{quote(password, safe='')}@" + elif password: + auth = f":{quote(password, safe='')}@" + else: + auth = "" + return f"{scheme}://{auth}{host}:{port}/{database}" + + +def create_advanced_memory_service(redis_url: str) -> AdvancedMemoryService: + """Create the long-term Advanced Memory service backed by Redis.""" + memory_ttl = os.getenv("M_TTL") + 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, + ) + return AdvancedMemoryService(config) + + +async def ask(runner: Runner, session_id: str, prompt: str) -> None: + """Send one message through the shared app and user identity.""" + print(f"\n📝 user: {prompt}") + async for event in runner.run_async( + user_id="redis-demo-user", + session_id=session_id, + new_message=Content(parts=[Part.from_text(text=prompt)]), + ): + if not event.content or not event.content.parts: + continue + for part in event.content.parts: + if part.function_call: + print(f"🔧 tool call: {part.function_call.name}({part.function_call.args})") + elif part.function_response: + print(f"📊 Tool Result: {part.function_response.response}") + elif not event.partial and part.text and not part.thought: + print(f"🤖 Assistant: {part.text}") + + +async def run_phase(phase: str) -> None: + """Run Runner A or Runner B against the same Redis user.""" + app_name = "advanced-memory-redis-demo" + redis_url = build_redis_url_from_environment() + runner = Runner( + app_name=app_name, + agent=create_agent(), + session_service=InMemorySessionService(), + memory_service=create_advanced_memory_service(redis_url), + ) + try: + queries = RUNNER_A_QUERIES if phase == "write" else RUNNER_B_QUERIES + runner_name = "A" if phase == "write" else "B" + for index, prompt in enumerate(queries): + print(f"\n----- Runner {runner_name}, query {index + 1} -----") + await ask(runner, f"redis-{phase}-session-{index}", prompt) + finally: + await runner.close() + + +def run_two_processes() -> None: + """Start fresh writer and reader processes to prove Redis persistence.""" + for phase in ("write", "read"): + print(f"\n{'=' * 20} {phase.upper()} PROCESS {'=' * 20}", flush=True) + subprocess.run( + [sys.executable, str(Path(__file__).resolve()), "--phase", phase], + check=True, + ) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--phase", choices=("write", "read")) + arguments = parser.parse_args() + if arguments.phase: + asyncio.run(run_phase(arguments.phase)) + else: + run_two_processes() diff --git a/examples/memory_service_with_advanced_memory_sql/.env b/examples/memory_service_with_advanced_memory_sql/.env new file mode 100644 index 000000000..81dbccf4a --- /dev/null +++ b/examples/memory_service_with_advanced_memory_sql/.env @@ -0,0 +1,13 @@ +# Model configuration +TRPC_AGENT_API_KEY= +TRPC_AGENT_BASE_URL= +TRPC_AGENT_MODEL_NAME= + +# Easy local test with SQLite. SQL_IS_ASYNC=false uses the built-in sqlite driver. +# SQL_URL=sqlite:///advanced-memory-sql-demo.db +# SQL_IS_ASYNC=false + +# For MySQL, replace SQL_URL and set SQL_IS_ASYNC=true: +SQL_URL= +SQL_IS_ASYNC=true +M_TTL=120 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..87dbac3f4 --- /dev/null +++ b/examples/memory_service_with_advanced_memory_sql/README.md @@ -0,0 +1,174 @@ +# Advanced Memory SQL 示例 + +本示例使用 SQL 保存 Advanced Memory,并验证同一用户的长期 memory 可以跨 Python 进程和不同 session 读取。 + +- SQL:`AdvancedMemoryService(storage_backend="sql")` + +```text +AdvancedMemoryService +└── SQL 保存长期 memory index 和 topic + +Runner +└── InMemorySessionService(仅用于运行示例) +``` + +## 配置 + +默认使用 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 的过期时间,单位为秒。 + +更多 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( + AdvancedCompactConfig( + 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, +) +``` + +其中: + +- 用户只需要配置长期 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。 + +## 运行结果(实测) + +```text + +==================== 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 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! + +----- 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': '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。 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..6fdd3d2f1 --- /dev/null +++ b/examples/memory_service_with_advanced_memory_sql/run_agent.py @@ -0,0 +1,122 @@ +#!/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 AdvancedCompactConfig +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")) + +RUNNER_A_QUERIES = [ + "Do you remember my name?", + "Do you remember my favorite color?", + "what is the weather like in paris?", + "Hello! My name is Alice. What's your name?", + "Do you remember my name?", + "Hello! My favorite color is blue. What's your favorite color?", + "Do you remember my favorite color?", +] + +RUNNER_B_QUERIES = [ + "Do you remember my name?", + "Do you remember my favorite color?", +] + + +def build_sql_url_from_environment() -> str: + """Use SQL_URL or build a MySQL URL from standard environment variables.""" + sql_url = os.getenv("SQL_URL") + if sql_url: + return sql_url + + user = quote(os.getenv("MYSQL_USER", "root"), safe="") + password = quote(os.getenv("MYSQL_PASSWORD", ""), safe="") + host = os.getenv("MYSQL_HOST", "127.0.0.1") + port = os.getenv("MYSQL_PORT", "3306") + database = os.getenv("MYSQL_DB", "trpc_agent_advanced_memory") + return f"mysql+aiomysql://{user}:{password}@{host}:{port}/{database}?charset=utf8mb4" + + +def sql_is_async() -> bool: + """Return whether the configured SQL driver is asynchronous.""" + return os.getenv("SQL_IS_ASYNC", "true").lower() in {"1", "true", "yes"} + + +def create_advanced_memory_service(sql_url: str) -> AdvancedMemoryService: + """Create the long-term Advanced Memory service backed by SQL.""" + memory_ttl = os.getenv("M_TTL") + 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, + ) + return AdvancedMemoryService(config) + + +async def run_phase(phase: str) -> None: + """Run Runner A or Runner B against the same SQL database.""" + sql_url = build_sql_url_from_environment() + runner = Runner( + app_name="advanced-memory-sql-demo", + agent=create_agent(), + session_service=InMemorySessionService(), + memory_service=create_advanced_memory_service(sql_url), + ) + try: + queries = RUNNER_A_QUERIES if phase == "write" else RUNNER_B_QUERIES + runner_name = "A" if phase == "write" else "B" + for index, prompt in enumerate(queries): + print(f"\n----- Runner {runner_name}, query {index + 1} -----") + print(f"📝 user: {prompt}") + async for event in runner.run_async( + user_id="sql-demo-user", + session_id=f"sql-{phase}-session-{index}", + new_message=Content(parts=[Part.from_text(text=prompt)]), + ): + if not event.content or not event.content.parts: + continue + for part in event.content.parts: + if part.function_call: + print(f"🔧 tool call: {part.function_call.name}({part.function_call.args})") + elif part.function_response: + print(f"📊 Tool Result: {part.function_response.response}") + elif not event.partial and part.text and not part.thought: + print(f"🤖 Assistant: {part.text}") + finally: + await runner.close() + + +def run_two_processes() -> None: + """Start independent writer and reader processes.""" + for phase in ("write", "read"): + print(f"\n{'=' * 20} {phase.upper()} PROCESS {'=' * 20}", flush=True) + subprocess.run( + [sys.executable, str(Path(__file__).resolve()), "--phase", phase], + check=True, + ) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--phase", choices=("write", "read")) + args = parser.parse_args() + if args.phase: + asyncio.run(run_phase(args.phase)) + else: + run_two_processes() diff --git a/examples/session_service_with_advanced_memory_redis/.env b/examples/session_service_with_advanced_memory_redis/.env new file mode 100644 index 000000000..4858f369a --- /dev/null +++ b/examples/session_service_with_advanced_memory_redis/.env @@ -0,0 +1,10 @@ +TRPC_AGENT_API_KEY= +TRPC_AGENT_BASE_URL= +TRPC_AGENT_MODEL_NAME= + +REDIS_USER= +REDIS_PASSWORD= +REDIS_HOST=127.0.0.1 +REDIS_PORT=6379 +REDIS_DB=0 +SESSION_ID=simple-demo diff --git a/examples/session_service_with_advanced_memory_redis/README.md b/examples/session_service_with_advanced_memory_redis/README.md new file mode 100644 index 000000000..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 3854dd2e1..000000000 --- a/tests/advanced_memory/test_advanced_memory_session_service.py +++ /dev/null @@ -1,221 +0,0 @@ -"""Tests for the standalone Advanced Memory SessionService.""" - -from __future__ import annotations - -import asyncio -import json -from pathlib import Path -from types import SimpleNamespace - -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig -from trpc_agent_sdk.events import Event -from trpc_agent_sdk.runners import Runner -from trpc_agent_sdk.sessions import AdvancedMemorySessionService -from trpc_agent_sdk.sessions import SessionServiceConfig -from trpc_agent_sdk.types import Content -from trpc_agent_sdk.types import Part - - -def _event(event_id: str, text: str) -> Event: - """Create a deterministic event for persistence tests.""" - return Event( - id=event_id, - invocation_id=f"invocation-{event_id}", - author="user", - content=Content(parts=[Part.from_text(text=text)]), - ) - - -def _config(root_dir: Path) -> AdvancedMemoryConfig: - """Disable model-driven background work for storage-only tests.""" - return AdvancedMemoryConfig( - root_dir=root_dir, - session_memory_enabled=False, - history_snip_enabled=False, - microcompact_enabled=False, - autocompact_enabled=False, - ) - - -async def test_session_service_persists_and_restores_events(tmp_path: Path) -> None: - """Ensure a new service instance can restore a complete transcript.""" - first = AdvancedMemorySessionService(config=_config(tmp_path)) - session = await first.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="demo-session", - state={ - "session-key": "session-value", - "app:theme": "dark", - "user:name": "alice", - }, - ) - await first.append_event(session, _event("event-1", "hello")) - metadata = json.loads((first.runtime.paths.session_dir(session.id) / "session.json").read_text(encoding="utf-8")) - assert metadata["state"] == {"session-key": "session-value"} - - second = AdvancedMemorySessionService(config=_config(tmp_path)) - restored = await second.get_session( - app_name="demo-app", - user_id="demo-user", - session_id="demo-session", - ) - - assert restored is not None - assert [event.id for event in restored.events] == ["event-1"] - assert restored.state["session-key"] == "session-value" - assert restored.state["app:theme"] == "dark" - assert restored.state["user:name"] == "alice" - - -async def test_session_id_collision_between_users_is_rejected(tmp_path: Path) -> None: - """Prevent different users from silently sharing one session directory.""" - service = AdvancedMemorySessionService(config=_config(tmp_path)) - await service.create_session( - app_name="demo-app", - user_id="user-a", - session_id="shared-session", - ) - - try: - await service.create_session( - app_name="demo-app", - user_id="user-b", - session_id="shared-session", - ) - except ValueError as exc: - assert "already used" in str(exc) - else: - raise AssertionError("Expected a cross-user session ID collision to fail") - - -async def test_delete_session_removes_persistent_session_data(tmp_path: Path) -> None: - """Ensure deleting a session removes its metadata and transcript directory.""" - service = AdvancedMemorySessionService(config=_config(tmp_path)) - session = await service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="delete-me", - state={ - "app:theme": "dark", - "user:name": "alice" - }, - ) - await service.append_event(session, _event("event-1", "hello")) - metadata = json.loads((service.runtime.paths.session_dir(session.id) / "session.json").read_text(encoding="utf-8")) - assert metadata["state"] == {} - - await service.delete_session( - app_name="demo-app", - user_id="demo-user", - session_id=session.id, - ) - - assert await service.get_session( - app_name="demo-app", - user_id="demo-user", - session_id=session.id, - ) is None - assert not service.runtime.paths.session_dir(session.id).exists() - - -async def test_ttl_cleanup_removes_expired_persistent_sessions(tmp_path: Path) -> None: - """Ensure configured session TTL removes idle session directories.""" - session_config = SessionServiceConfig(ttl=SessionServiceConfig.create_ttl_config( - ttl_seconds=1, - cleanup_interval_seconds=0.05, - )) - service = AdvancedMemorySessionService( - config=_config(tmp_path), - session_config=session_config, - ) - session = await service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="expires", - ) - - await asyncio.sleep(1.1) - - assert await service.get_session( - app_name="demo-app", - user_id="demo-user", - session_id=session.id, - ) is None - await service.close() - - -async def test_runner_binds_standalone_session_service(tmp_path: Path) -> None: - """Ensure Runner installs Advanced callbacks without a memory service.""" - service = AdvancedMemorySessionService(config=_config(tmp_path)) - agent = SimpleNamespace(name="test-agent", tools=[], before_model_callback=None) - - runner = Runner( - app_name="demo-app", - agent=agent, - session_service=service, - ) - - assert runner.session_service.delegate is service - assert service.integration is not None - - -def test_session_service_can_cross_event_loops(tmp_path: Path) -> None: - """Ensure deferred Runner work can share the service's file lock.""" - service = AdvancedMemorySessionService(config=_config(tmp_path)) - - async def create() -> None: - await service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="demo-session", - ) - - async def append() -> None: - session = await service.get_session( - app_name="demo-app", - user_id="demo-user", - session_id="demo-session", - ) - assert session is not None - await service.append_event(session, _event("event-1", "hello")) - - asyncio.run(create()) - asyncio.run(append()) - - -def test_transcript_decorator_can_cross_event_loops(tmp_path: Path) -> None: - """Ensure the wrapped service is safe for deferred-worker event loops.""" - service = AdvancedMemorySessionService(config=_config(tmp_path)) - runner = Runner( - app_name="demo-app", - agent=SimpleNamespace( - name="test-agent", - tools=[], - before_model_callback=None, - get_subagents=lambda: [], - ), - session_service=service, - ) - - async def create() -> None: - await runner.session_service.create_session( - app_name="demo-app", - user_id="demo-user", - session_id="wrapped-session", - ) - - async def append() -> None: - session = await runner.session_service.get_session( - app_name="demo-app", - user_id="demo-user", - session_id="wrapped-session", - ) - assert session is not None - await runner.session_service.append_event(session, _event("event-1", "hello")) - - asyncio.run(create()) - asyncio.run(append()) - records = asyncio.run(service.runtime.transcripts.read_all("wrapped-session")) - assert [record["event_id"] for record in records if record.get("kind") == "event"] == ["event-1"] - asyncio.run(runner.close()) diff --git a/tests/advanced_memory/test_advanced_memory_tools.py b/tests/advanced_memory/test_advanced_memory_tools.py index 12c83705b..750777d1b 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 AdvancedCompactConfig +from trpc_agent_sdk.advanced_memory import AdvancedMemoryPaths from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime from trpc_agent_sdk.tools import AdvancedMemoryTools from trpc_agent_sdk.tools import create_advanced_memory_tools @@ -14,10 +17,10 @@ def _runtime(tmp_path: Path) -> AdvancedMemoryRuntime: """Create a test runtime with long-term memory enabled.""" - return AdvancedMemoryRuntime.create(AdvancedMemoryConfig( + return AdvancedMemoryRuntime.create(AdvancedCompactConfig( enabled=True, root_dir=tmp_path, - )) + )).for_scope("demo-app", "demo-user") async def test_save_read_and_update_memory_index(tmp_path: Path) -> None: @@ -68,6 +71,34 @@ async def test_save_memory_rejects_unknown_type(tmp_path: Path) -> None: ) +@pytest.mark.parametrize( + ("storage_backend", "expected_prefix"), + (("redis", "advanced-memory://redis/"), ("sql", "advanced-memory://sql/")), +) +async def test_list_memory_index_reports_backend_storage_reference( + storage_backend: str, + expected_prefix: str, +) -> None: + """Avoid exposing a local filesystem path for external memory stores.""" + config = 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, + ) + 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_memory_context.py b/tests/advanced_memory/test_memory_context.py index 965be4f65..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, )) @@ -104,49 +105,65 @@ async def test_long_term_memory_index_is_injected_once(tmp_path: Path) -> None: assert "secrets, credentials, tokens, and other sensitive data" in instruction -async def test_unified_setup_installs_complete_pipeline_in_order(tmp_path: Path) -> None: - """Ensure unified setup installs the five components in order.""" +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( + 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_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 == { @@ -156,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 8a854da8f..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,14 +32,14 @@ async def select(self, query, candidates, ctx, *, limit): async def test_preloader_injects_selected_topic_with_budget(tmp_path: Path) -> None: """Ensure selected topic content is rendered and bounded.""" runtime = AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( + AdvancedCompactConfig( enabled=True, root_dir=tmp_path, preload_memory_enabled=True, preload_memory_max_chars=200, )) await runtime.initialize() - await runtime.long_term_memory.write_topic( + await runtime.for_scope("demo-app", "demo-user").long_term_memory.write_topic( "project.md", MemoryDocument( name="Project", @@ -48,10 +48,15 @@ async def test_preloader_injects_selected_topic_with_budget(tmp_path: Path) -> N content="important project details", ), ) + ctx = SimpleNamespace(session=SimpleNamespace( + app_name="demo-app", + user_id="demo-user", + id="session-a", + )) result = await MemoryPreloader(runtime, _FakeSelector()).preload( "What is relevant?", - SimpleNamespace(), + ctx, ) assert result is not None @@ -63,14 +68,14 @@ async def test_preloader_injects_selected_topic_with_budget(tmp_path: Path) -> N async def test_preloader_marks_truncated_content(tmp_path: Path) -> None: """Tell the main model when the configured content budget truncated a topic.""" runtime = AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( + AdvancedCompactConfig( enabled=True, root_dir=tmp_path, preload_memory_enabled=True, preload_memory_max_chars=12, )) await runtime.initialize() - await runtime.long_term_memory.write_topic( + await runtime.for_scope("demo-app", "demo-user").long_term_memory.write_topic( "project.md", MemoryDocument( name="Project", @@ -79,10 +84,15 @@ async def test_preloader_marks_truncated_content(tmp_path: Path) -> None: content="important project details", ), ) + ctx = SimpleNamespace(session=SimpleNamespace( + app_name="demo-app", + user_id="demo-user", + id="session-a", + )) result = await MemoryPreloader(runtime, _FakeSelector()).preload( "What is relevant?", - SimpleNamespace(), + ctx, ) assert result is not None @@ -93,13 +103,13 @@ async def test_preloader_marks_truncated_content(tmp_path: Path) -> None: async def test_preloader_failure_is_best_effort(tmp_path: Path) -> None: """Return no prompt content when relevance screening fails.""" runtime = AdvancedMemoryRuntime.create( - AdvancedMemoryConfig( + AdvancedCompactConfig( enabled=True, root_dir=tmp_path, preload_memory_enabled=True, )) await runtime.initialize() - await runtime.long_term_memory.write_topic( + await runtime.for_scope("demo-app", "demo-user").long_term_memory.write_topic( "project.md", MemoryDocument( name="Project", diff --git a/tests/advanced_memory/test_redis_stores.py b/tests/advanced_memory/test_redis_stores.py new file mode 100644 index 000000000..b681bbc9f --- /dev/null +++ b/tests/advanced_memory/test_redis_stores.py @@ -0,0 +1,142 @@ +"""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 new file mode 100644 index 000000000..b29745c55 --- /dev/null +++ b/tests/advanced_memory/test_sql_stores.py @@ -0,0 +1,113 @@ +"""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 index d529e3e9c..e4f025cd5 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 @@ -11,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: @@ -33,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 @@ -44,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: @@ -52,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() @@ -66,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)) @@ -150,6 +169,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 +253,40 @@ 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 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, @@ -251,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, @@ -268,6 +345,61 @@ 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, + 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)) @@ -286,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 77% rename from tests/advanced_memory/test_autocompact.py rename to tests/sessions/compact/test_autocompact.py index af4e77a33..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, @@ -62,7 +63,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 +84,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"), ) @@ -109,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, @@ -122,7 +186,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", @@ -137,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( @@ -435,7 +533,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/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 94% rename from tests/advanced_memory/test_session_memory_extractor.py rename to tests/sessions/compact/test_session_memory_extractor.py index 240e1f2fc..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, @@ -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 @@ -157,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, @@ -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/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 90% rename from tests/advanced_memory/test_tool_result_budget.py rename to tests/sessions/compact/test_tool_result_budget.py index 4df8d9036..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, @@ -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(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) diff --git a/tests/advanced_memory/test_transcript_session_service.py b/tests/sessions/compact/test_transcript_session_service.py similarity index 71% rename from tests/advanced_memory/test_transcript_session_service.py rename to tests/sessions/compact/test_transcript_session_service.py index a23217011..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,14 +37,14 @@ 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) 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" @@ -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") @@ -65,13 +67,13 @@ 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"] 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")) @@ -79,7 +81,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" @@ -87,23 +89,23 @@ 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")) - 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" 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,13 +117,22 @@ 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) 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/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 c346cd2d3..658252342 100644 --- a/trpc_agent_sdk/advanced_memory/__init__.py +++ b/trpc_agent_sdk/advanced_memory/__init__.py @@ -3,95 +3,46 @@ # 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 ._preload_memory import MemoryCandidate from ._preload_memory import MemoryPreloader from ._preload_memory import MemoryRelevanceSelector from ._preload_memory import ModelMemoryRelevanceSelector from ._preload_memory import select_relevant_memory_filenames -from ._runtime import AdvancedMemoryRuntime -from ._session_memory import build_session_memory_prompt -from ._session_memory import ForkedSessionMemoryGenerator -from ._session_memory import has_session_memory_content -from ._session_memory import limit_session_memory_document -from ._session_memory import SessionMemoryExtractionInput -from ._session_memory import SessionMemoryExtractionResult -from ._session_memory import SessionMemoryExtractor -from ._session_service import TranscriptSessionService -from ._integration import AdvancedContextManagement -from ._integration import AdvancedMemoryIntegration -from ._integration import setup_advanced_memory -from ._integration import setup_context_management -from ._storage import LongTermMemoryStore -from ._storage import SessionMemoryStore -from ._storage import ToolResultStore -from ._storage import TranscriptStore -from ._tool_result_budget import setup_tool_result_budget -from ._tool_result_budget import ToolResultBudget -from ._tool_result_budget import ToolResultBudgetCallback -from ._tool_result_budget import ToolResultBudgetResult -from ._transcript import TRANSCRIPT_SCHEMA_VERSION -from ._token_budget import ContextBudget -from ._token_budget import ContextTokenEstimate -from ._token_budget import HeuristicTokenEstimator -from ._token_budget import ModelContextWindowResolver -from ._token_budget import TokenContextTracker -from ._token_budget import TokenEstimator +from ._storage_backend import AdvancedMemoryStorageBackend +from ._storage_backend import LocalAdvancedMemoryStorageBackend __all__ = [ - "AutoCompact", - "AutoCompactCallback", - "AutoCompactResult", - "AdvancedMemoryConfig", - "AdvancedContextManagement", - "AdvancedMemoryIntegration", + "AdvancedMemoryStorageBackend", + "AdvancedCompactConfig", + "LongTermMemoryIntegration", "AdvancedMemoryPaths", "AdvancedMemoryRuntime", - "ContextBudget", - "ContextTokenEstimate", - "build_session_memory_prompt", - "content_signature", - "estimate_request_chars", - "ForkedLegacySummaryGenerator", - "ForkedSessionMemoryGenerator", - "has_session_memory_content", - "HistorySnip", - "HistorySnipCallback", - "HistorySnipResult", - "HeuristicTokenEstimator", + "ScopedAdvancedMemoryRuntime", "LongTermMemoryStore", + "LocalAdvancedMemoryStorageBackend", "LongTermMemoryContext", "LongTermMemoryContextCallback", "MemoryDocument", + "MemoryScope", "MemoryIndexEntry", "MemoryType", "MemoryCandidate", @@ -101,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 947db2330..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 @@ -32,17 +32,24 @@ def runtime(self) -> AdvancedMemoryRuntime: """Return the runtime bound to this long-term memory context.""" return self._runtime - async def apply(self, request: "LlmRequest") -> bool: + async def apply(self, request: "LlmRequest", ctx: "InvocationContext | None" = None) -> bool: """Append the MEMORY.md index and on-demand read guidance.""" - config = self._runtime.config + runtime = self._runtime.for_session(ctx.session) if ctx is not None else self._runtime + config = runtime.config if not config.enabled or not config.long_term_memory_injection_enabled: return False - await self._runtime.initialize() + await runtime.initialize() existing_instruction = (str(request.config.system_instruction) if request.config is not None and request.config.system_instruction else "") if LONG_TERM_MEMORY_MARKER in existing_instruction: return False - index = await self._runtime.long_term_memory.read_index() + index = await runtime.long_term_memory.read_index() + focus_instruction = (config.memory_focus_instruction or "").strip() + custom_focus = ("\n\n## Custom memory focus\n" + "The following is an additional application-level memory preference. " + "Give it extra attention when deciding whether stable, explicit information " + "is worth saving, while still following the safety and quality rules above:\n" + f"{focus_instruction}\n" if focus_instruction else "") instruction = ( f"{LONG_TERM_MEMORY_MARKER}\n" "The following is a bounded index of this project's long-term memory. It is a trusted cross-session " @@ -66,13 +73,15 @@ async def apply(self, request: "LlmRequest") -> bool: "Do not save temporary task details, information reconstructable from current code, unverified guesses, " "duplicates, the model's own reasoning, or secrets, credentials, tokens, and other sensitive data. " "Do not write information that is uncertain, useful only in the current conversation, or not clearly " - "worth preserving.\n\n" + f"worth preserving.{custom_focus}\n\n" "save_memory writes both the detail file and the index. Pass a stable filename and concise " "name/description/summary, and use one of user, feedback, project, or reference for memory_type. " "Keep the description short and general; put detailed information in content. " "If save_memory is unavailable, do not claim that the information was saved.\n" - f"Memory directory: {self._runtime.paths.memory_dir}\n" - f"Index file: {self._runtime.paths.memory_index_path}\n" + f"Memory directory: " + f"{runtime.paths.memory_dir if config.storage_backend == 'local' else config.storage_backend.upper()}\n" + f"Index file: " + f"{runtime.paths.storage_reference('memory_index')}\n" f"\n{index.rstrip()}\n\n" f"") request.append_instructions([instruction]) @@ -95,8 +104,7 @@ def memory_context(self) -> LongTermMemoryContext: async def __call__(self, ctx: "InvocationContext", request: "LlmRequest") -> None: """Inject the long-term memory index before a model request.""" - del ctx - await self._memory_context.apply(request) + await self._memory_context.apply(request, ctx) return None diff --git a/trpc_agent_sdk/advanced_memory/_paths.py b/trpc_agent_sdk/advanced_memory/_paths.py deleted file mode 100644 index da1a41e07..000000000 --- a/trpc_agent_sdk/advanced_memory/_paths.py +++ /dev/null @@ -1,101 +0,0 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. -# -# Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Safe path resolution for the independent memory mechanism.""" - -from __future__ import annotations - -import hashlib -import re -from dataclasses import dataclass -from pathlib import Path - -from ._config import AdvancedMemoryConfig - -_SAFE_COMPONENT_PATTERN = re.compile(r"[^A-Za-z0-9._-]+") - - -def _safe_component(value: str, *, field_name: str) -> str: - """Convert an external identifier into a safe path component.""" - normalized = _SAFE_COMPONENT_PATTERN.sub("_", value.strip()).strip("._") - if not normalized: - raise ValueError(f"{field_name} must contain at least one safe character") - return normalized - - -def _collision_safe_component(value: str, *, field_name: str) -> str: - """Add a digest when sanitization could cause path collisions.""" - stripped = value.strip() - normalized = _safe_component(stripped, field_name=field_name) - if normalized == stripped: - return normalized - digest = hashlib.sha256(stripped.encode("utf-8")).hexdigest()[:12] - return f"{normalized}-{digest}" - - -@dataclass(frozen=True) -class AdvancedMemoryPaths: - """Build all disk paths for long-term and session memory.""" - - config: AdvancedMemoryConfig - - @property - def memory_dir(self) -> Path: - """Return the long-term memory directory.""" - return self.config.root_dir / self.config.memory_dir_name - - @property - def session_root_dir(self) -> Path: - """Return the root directory for session memory.""" - return self.config.root_dir / self.config.session_dir_name - - @property - def memory_index_path(self) -> Path: - """Return the long-term memory index path.""" - return self.memory_dir / self.config.memory_index_name - - def memory_topic_path(self, topic_name: str) -> Path: - """Return a safe path for a long-term memory topic.""" - safe_name = _collision_safe_component(topic_name, field_name="topic_name") - if not safe_name.lower().endswith(".md"): - safe_name = f"{safe_name}.md" - if safe_name == self.config.memory_index_name: - raise ValueError("Topic file cannot overwrite the memory index") - return self.memory_dir / safe_name - - def session_dir(self, session_id: str) -> Path: - """Return the isolated storage directory for a session.""" - return self.session_root_dir / _collision_safe_component( - session_id, - field_name="session_id", - ) - - def transcript_path(self, session_id: str) -> Path: - """Return the transcript path for a session.""" - return self.session_dir(session_id) / self.config.transcript_name - - def session_memory_path(self, session_id: str) -> Path: - """Return the session memory path for a session.""" - return self.session_dir(session_id) / self.config.session_memory_name - - def tool_results_dir(self, session_id: str) -> Path: - """Return the large tool-result directory for a session.""" - return self.session_dir(session_id) / "tool-results" - - def tool_result_path(self, session_id: str, result_id: str) -> Path: - """Return a safe JSON path for a large tool result.""" - safe_result_id = _collision_safe_component(result_id, field_name="result_id") - return self.tool_results_dir(session_id) / f"{safe_result_id}.json" - - def ensure_base_directories(self) -> None: - """Create the long-term and session memory directories.""" - self.memory_dir.mkdir(parents=True, exist_ok=True) - self.session_root_dir.mkdir(parents=True, exist_ok=True) - - def ensure_session_directory(self, session_id: str) -> Path: - """Create and return a session's storage directory.""" - path = self.session_dir(session_id) - path.mkdir(parents=True, exist_ok=True) - return path diff --git a/trpc_agent_sdk/advanced_memory/_preload_memory.py b/trpc_agent_sdk/advanced_memory/_preload_memory.py index 4f7f92eee..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 @@ -228,18 +227,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 +247,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 +272,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..f49479063 --- /dev/null +++ b/trpc_agent_sdk/advanced_memory/_redis_stores.py @@ -0,0 +1,5 @@ +"""Redis stores owned by long-term Advanced Memory.""" + +from trpc_agent_sdk.sessions.compact._redis_stores import RedisLongTermMemoryStore + +__all__ = ["RedisLongTermMemoryStore"] diff --git a/trpc_agent_sdk/advanced_memory/_runtime.py b/trpc_agent_sdk/advanced_memory/_runtime.py deleted file mode 100644 index c26def35f..000000000 --- a/trpc_agent_sdk/advanced_memory/_runtime.py +++ /dev/null @@ -1,53 +0,0 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. -# -# Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Unified runtime entry point for the independent memory mechanism.""" - -from __future__ import annotations - -from dataclasses import dataclass - -from ._config import AdvancedMemoryConfig -from ._coordination import SessionOperationCoordinator -from ._paths import AdvancedMemoryPaths -from ._storage import LongTermMemoryStore -from ._storage import SessionMemoryStore -from ._storage import ToolResultStore -from ._storage import TranscriptStore - - -@dataclass(frozen=True) -class AdvancedMemoryRuntime: - """Aggregate configuration, paths, and the three storage objects.""" - - config: AdvancedMemoryConfig - paths: AdvancedMemoryPaths - coordination: SessionOperationCoordinator - long_term_memory: LongTermMemoryStore - session_memory: SessionMemoryStore - tool_results: ToolResultStore - transcripts: TranscriptStore - - @classmethod - def create(cls, config: AdvancedMemoryConfig | None = None) -> "AdvancedMemoryRuntime": - """Create a runtime isolated from the legacy mechanism.""" - resolved_config = config or AdvancedMemoryConfig() - paths = AdvancedMemoryPaths(resolved_config) - return cls( - config=resolved_config, - paths=paths, - coordination=SessionOperationCoordinator(), - long_term_memory=LongTermMemoryStore(resolved_config, paths), - session_memory=SessionMemoryStore(resolved_config, paths), - tool_results=ToolResultStore(resolved_config, paths), - transcripts=TranscriptStore(resolved_config, paths), - ) - - async def initialize(self) -> bool: - """Create memory directories only when the mechanism is enabled.""" - if not self.config.enabled: - return False - await self.long_term_memory.initialize() - return True diff --git a/trpc_agent_sdk/advanced_memory/_sql_stores.py b/trpc_agent_sdk/advanced_memory/_sql_stores.py new file mode 100644 index 000000000..d9173b3ce --- /dev/null +++ b/trpc_agent_sdk/advanced_memory/_sql_stores.py @@ -0,0 +1,5 @@ +"""SQL stores owned by long-term Advanced Memory.""" + +from trpc_agent_sdk.sessions.compact._sql_stores import SqlLongTermMemoryStore + +__all__ = ["SqlLongTermMemoryStore"] diff --git a/trpc_agent_sdk/advanced_memory/_storage.py b/trpc_agent_sdk/advanced_memory/_storage.py index 1fb43591c..0e0a957ae 100644 --- a/trpc_agent_sdk/advanced_memory/_storage.py +++ b/trpc_agent_sdk/advanced_memory/_storage.py @@ -1,331 +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 tempfile -import threading -from collections.abc import Mapping -from dataclasses import replace -from datetime import datetime -from datetime import timezone -from pathlib import Path -from typing import Any - -from ._config import AdvancedMemoryConfig -from ._formats import MemoryDocument -from ._formats import MemoryIndexEntry -from ._formats import SessionMemoryDocument -from ._paths import AdvancedMemoryPaths - - -def _atomic_write_text(path: Path, content: str, *, encoding: str) -> None: - """Atomically replace a text file using a temporary sibling file.""" - path.parent.mkdir(parents=True, exist_ok=True) - file_descriptor, temporary_name = tempfile.mkstemp(prefix=f".{path.name}.", dir=path.parent) - try: - with os.fdopen(file_descriptor, "w", encoding=encoding) as temporary_file: - temporary_file.write(content) - temporary_file.flush() - os.fsync(temporary_file.fileno()) - os.replace(temporary_name, path) - except BaseException: - try: - os.unlink(temporary_name) - except FileNotFoundError: - pass - raise - - -class LongTermMemoryStore: - """Manage MEMORY.md and its detail files in the same directory.""" - - def __init__(self, config: AdvancedMemoryConfig, paths: AdvancedMemoryPaths | None = None) -> None: - """Initialize long-term storage without changing legacy memory.""" - self._config = config - self._paths = paths or AdvancedMemoryPaths(config) - - @property - def index_path(self) -> Path: - """Return the disk path for MEMORY.md.""" - return self._paths.memory_index_path - - async def initialize(self) -> None: - """Create the memory directory and an empty index.""" - await asyncio.to_thread(self._initialize_sync) - - def _initialize_sync(self) -> None: - """Synchronously create the memory directory and empty index.""" - self._paths.ensure_base_directories() - if not self.index_path.exists(): - _atomic_write_text(self.index_path, "", encoding=self._config.encoding) - - async def read_index(self) -> str: - """Read only the configured prefix of MEMORY.md.""" - return await asyncio.to_thread(self._read_index_sync) - - def _read_index_sync(self) -> str: - """Synchronously read MEMORY.md within configured limits.""" - if not self.index_path.exists(): - return "" - with self.index_path.open("r", encoding=self._config.encoding) as index_file: - lines: list[str] = [] - used_bytes = 0 - for _ in range(self._config.memory_index_max_lines): - line = index_file.readline() - if not line: - break - line_bytes = len(line.encode(self._config.encoding)) - if used_bytes + line_bytes > self._config.memory_index_max_bytes: - break - lines.append(line) - used_bytes += line_bytes - return "".join(lines) - - async def write_index(self, entries: list[MemoryIndexEntry]) -> None: - """Atomically write MEMORY.md in the standard index format.""" - content = "\n".join(entry.to_markdown() for entry in entries) - if content: - content += "\n" - await asyncio.to_thread(self._write_index_sync, content) - - def _write_index_sync(self, content: str) -> None: - """Synchronously write MEMORY.md; read_index applies prompt-size limits.""" - _atomic_write_text(self.index_path, content, encoding=self._config.encoding) - - async def read_topic(self, topic_name: str) -> str | None: - """Read a detail memory topic, returning None if absent.""" - path = self._paths.memory_topic_path(topic_name) - return await asyncio.to_thread(self._read_optional_text, path) - - async def read_topic_frontmatter(self, topic_name: str) -> str | None: - """Read only the frontmatter of a detail memory topic.""" - path = self._paths.memory_topic_path(topic_name) - return await asyncio.to_thread(self._read_frontmatter, path) - - def _read_optional_text(self, path: Path) -> str | None: - """Synchronously read an optional text file.""" - if not path.exists(): - return None - return path.read_text(encoding=self._config.encoding) - - def _read_frontmatter(self, path: Path) -> str | None: - """Synchronously read a topic's bounded frontmatter block.""" - if not path.exists(): - return None - lines: list[str] = [] - with path.open(encoding=self._config.encoding) as file: - for line in file: - lines.append(line) - if len(lines) > 1 and line.rstrip("\r\n") == "---": - break - return "".join(lines) - - async def write_topic(self, topic_name: str, document: MemoryDocument) -> Path: - """Atomically write a detail memory file with frontmatter.""" - path = self._paths.memory_topic_path(topic_name) - document = replace(document, updated_at=datetime.now(timezone.utc)) - await asyncio.to_thread( - _atomic_write_text, - path, - document.to_markdown(), - encoding=self._config.encoding, - ) - return path - - async def list_topics(self) -> list[Path]: - """List detail memory files by name, excluding MEMORY.md.""" - return await asyncio.to_thread(self._list_topics_sync) - - def _list_topics_sync(self) -> list[Path]: - """Synchronously list all detail memory files.""" - if not self._paths.memory_dir.exists(): - return [] - return sorted( - (path for path in self._paths.memory_dir.glob("*.md") if path.name != self._config.memory_index_name), - key=lambda path: path.name, - ) - - -class SessionMemoryStore: - """Manage an isolated structured Markdown summary per session.""" - - def __init__(self, config: AdvancedMemoryConfig, paths: AdvancedMemoryPaths | None = None) -> None: - """Initialize session memory storage.""" - self._config = config - self._paths = paths or AdvancedMemoryPaths(config) - - async def read(self, session_id: str) -> str | None: - """Read session memory, returning None if absent.""" - path = self._paths.session_memory_path(session_id) - return await asyncio.to_thread(self._read_sync, path) - - def _read_sync(self, path: Path) -> str | None: - """Synchronously read session memory.""" - if not path.exists(): - return None - return path.read_text(encoding=self._config.encoding) - - async def write(self, session_id: str, document: SessionMemoryDocument) -> Path: - """Atomically write session memory using the fixed section template.""" - path = self._paths.session_memory_path(session_id) - await asyncio.to_thread( - _atomic_write_text, - path, - document.to_markdown(), - encoding=self._config.encoding, - ) - return path - - -class ToolResultStore: - """Persist complete tool results that exceed the context budget.""" - - def __init__(self, config: AdvancedMemoryConfig, paths: AdvancedMemoryPaths | None = None) -> None: - """Initialize large tool-result storage.""" - self._config = config - self._paths = paths or AdvancedMemoryPaths(config) - - async def write(self, session_id: str, result_id: str, serialized_result: str) -> Path: - """Atomically write a complete tool result and return its disk path.""" - path = self._paths.tool_result_path(session_id, result_id) - await asyncio.to_thread( - _atomic_write_text, - path, - serialized_result, - encoding=self._config.encoding, - ) - return path - - async def read(self, session_id: str, result_id: str) -> str | None: - """Read a persisted complete tool result.""" - path = self._paths.tool_result_path(session_id, result_id) - return await asyncio.to_thread(self._read_sync, path) - - def _read_sync(self, path: Path) -> str | None: - """Synchronously read an optional complete tool-result file.""" - if not path.exists(): - return None - return path.read_text(encoding=self._config.encoding) - - -class TranscriptStore: - """Store complete per-session records as append-only JSONL.""" - - def __init__(self, config: AdvancedMemoryConfig, paths: AdvancedMemoryPaths | None = None) -> None: - """Initialize transcript storage and its process-local write lock.""" - self._config = config - self._paths = paths or AdvancedMemoryPaths(config) - self._write_lock = threading.Lock() - self._seen_unique_values: dict[tuple[Path, str], set[str]] = {} - - async def append(self, session_id: str, record: Mapping[str, Any]) -> Path: - """Append one JSON-serializable record to a session transcript.""" - path = self._paths.transcript_path(session_id) - payload = dict(record) - payload.setdefault("recorded_at", datetime.now(timezone.utc).isoformat()) - serialized = json.dumps(payload, ensure_ascii=False, separators=(",", ":"), default=str) - await asyncio.to_thread(self._append_sync, path, serialized) - return path - - def _append_sync(self, path: Path, serialized: str) -> None: - """Synchronously append one transcript line under the write lock.""" - path.parent.mkdir(parents=True, exist_ok=True) - with self._write_lock: - self._append_serialized_unlocked(path, serialized) - - def _append_serialized_unlocked(self, path: Path, serialized: str) -> None: - """Append one serialized line while the caller holds the lock.""" - with path.open("a", encoding=self._config.encoding) as transcript_file: - transcript_file.write(serialized) - transcript_file.write("\n") - transcript_file.flush() - if self._config.transcript_fsync: - os.fsync(transcript_file.fileno()) - - async def append_unique( - self, - session_id: str, - record: Mapping[str, Any], - *, - unique_key: str, - ) -> tuple[Path, bool]: - """Append a transcript record after de-duplicating by a field.""" - path = self._paths.transcript_path(session_id) - payload = dict(record) - unique_value = payload.get(unique_key) - if not isinstance(unique_value, str) or not unique_value: - raise ValueError(f"Transcript unique key {unique_key!r} must be a non-empty string") - payload.setdefault("recorded_at", datetime.now(timezone.utc).isoformat()) - serialized = json.dumps(payload, ensure_ascii=False, separators=(",", ":"), default=str) - appended = await asyncio.to_thread( - self._append_unique_sync, - path, - serialized, - unique_key, - unique_value, - ) - return path, appended - - def _append_unique_sync( - self, - path: Path, - serialized: str, - unique_key: str, - unique_value: str, - ) -> bool: - """Load de-duplication state and append only new records.""" - path.parent.mkdir(parents=True, exist_ok=True) - cache_key = (path, unique_key) - with self._write_lock: - seen_values = self._seen_unique_values.get(cache_key) - if seen_values is None: - seen_values = self._load_unique_values_unlocked(path, unique_key) - self._seen_unique_values[cache_key] = seen_values - if unique_value in seen_values: - return False - self._append_serialized_unlocked(path, serialized) - seen_values.add(unique_value) - return True - - def _load_unique_values_unlocked(self, path: Path, unique_key: str) -> set[str]: - """Load existing de-duplication values while holding the lock.""" - if not path.exists(): - return set() - values: set[str] = set() - with path.open("r", encoding=self._config.encoding) as transcript_file: - for line in transcript_file: - if not line.strip(): - continue - parsed = json.loads(line) - if isinstance(parsed, dict) and isinstance(parsed.get(unique_key), str): - values.add(parsed[unique_key]) - return values - - async def read_all(self, session_id: str) -> list[dict[str, Any]]: - """Read all transcript records for a session in write order.""" - path = self._paths.transcript_path(session_id) - return await asyncio.to_thread(self._read_all_sync, path) - - def _read_all_sync(self, path: Path) -> list[dict[str, Any]]: - """Parse a consistent transcript snapshot under the file lock.""" - with self._write_lock: - if not path.exists(): - return [] - records: list[dict[str, Any]] = [] - with path.open("r", encoding=self._config.encoding) as transcript_file: - for line_number, line in enumerate(transcript_file, start=1): - if not line.strip(): - continue - parsed = json.loads(line) - if not isinstance(parsed, dict): - raise ValueError(f"Transcript line {line_number} is not a JSON object") - records.append(parsed) - return records +__all__ = ["LongTermMemoryStore"] 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..19a5b2729 --- /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 trpc_agent_sdk.sessions.compact._paths import MemoryScope +from trpc_agent_sdk.sessions.compact._runtime import ScopedAdvancedMemoryRuntime + + +class AdvancedMemoryStorageBackend(Protocol): + """Create storage views isolated to an application user.""" + + def for_scope(self, scope: MemoryScope) -> ScopedAdvancedMemoryRuntime: + """Return the tenant-bound storage view.""" + + +class LocalAdvancedMemoryStorageBackend: + """Adapt the file-backed runtime to the storage backend boundary.""" + + def __init__(self, runtime: object) -> None: + self._runtime = runtime + + def for_scope(self, scope: MemoryScope) -> ScopedAdvancedMemoryRuntime: + """Return a file-backed scope without exposing local path mechanics.""" + return self._runtime.for_scope(scope.app_name, scope.user_id) diff --git a/trpc_agent_sdk/evaluation/_eval_session_service.py b/trpc_agent_sdk/evaluation/_eval_session_service.py index 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 8cc2c97f8..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. - """ - return None + """Release service-owned local or external storage resources.""" + await self._runtime.close() diff --git a/trpc_agent_sdk/runners.py b/trpc_agent_sdk/runners.py index 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 feb14fb14..000000000 --- a/trpc_agent_sdk/sessions/_advanced_memory_session_service.py +++ /dev/null @@ -1,403 +0,0 @@ -# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. -# -# Copyright (C) 2026 Tencent. All rights reserved. -# -# tRPC-Agent-Python is licensed under Apache-2.0. -"""Session service backed by Advanced Memory transcript storage.""" - -from __future__ import annotations - -import asyncio -import json -import os -import shutil -import tempfile -import time -import uuid -from pathlib import Path -from typing import Any -from typing import Optional - -from trpc_agent_sdk.abc import ListSessionsResponse -from trpc_agent_sdk.context import AgentContext -from trpc_agent_sdk.context import InvocationContext -from trpc_agent_sdk.events import Event -from trpc_agent_sdk.advanced_memory import AdvancedMemoryConfig -from trpc_agent_sdk.advanced_memory import AdvancedMemoryRuntime -from trpc_agent_sdk.advanced_memory._coordination import CrossLoopLock -from trpc_agent_sdk.advanced_memory._transcript import build_event_transcript_record -from trpc_agent_sdk.advanced_memory._transcript import find_last_event_id - -from ._base_session_service import BaseSessionService -from ._session import Session -from ._types import SessionServiceConfig -from ._utils import extract_state_delta -from ._utils import merge_state - - -class _AdvancedMemorySessionBackend(BaseSessionService): - """Persist Session metadata while TranscriptSessionService persists events.""" - - def __init__(self, runtime: AdvancedMemoryRuntime, session_config: SessionServiceConfig | None = None) -> None: - super().__init__(session_config=session_config) - self._runtime = runtime - self._lock = CrossLoopLock() - self._cleanup_task: asyncio.Task[None] | None = None - self._cleanup_stop_event: asyncio.Event | None = None - self._transcript_enabled = True - self._start_cleanup_task() - - def set_transcript_enabled(self, enabled: bool) -> None: - """Enable or disable transcript persistence for this backend.""" - self._transcript_enabled = enabled - - def _metadata_path(self, session_id: str) -> Path: - return self._runtime.paths.session_dir(session_id) / "session.json" - - @property - def _state_path(self) -> Path: - return self._runtime.paths.session_root_dir / "_state.json" - - async def _write_session(self, session: Session) -> None: - payload = session.model_dump(mode="json", by_alias=True, exclude={"events", "historical_events"}) - payload["state"] = extract_state_delta(session.state).session_state - path = self._metadata_path(session.id) - await asyncio.to_thread(self._write_json, path, payload, self._runtime.config.encoding) - - @staticmethod - def _write_json(path: Path, payload: dict[str, Any], encoding: str) -> None: - path.parent.mkdir(parents=True, exist_ok=True) - file_descriptor, temporary_name = tempfile.mkstemp(prefix=f".{path.name}.", dir=path.parent) - try: - with os.fdopen(file_descriptor, "w", encoding=encoding) as temporary_file: - temporary_file.write(json.dumps(payload, ensure_ascii=False, separators=(",", ":"))) - temporary_file.flush() - os.fsync(temporary_file.fileno()) - os.replace(temporary_name, path) - except BaseException: - try: - os.unlink(temporary_name) - except FileNotFoundError: - pass - raise - - async def _read_session(self, session_id: str) -> Session | None: - path = self._metadata_path(session_id) - if not path.exists(): - return None - payload = await asyncio.to_thread(path.read_text, encoding=self._runtime.config.encoding) - await asyncio.to_thread(path.touch) - return Session.model_validate(json.loads(payload)) - - def _start_cleanup_task(self) -> None: - """Start persistent session cleanup when TTL is enabled.""" - if not self.session_config.need_ttl_expire() or self._cleanup_task is not None: - return - try: - loop = asyncio.get_running_loop() - except RuntimeError: - return - self._cleanup_stop_event = asyncio.Event() - self._cleanup_task = loop.create_task(self._cleanup_loop()) - - async def _cleanup_loop(self) -> None: - """Periodically remove expired session directories.""" - assert self._cleanup_stop_event is not None - try: - while not self._cleanup_stop_event.is_set(): - try: - await asyncio.wait_for( - self._cleanup_stop_event.wait(), - timeout=self.session_config.ttl.cleanup_interval_seconds, - ) - break - except asyncio.TimeoutError: - async with self._lock: - await asyncio.to_thread(self._cleanup_expired_sessions) - except asyncio.CancelledError: - raise - - def _cleanup_expired_sessions(self) -> None: - """Delete session directories idle longer than the configured TTL.""" - cutoff = time.time() - self.session_config.ttl.ttl_seconds - root = self._runtime.paths.session_root_dir - if not root.exists(): - return - for metadata_path in root.glob("*/session.json"): - try: - if metadata_path.stat().st_mtime < cutoff: - shutil.rmtree(metadata_path.parent, ignore_errors=True) - except FileNotFoundError: - continue - - async def _stop_cleanup_task(self) -> None: - """Stop the background TTL cleanup task.""" - task = self._cleanup_task - self._cleanup_task = None - if task is None: - return - if self._cleanup_stop_event is not None: - self._cleanup_stop_event.set() - task.cancel() - await asyncio.gather(task, return_exceptions=True) - self._cleanup_stop_event = None - - async def _read_global_state(self) -> dict[str, dict[str, Any]]: - if not self._state_path.exists(): - return {"app": {}, "user": {}} - payload = await asyncio.to_thread( - self._state_path.read_text, - encoding=self._runtime.config.encoding, - ) - parsed = json.loads(payload) - return { - "app": dict(parsed.get("app", {})), - "user": dict(parsed.get("user", {})), - } - - async def _write_global_state(self, state: dict[str, dict[str, Any]]) -> None: - await asyncio.to_thread(self._write_json, self._state_path, state, self._runtime.config.encoding) - - async def _restore_events(self, session: Session) -> Session: - records = await self._runtime.transcripts.read_all(session.id) - events: list[Event] = [] - for record in records: - event_payload = record.get("event") - if record.get("kind") != "event" or not isinstance(event_payload, dict): - continue - events.append(Event.model_validate(event_payload)) - session.events = events - if events: - session.last_update_time = events[-1].timestamp - return session - - async def create_session( - self, - *, - app_name: str, - user_id: str, - state: Optional[dict[str, Any]] = None, - session_id: Optional[str] = None, - agent_context: Optional[AgentContext] = None, - ) -> Session: - self._start_cleanup_task() - resolved_id = session_id.strip() if session_id and session_id.strip() else str(uuid.uuid4()) - state_delta = extract_state_delta(state) - session = Session( - id=resolved_id, - app_name=app_name, - user_id=user_id, - state=state_delta.session_state, - save_key=f"{app_name}/{user_id}", - ) - async with self._lock: - await self._runtime.initialize() - existing = await self._read_session(resolved_id) - if existing is not None and (existing.app_name != app_name or existing.user_id != user_id): - raise ValueError(f"Session ID {resolved_id!r} is already used by another app or user") - global_state = await self._read_global_state() - global_state["app"].setdefault(app_name, {}).update(state_delta.app_state_delta) - global_state["user"].setdefault(f"{app_name}/{user_id}", {}).update(state_delta.user_state_delta) - await self._write_global_state(global_state) - await self._write_session(session) - session.state = merge_state( - extract_state_delta(session.state), - need_copy=True, - ) - session.state.update({f"app:{key}": value for key, value in global_state["app"].get(app_name, {}).items()}) - session.state.update({ - f"user:{key}": value - for key, value in global_state["user"].get(f"{app_name}/{user_id}", {}).items() - }) - return session - - async def get_session( - self, - *, - app_name: str, - user_id: str, - session_id: str, - agent_context: Optional[AgentContext] = None, - ) -> Session | None: - self._start_cleanup_task() - async with self._lock: - session = await self._read_session(session_id) - if session is None or session.app_name != app_name or session.user_id != user_id: - return None - global_state = await self._read_global_state() - app_state = global_state["app"].get(app_name, {}) - user_state = global_state["user"].get(f"{app_name}/{user_id}", {}) - session.state = merge_state( - extract_state_delta(session.state), - need_copy=True, - ) - session.state.update({f"app:{key}": value for key, value in app_state.items()}) - session.state.update({f"user:{key}": value for key, value in user_state.items()}) - return self.filter_events(await self._restore_events(session), need_copy=True) - - async def list_sessions( - self, - *, - app_name: str, - user_id: Optional[str] = None, - ) -> ListSessionsResponse: - self._start_cleanup_task() - if not self._runtime.paths.session_root_dir.exists(): - return ListSessionsResponse() - sessions: list[Session] = [] - for path in await asyncio.to_thread(lambda: list(self._runtime.paths.session_root_dir.glob("*/session.json"))): - try: - session = await asyncio.to_thread(lambda path=path: Session.model_validate( - json.loads(path.read_text(encoding=self._runtime.config.encoding)))) - except (OSError, ValueError, TypeError): - continue - if session.app_name == app_name and (user_id is None or session.user_id == user_id): - session.events = [] - session.historical_events = [] - sessions.append(session) - return ListSessionsResponse(sessions=sessions) - - async def delete_session(self, *, app_name: str, user_id: str, session_id: str) -> None: - self._start_cleanup_task() - session = await self.get_session( - app_name=app_name, - user_id=user_id, - session_id=session_id, - ) - if session is not None: - async with self._lock: - await asyncio.to_thread(shutil.rmtree, self._runtime.paths.session_dir(session_id), True) - - async def append_event(self, session: Session, event: Event) -> Event: - self._start_cleanup_task() - async with self._lock: - persisted = await super().append_event(session, event) - if not event.partial: - state_delta = extract_state_delta(event.actions.state_delta if event.actions else None) - if state_delta.app_state_delta or state_delta.user_state_delta: - global_state = await self._read_global_state() - global_state["app"].setdefault(session.app_name, {}).update(state_delta.app_state_delta) - global_state["user"].setdefault(f"{session.app_name}/{session.user_id}", - {}).update(state_delta.user_state_delta) - await self._write_global_state(global_state) - session.state.update({f"app:{key}": value for key, value in state_delta.app_state_delta.items()}) - session.state.update({f"user:{key}": value for key, value in state_delta.user_state_delta.items()}) - await self._write_session(session) - if not event.partial and self._transcript_enabled: - records = await self._runtime.transcripts.read_all(session.id) - record = build_event_transcript_record( - session, - persisted, - parent_event_id=find_last_event_id(records), - ) - await self._runtime.transcripts.append_unique( - session.id, - record, - unique_key="event_id", - ) - return persisted - - async def update_session(self, session: Session) -> None: - self._start_cleanup_task() - async with self._lock: - await self._write_session(session) - - async def create_session_summary( - self, - session: Session, - ctx: InvocationContext | None = None, - ) -> None: - await super().create_session_summary(session, ctx=ctx) - await self.update_session(session) - - async def get_session_summary(self, session: Session) -> str | None: - return await super().get_session_summary(session) - - async def close(self) -> None: - await self._stop_cleanup_task() - - -class AdvancedMemorySessionService(BaseSessionService): - """Persist sessions and raw events in the Advanced Memory directory.""" - - def __init__( - self, - runtime: AdvancedMemoryRuntime | None = None, - *, - config: AdvancedMemoryConfig | None = None, - session_config: SessionServiceConfig | None = None, - preload_memory_model: Any | None = None, - ) -> None: - if runtime is not None and config is not None and runtime.config != config: - raise ValueError("runtime and config must describe the same Advanced Memory configuration") - self._runtime = runtime or AdvancedMemoryRuntime.create(config) - self._preload_memory_model = preload_memory_model - self._backend = _AdvancedMemorySessionBackend(self._runtime, session_config=session_config) - self._integration: Any | None = None - self._bound_agent: Any | None = None - super().__init__(session_config=session_config) - - @property - def runtime(self) -> AdvancedMemoryRuntime: - """Return the Advanced Memory runtime used by this service.""" - return self._runtime - - @property - def integration(self) -> Any | None: - """Return the Advanced Memory binding, when attached to a Runner.""" - return self._integration - - @property - def backend(self) -> BaseSessionService: - """Return the persistent backend used by the transcript decorator.""" - return self._backend - - def bind(self, agent: Any) -> BaseSessionService: - """Install Advanced Memory callbacks and return the wrapped service.""" - from trpc_agent_sdk.advanced_memory import setup_advanced_memory - - if self._integration is not None: - if agent is not self._bound_agent: - raise ValueError("AdvancedMemorySessionService is already bound to another agent") - return self._integration.session_service - integration = setup_advanced_memory( - agent, - self, - self._runtime, - preload_memory_model=self._preload_memory_model, - ) - self._backend.set_transcript_enabled(False) - self._integration = integration - self._bound_agent = agent - return self._integration.session_service - - async def create_session(self, **kwargs: Any) -> Session: - return await self._backend.create_session(**kwargs) - - async def get_session(self, **kwargs: Any) -> Session | None: - return await self._backend.get_session(**kwargs) - - async def list_sessions(self, **kwargs: Any) -> ListSessionsResponse: - return await self._backend.list_sessions(**kwargs) - - async def delete_session(self, **kwargs: Any) -> None: - await self._backend.delete_session(**kwargs) - - async def append_event(self, session: Session, event: Event) -> Event: - return await self._backend.append_event(session, event) - - async def update_session(self, session: Session) -> None: - await self._backend.update_session(session) - - async def create_session_summary( - self, - session: Session, - ctx: InvocationContext | None = None, - ) -> None: - await self._backend.create_session_summary(session, ctx=ctx) - - async def get_session_summary(self, session: Session) -> str | None: - return await self._backend.get_session_summary(session) - - async def close(self) -> None: - await self._backend.close() diff --git a/trpc_agent_sdk/sessions/_base_session_service.py b/trpc_agent_sdk/sessions/_base_session_service.py index 979523f46..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 4333ffeb1..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,18 +396,41 @@ 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 + # AsyncSession cannot perform an implicit refresh when an ORM + # attribute is accessed after commit. Keep committed values available + # because this service reads StorageSession state after committing. + kwargs.setdefault("expire_on_commit", False) self._sql_storage = SqlStorage(is_async=is_async, db_url=db_url, metadata=SessionStorageBase.metadata, **kwargs) self.__cleanup_task: Optional[asyncio.Task] = None self.__cleanup_stop_event: Optional[asyncio.Event] = None self._start_cleanup_task() + @property + def db_url(self) -> str: + """Return the configured SQL connection URL.""" + return self._db_url + + @property + def is_async(self) -> bool: + """Return whether this service uses asynchronous SQL sessions.""" + return self._is_async + @override async def create_session( self, @@ -529,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: @@ -543,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 @@ -651,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() @@ -704,6 +763,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 +777,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 +794,10 @@ async def _get_session(self, sql_session: SqlSession, app_name: str, user_id: st storage_session.update_time = func.now() await self._sql_storage.commit(sql_session) + # Assigning a SQL expression expires the server-generated timestamp + # even when expire_on_commit=False. Refresh it before callers access + # update_time outside SQLAlchemy's async greenlet. + await self._sql_storage.refresh(sql_session, storage_session) return storage_session diff --git a/trpc_agent_sdk/sessions/compact/__init__.py b/trpc_agent_sdk/sessions/compact/__init__.py new file mode 100644 index 000000000..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 71% rename from trpc_agent_sdk/advanced_memory/_autocompact.py rename to trpc_agent_sdk/sessions/compact/_autocompact.py index e7efb436a..045d3052d 100644 --- a/trpc_agent_sdk/advanced_memory/_autocompact.py +++ b/trpc_agent_sdk/sessions/compact/_autocompact.py @@ -8,6 +8,7 @@ from __future__ import annotations import asyncio +import copy import hashlib import json import re @@ -18,6 +19,7 @@ from typing import TYPE_CHECKING from trpc_agent_sdk.agents import LlmAgent +from trpc_agent_sdk.events import Event from trpc_agent_sdk.models import LlmResponse from trpc_agent_sdk.runners import Runner from trpc_agent_sdk.sessions import InMemorySessionService @@ -26,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 @@ -35,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 = ( @@ -68,6 +73,8 @@ class AutoCompactRecord: boundary_occurrence: int summary: str source: str + boundary_event_id: str | None = None + compaction_id: str | None = None @dataclass @@ -226,31 +233,45 @@ 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"] = {} @property 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.""" - 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) @@ -269,12 +290,16 @@ 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": 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: @@ -286,11 +311,16 @@ 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.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, @@ -307,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], @@ -372,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( @@ -419,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]) @@ -435,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.""" @@ -482,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, @@ -537,6 +694,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 @@ -556,6 +734,9 @@ async def apply( token_budget_before = tracker.budget(request, ctx) token_mode = token_budget_before.token_mode_enabled request_tokens_before = token_budget_before.estimate.tokens + comparison_tokens_before = ( + tracker.estimate_request_tokens(request) if token_mode else None + ) blocking_reached = (request_tokens_before >= token_budget_before.blocking_threshold_tokens if token_mode else request_chars_before >= config.autocompact_blocking_chars) if state.consecutive_failures >= config.autocompact_max_failures and blocking_reached: @@ -594,38 +775,46 @@ async def apply( original_contents = [content.model_copy(deep=True) for content in request.contents] try: compact_record: AutoCompactRecord | None = None - session_memory = await self._latest_session_memory_record(session_id) + if (self._session_memory_extractor is not None 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( @@ -648,22 +837,27 @@ async def apply( ), boundary_index=boundary_index, source="legacy", + boundary_event_id=self._legacy_boundary_event_id(ctx), ) request_chars_after = estimate_request_chars(request) - if request_chars_after >= request_chars_before: - 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 @@ -675,9 +869,9 @@ async def apply( request_chars_before, request_chars_after, 0, - request_tokens_before=request_tokens_before if token_mode else None, - request_tokens_after=(token_budget_after.estimate.tokens if token_mode else None), - token_source=token_budget_after.estimate.source if token_mode else None, + request_tokens_before=comparison_tokens_before if token_mode else None, + request_tokens_after=comparison_tokens_after if token_mode else None, + token_source="estimated" if token_mode else None, ) except Exception as exc: # noqa: BLE001 request.contents = original_contents 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 80% rename from trpc_agent_sdk/advanced_memory/_config.py rename to trpc_agent_sdk/sessions/compact/_config.py index 975875d3e..456ff3dcd 100644 --- a/trpc_agent_sdk/advanced_memory/_config.py +++ b/trpc_agent_sdk/sessions/compact/_config.py @@ -12,6 +12,9 @@ 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", @@ -96,11 +99,22 @@ 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) + 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 +123,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 @@ -162,9 +177,40 @@ 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 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, 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 89% rename from trpc_agent_sdk/advanced_memory/_history_snip.py rename to trpc_agent_sdk/sessions/compact/_history_snip.py index 67d8ff1a6..72f959320 100644 --- a/trpc_agent_sdk/advanced_memory/_history_snip.py +++ b/trpc_agent_sdk/sessions/compact/_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/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 86% rename from trpc_agent_sdk/advanced_memory/_microcompact.py rename to trpc_agent_sdk/sessions/compact/_microcompact.py index eeaabdd36..ec1ae94d5 100644 --- a/trpc_agent_sdk/advanced_memory/_microcompact.py +++ b/trpc_agent_sdk/sessions/compact/_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/sessions/compact/_paths.py b/trpc_agent_sdk/sessions/compact/_paths.py new file mode 100644 index 000000000..87a473c17 --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/_paths.py @@ -0,0 +1,208 @@ +# 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 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/sessions/compact/_runtime.py b/trpc_agent_sdk/sessions/compact/_runtime.py new file mode 100644 index 000000000..e0bd6d24f --- /dev/null +++ b/trpc_agent_sdk/sessions/compact/_runtime.py @@ -0,0 +1,258 @@ +# 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 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.""" + + 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 + + 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.""" + 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.""" + 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 + + @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 + + 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) diff --git a/trpc_agent_sdk/advanced_memory/_session_memory.py b/trpc_agent_sdk/sessions/compact/_session_memory.py similarity index 81% rename from trpc_agent_sdk/advanced_memory/_session_memory.py rename to trpc_agent_sdk/sessions/compact/_session_memory.py index 39f9ab454..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]], @@ -607,14 +658,64 @@ 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) + 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_id: str, + session: "SessionABC", included_records: list[dict[str, Any]], document: SessionMemoryDocument, context_tokens: int | None, @@ -634,8 +735,37 @@ async def _persist_checkpoint( document.key_results, document.worklog, ) - await self._runtime.transcripts.append_unique( - session_id, + 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", @@ -661,12 +791,20 @@ 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) - 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( @@ -681,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 @@ -699,7 +835,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 +851,12 @@ 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) + 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.id, + session, included, document, context_tokens, diff --git a/trpc_agent_sdk/advanced_memory/_session_service.py b/trpc_agent_sdk/sessions/compact/_session_service.py similarity index 73% rename from trpc_agent_sdk/advanced_memory/_session_service.py rename to trpc_agent_sdk/sessions/compact/_session_service.py index a8cccd52d..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.""" @@ -73,30 +85,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 +157,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,27 +186,37 @@ 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: """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 87% rename from trpc_agent_sdk/advanced_memory/_tool_result_budget.py rename to trpc_agent_sdk/sessions/compact/_tool_result_budget.py index 73a710aa2..7181584e5 100644 --- a/trpc_agent_sdk/advanced_memory/_tool_result_budget.py +++ b/trpc_agent_sdk/sessions/compact/_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,17 @@ 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( + 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, @@ -217,8 +231,8 @@ 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.", - "path": str(persisted_path), + "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, @@ -281,11 +295,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 +316,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 +342,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 +430,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/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/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..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 @@ -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] = [] @@ -43,9 +48,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 +71,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 +96,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 +110,13 @@ async def save_memory( memory_type=resolved_type, content=content, ) - async with self._index_lock: - path = await self._runtime.long_term_memory.write_topic( + runtime = self._runtime_for_context(tool_context) + async with self._index_lock(runtime): + path = await runtime.long_term_memory.write_topic( filename, document, ) - entries = _parse_index(await self._runtime.long_term_memory.read_index()) + entries = _parse_index(await runtime.long_term_memory.read_index()) new_entry = MemoryIndexEntry( name=name, filename=path.name, @@ -100,19 +124,19 @@ async def save_memory( ) entries = [entry for entry in entries if entry.filename != new_entry.filename] entries.insert(0, new_entry) - await self._runtime.long_term_memory.write_index(entries) - updated_at = parse_memory_updated_at(await self._runtime.long_term_memory.read_topic(filename) or "") + await runtime.long_term_memory.write_index(entries) + updated_at = parse_memory_updated_at(await runtime.long_term_memory.read_topic(filename) or "") return { "saved": True, "filename": path.name, - "path": str(path), + "path": runtime.paths.storage_reference("memory_topic", topic_name=path.name), "memory_type": resolved_type.value, "updated_at": updated_at.isoformat() if updated_at is not None else None, } - async def read_memory(self, filename: str) -> dict: + async def read_memory(self, filename: str, tool_context: Any | None = None) -> dict: """Read a complete long-term memory by its filename in MEMORY.md.""" - content = await self._runtime.long_term_memory.read_topic(filename) + content = await self._runtime_for_context(tool_context).long_term_memory.read_topic(filename) if content is None: return {"found": False, "filename": filename} updated_at = parse_memory_updated_at(content) @@ -133,11 +157,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: - """Return the current long-term memory index and its disk path.""" + async def list_memory_index(self, tool_context: Any | None = None) -> dict: + """Return the current long-term memory index and its storage reference.""" + runtime = self._runtime_for_context(tool_context) return { - "index_path": str(self._runtime.paths.memory_index_path), - "index": await self._runtime.long_term_memory.read_index(), + "index_path": _memory_index_reference(runtime), + "index": await runtime.long_term_memory.read_index(), }