diff --git a/pyproject.toml b/pyproject.toml index db0e78283..05f68bd2b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -123,6 +123,7 @@ dev = [ "numpy>=2.2.5", "pytest", "pytest-asyncio", + "fakeredis>=2.31.0", "black", "flake8", "yapf", diff --git a/requirements-test.txt b/requirements-test.txt index c54dad067..0b9d6115f 100644 --- a/requirements-test.txt +++ b/requirements-test.txt @@ -17,6 +17,7 @@ rich>=13.0.0 greenlet aiosqlite redis>=6.2.0 +fakeredis>=2.31.0 langgraph google-genai>=1.24.0 diff --git a/tests/memory/test_redis_memory_service.py b/tests/memory/test_redis_memory_service.py index 3a48f3bce..b23c33b77 100644 --- a/tests/memory/test_redis_memory_service.py +++ b/tests/memory/test_redis_memory_service.py @@ -16,16 +16,14 @@ import time from contextlib import asynccontextmanager -from typing import Any, Optional +from typing import Optional from unittest.mock import AsyncMock, MagicMock, patch -import pytest - from trpc_agent_sdk.abc import MemoryServiceConfig from trpc_agent_sdk.events import Event from trpc_agent_sdk.memory._redis_memory_service import RedisMemoryService from trpc_agent_sdk.sessions import Session -from trpc_agent_sdk.types import Content, Part, SearchMemoryResponse, Ttl +from trpc_agent_sdk.types import Content, Part, SearchMemoryResponse # --------------------------------------------------------------------------- @@ -101,6 +99,17 @@ def test_default_init(self, MockRedisStorage): svc = RedisMemoryService(db_url="redis://localhost", memory_service_config=_make_config_no_ttl()) assert svc.enabled is True MockRedisStorage.assert_called_once() + assert MockRedisStorage.call_args.kwargs["decode_responses"] is True + + @patch("trpc_agent_sdk.memory._redis_memory_service.RedisStorage") + def test_explicit_bytes_responses_are_preserved(self, MockRedisStorage): + MockRedisStorage.return_value = MagicMock() + RedisMemoryService( + db_url="redis://localhost", + decode_responses=False, + memory_service_config=_make_config_no_ttl(), + ) + assert MockRedisStorage.call_args.kwargs["decode_responses"] is False @patch("trpc_agent_sdk.memory._redis_memory_service.RedisStorage") def test_passes_is_async(self, MockRedisStorage): diff --git a/tests/sessions/replay_cases.py b/tests/sessions/replay_cases.py new file mode 100644 index 000000000..9303ef8a8 --- /dev/null +++ b/tests/sessions/replay_cases.py @@ -0,0 +1,176 @@ +# 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. +"""Standard input traces for Session/Memory/Summary replay tests.""" + +from __future__ import annotations + +from dataclasses import dataclass +from dataclasses import field +from typing import Any + + +@dataclass(frozen=True) +class ReplayOperation: + """One backend-independent operation in a replay trace.""" + + kind: str + payload: dict[str, Any] = field(default_factory=dict) + + +@dataclass(frozen=True) +class ReplayCase: + """A complete trace and the memory queries evaluated afterwards.""" + + case_id: str + description: str + operations: tuple[ReplayOperation, ...] + memory_queries: tuple[str, ...] = () + + +def _text(author: str, text: str, **extra: Any) -> ReplayOperation: + return ReplayOperation("event", {"event_type": "text", "author": author, "text": text, **extra}) + + +def _state(delta: dict[str, Any]) -> ReplayOperation: + return ReplayOperation("event", {"event_type": "state", "author": "agent", "state_delta": delta}) + + +def _dialogue(prefix: str, turns: int) -> tuple[ReplayOperation, ...]: + operations: list[ReplayOperation] = [] + for turn in range(1, turns + 1): + operations.extend((_text("user", + f"{prefix} user turn {turn}"), _text("assistant", f"{prefix} assistant turn {turn}"))) + return tuple(operations) + + +REPLAY_CASES: tuple[ReplayCase, ...] = ( + ReplayCase( + case_id="single_turn", + description="single user/assistant exchange", + operations=(_text("user", "hello replay"), _text("assistant", "hello user")), + ), + ReplayCase( + case_id="multi_turn", + description="three consecutive conversation turns", + operations=_dialogue("multi", 3), + ), + ReplayCase( + case_id="tool_call", + description="function call followed by function response", + operations=( + _text("user", "search for replay consistency"), + ReplayOperation( + "event", { + "event_type": "function_call", + "author": "assistant", + "name": "search", + "call_id": "search-call-1", + "args": { + "query": "replay consistency" + }, + }), + ReplayOperation( + "event", { + "event_type": "function_response", + "author": "tool", + "name": "search", + "call_id": "search-call-1", + "response": { + "result": "consistent" + }, + }), + _text("assistant", "the search result is consistent"), + ), + ), + ReplayCase( + case_id="state_overwrite", + description="repeated session state writes and overwrite", + operations=( + _state({ + "phase": "draft", + "counter": 1 + }), + _state({ + "phase": "review", + "counter": 2 + }), + _state({ + "phase": "done", + "counter": 3 + }), + ), + ), + ReplayCase( + case_id="scoped_state", + description="application, user, session, and temporary state", + operations=( + _state({ + "app:release": "2026", + "user:language": "python", + "session_flag": "active" + }), + _state({ + "user:language": "go", + "temp:request_id": "ephemeral", + "session_flag": "complete" + }), + ), + ), + ReplayCase( + case_id="memory_roundtrip", + description="store and retrieve one user preference memory", + operations=( + _text("user", "preference-token-oolong means I prefer oolong tea"), + _text("assistant", "I will remember that preference"), + ReplayOperation("store_memory"), + ), + memory_queries=("preference-token-oolong", ), + ), + ReplayCase( + case_id="summary_create", + description="create a deterministic summary for a long dialogue", + operations=(*_dialogue("summary-create", 4), + ReplayOperation("summarize", {"summary_text": "summary version one"})), + ), + ReplayCase( + case_id="summary_update", + description="preserve an old summary on failure, then replace it", + operations=(*_dialogue("summary-update-initial", 4), + ReplayOperation("summarize", {"summary_text": "summary update version one"}), + *_dialogue("summary-update-later", 3), ReplayOperation("summarize_failure"), + ReplayOperation("summarize", {"summary_text": "summary update version two"})), + ), + ReplayCase( + case_id="summary_truncation", + description="summary, retained events, and new events restore context together", + operations=(*_dialogue("truncation-old", 5), + ReplayOperation("summarize", {"summary_text": "compressed historical context"}), + _text("user", "follow-up after compression"), _text("assistant", "answer after compression")), + ), + ReplayCase( + case_id="partial_retry", + description="interrupted write leaves no dirty state and retry is stored once", + operations=( + _text("user", "produce a streamed answer"), + _text( + "assistant", + "unfinished", + event_id="retry-event", + partial=True, + state_delta={ + "recovery_status": "dirty", + "temp:retry_buffer": "unfinished", + }, + ), + _text( + "assistant", + "finished answer", + event_id="retry-event", + state_delta={"recovery_status": "complete"}, + ), + ), + ), +) diff --git a/tests/sessions/replay_consistency_design.md b/tests/sessions/replay_consistency_design.md new file mode 100644 index 000000000..a0a862069 --- /dev/null +++ b/tests/sessions/replay_consistency_design.md @@ -0,0 +1,3 @@ +# Replay 一致性设计说明 + +框架用同一轨迹驱动 InMemory 和 SQLite 基线,真实 SQL/Redis 由环境变量接入。Redis 无服务时以 fakeredis 客户端注入正式 RedisStorage,不在 harness 内重写语义。快照只归一化自动 ID、动态时间、文本空白和字典顺序;事件、state、memory 及 summary 的 session 归属、覆盖关系、回放版本严格比较。SDK 没有 summary version,harness 按成功写入次数计数。allowed_diff 只匹配精确路径。测试重建并读取仓库 JSON 基线,CLI 仅显式 `--output` 写文件。 diff --git a/tests/sessions/replay_harness.py b/tests/sessions/replay_harness.py new file mode 100644 index 000000000..43b8b21e5 --- /dev/null +++ b/tests/sessions/replay_harness.py @@ -0,0 +1,983 @@ +# 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. +"""Reusable Session/Memory/Summary replay consistency harness.""" + +from __future__ import annotations + +import argparse +import asyncio +import copy +import json +import math +import os +import re +import time +from dataclasses import asdict +from dataclasses import dataclass +from pathlib import Path +from tempfile import TemporaryDirectory +from typing import Any +from typing import Callable +from typing import Optional +from unittest.mock import patch + +from trpc_agent_sdk.abc import MemoryServiceConfig +from trpc_agent_sdk.events import Event +from trpc_agent_sdk.memory import InMemoryMemoryService +from trpc_agent_sdk.memory import RedisMemoryService +from trpc_agent_sdk.memory import SqlMemoryService +from trpc_agent_sdk.models import LlmResponse +from trpc_agent_sdk.sessions import InMemorySessionService +from trpc_agent_sdk.sessions import RedisSessionService +from trpc_agent_sdk.sessions import Session +from trpc_agent_sdk.sessions import SessionServiceConfig +from trpc_agent_sdk.sessions import SessionSummarizer +from trpc_agent_sdk.sessions import SqlSessionService +from trpc_agent_sdk.sessions import SummarizerSessionManager +from trpc_agent_sdk.storage import RedisStorage +from trpc_agent_sdk.types import Content +from trpc_agent_sdk.types import EventActions +from trpc_agent_sdk.types import FunctionCall +from trpc_agent_sdk.types import FunctionResponse +from trpc_agent_sdk.types import Part + +from .replay_cases import REPLAY_CASES +from .replay_cases import ReplayCase +from .replay_cases import ReplayOperation + +APP_NAME = "replay-consistency" +# Fixed event timestamps keep replay snapshots deterministic. TTL is disabled in +# the standard replay configuration, so this is not a storage commit timestamp. +BASE_TIMESTAMP = 4102444800.0 + + +class DeterministicSummaryModel: + """A no-network model that returns the summary selected by the trace.""" + + name = "deterministic-replay-summary" + + def __init__(self) -> None: + self.next_summary = "" + self.fail_next = False + self.failure_count = 0 + + async def generate_async(self, _request: Any, stream: bool = False, ctx: Any = None): + del stream, ctx + if self.fail_next: + self.fail_next = False + self.failure_count += 1 + raise RuntimeError("deterministic summary failure") + yield LlmResponse(content=Content(parts=[Part.from_text(text=self.next_summary)])) + + +class _FakeRedisStorage(RedisStorage): + """Use the SDK RedisStorage with a fakeredis client and no connection pool.""" + + def __init__(self, client: Any) -> None: + super().__init__(redis_url="redis://replay-mock", is_async=False) + self._client = client + self.closed = False + + async def create_redis_engine(self) -> None: + """The injected fakeredis client does not need a connection pool.""" + + async def create_redis_session(self) -> Any: + if self.closed: + raise RuntimeError("fakeredis storage is closed") + return self._client + + async def close(self) -> None: + # The injected fakeredis client owns the connection; the base class has + # no pool to disconnect, but keeping its close contract makes this fake + # safe if pool setup changes in the future. + await super().close() + self._client.close() + self.closed = True + + +@dataclass +class BackendBundle: + """Session, Memory, and Summary services operated as one backend.""" + + name: str + session_service: Any + memory_service: Any + summary_manager: SummarizerSessionManager + summary_model: DeterministicSummaryModel + summary_versions: dict[str, int] + cleanup: Optional[Callable[[], None]] = None + + async def close(self) -> None: + try: + try: + await self.memory_service.close() + finally: + await self.session_service.close() + finally: + if self.cleanup: + self.cleanup() + + +async def _close_backends(backends: list[BackendBundle]) -> None: + """Close backend bundles in reverse construction order.""" + + for backend in reversed(backends): + await backend.close() + + +@dataclass(frozen=True) +class AllowedDiff: + """A narrowly-scoped, documented backend difference.""" + + path: str + reason: str + + +@dataclass +class DiffEntry: + """One field-level difference with replay location metadata.""" + + case_id: str + left_backend: str + right_backend: str + session_id: str + path: str + left_value: Any + right_value: Any + allowed: bool + reason: Optional[str] = None + event_index: Optional[int] = None + summary_id: Optional[str] = None + + +ALLOWED_DIFFS: tuple[AllowedDiff, ...] = ( + AllowedDiff("$.session.last_update_time", "Persistent backends use their storage commit clock."), + AllowedDiff("$.summary.summary_timestamp", "Summary metadata is stamped independently by each manager."), +) + + +def _session_config(enable_ttl: bool = False) -> SessionServiceConfig: + config = SessionServiceConfig(store_historical_events=True) + if not enable_ttl: + config.clean_ttl_config() + return config + + +def _memory_config(enable_ttl: bool = False) -> MemoryServiceConfig: + config = MemoryServiceConfig(enabled=True) + if not enable_ttl: + config.clean_ttl_config() + return config + + +def _summary_components() -> tuple[DeterministicSummaryModel, SummarizerSessionManager]: + model = DeterministicSummaryModel() + summarizer = SessionSummarizer( + model=model, # type: ignore[arg-type] + check_summarizer_functions=[lambda _session: True], + keep_recent_count=2, + start_by_user_turn=True, + ) + manager = SummarizerSessionManager(model=model, summarizer=summarizer, + auto_summarize=True) # type: ignore[arg-type] + return model, manager + + +async def create_in_memory_backend() -> BackendBundle: + """Create the dependency-free lightweight backend.""" + + model, manager = _summary_components() + session_service = InMemorySessionService(summarizer_manager=manager, session_config=_session_config()) + memory_service = InMemoryMemoryService(memory_service_config=_memory_config(), enabled=True) + return BackendBundle("in_memory", session_service, memory_service, manager, model, {}) + + +async def create_sqlite_backend(name: str = "sqlite") -> BackendBundle: + """Create an isolated SQLite persistence backend.""" + + temporary_directory = TemporaryDirectory(prefix=f"trpc-replay-{name}-") + database_path = Path(temporary_directory.name) / "replay.db" + try: + bundle = await create_sql_backend(f"sqlite:///{database_path.as_posix()}", name=name) + except Exception: + temporary_directory.cleanup() + raise + bundle.cleanup = temporary_directory.cleanup + return bundle + + +async def create_sql_backend(db_url: str, name: str = "sql") -> BackendBundle: + """Create a SQL persistence backend from a sync SQLAlchemy URL.""" + + model, manager = _summary_components() + session_service: Optional[SqlSessionService] = None + memory_service: Optional[SqlMemoryService] = None + try: + session_service = SqlSessionService( + db_url=db_url, + summarizer_manager=manager, + session_config=_session_config(), + is_async=False, + ) + memory_service = SqlMemoryService( + db_url=db_url, + enabled=True, + memory_service_config=_memory_config(), + is_async=False, + ) + await session_service._sql_storage.create_sql_engine() # pylint: disable=protected-access + await memory_service._sql_storage.create_sql_engine() # pylint: disable=protected-access + return BackendBundle(name, session_service, memory_service, manager, model, {}) + except Exception: + # A factory failure must not leak an already-created SQL engine. + if memory_service is not None: + try: + await memory_service.close() + except Exception: # pylint: disable=broad-except + pass + if session_service is not None: + try: + await session_service.close() + except Exception: # pylint: disable=broad-except + pass + raise + + +async def create_redis_backend(redis_url: str) -> BackendBundle: + """Create an optional Redis integration backend.""" + + model, manager = _summary_components() + session_service = RedisSessionService( + db_url=redis_url, + summarizer_manager=manager, + session_config=_session_config(), + is_async=False, + ) + memory_service = RedisMemoryService( + db_url=redis_url, + enabled=True, + memory_service_config=_memory_config(), + is_async=False, + ) + return BackendBundle("redis", session_service, memory_service, manager, model, {}) + + +async def create_mock_redis_backend(enable_ttl: bool = False) -> BackendBundle: + """Create Redis services with real RedisStorage backed by fakeredis.""" + + try: + import fakeredis + except ImportError as exc: # pragma: no cover - exercised by an optional integration + raise RuntimeError("fakeredis is required for the Redis replay fallback") from exc + + model, manager = _summary_components() + server = fakeredis.FakeServer() + + def storage_factory(*args: Any, **kwargs: Any) -> _FakeRedisStorage: + if args: + raise TypeError(f"Unexpected positional RedisStorage arguments: {args}") + unexpected = set(kwargs) - {"redis_url", "is_async", "decode_responses"} + if unexpected: + raise TypeError(f"Unsupported RedisStorage arguments in fakeredis fallback: {sorted(unexpected)}") + if kwargs.get("is_async", False): + raise ValueError("The replay fakeredis fallback only supports synchronous RedisStorage") + client = fakeredis.FakeRedis( + server=server, + decode_responses=kwargs.get("decode_responses", True), + ) + return _FakeRedisStorage(client) + + session_service: Optional[RedisSessionService] = None + memory_service: Optional[RedisMemoryService] = None + try: + with patch("trpc_agent_sdk.sessions._redis_session_service.RedisStorage", + side_effect=storage_factory), patch("trpc_agent_sdk.memory._redis_memory_service.RedisStorage", + side_effect=storage_factory): + session_service = RedisSessionService( + db_url="redis://replay-mock", + summarizer_manager=manager, + session_config=_session_config(enable_ttl), + is_async=False, + ) + memory_service = RedisMemoryService( + db_url="redis://replay-mock", + enabled=True, + memory_service_config=_memory_config(enable_ttl), + is_async=False, + ) + except Exception: + if memory_service is not None: + try: + await memory_service.close() + except Exception: # pylint: disable=broad-except + pass + if session_service is not None: + try: + await session_service.close() + except Exception: # pylint: disable=broad-except + pass + raise + assert session_service is not None and memory_service is not None + return BackendBundle("redis_mock", session_service, memory_service, manager, model, {}) + + +def _identity(case: ReplayCase) -> tuple[str, str]: + return f"user-{case.case_id}", f"session-{case.case_id}" + + +def _app_name(case: ReplayCase) -> str: + return f"{APP_NAME}-{case.case_id}" + + +def _event_from_operation(case: ReplayCase, operation: ReplayOperation, sequence: int) -> Event: + payload = operation.payload + event_type = payload["event_type"] + event_id = payload.get("event_id", f"{case.case_id}-event-{sequence:02d}") + content = None + if event_type == "text": + content = Content(parts=[Part.from_text(text=payload["text"])]) + elif event_type == "function_call": + function_call = FunctionCall(id=payload["call_id"], name=payload["name"], args=payload["args"]) + content = Content(parts=[Part(function_call=function_call)]) + elif event_type == "function_response": + function_response = FunctionResponse( + id=payload["call_id"], + name=payload["name"], + response=payload["response"], + ) + content = Content(parts=[Part(function_response=function_response)]) + elif event_type != "state": + raise ValueError(f"Unsupported replay event type: {event_type}") + + return Event( + id=event_id, + invocation_id=f"{case.case_id}-invocation-{sequence:02d}", + author=payload["author"], + content=content, + actions=EventActions(state_delta=copy.deepcopy(payload.get("state_delta", {}))), + timestamp=BASE_TIMESTAMP + sequence, + partial=payload.get("partial", False), + turn_complete=not payload.get("partial", False), + ) + + +async def execute_case(bundle: BackendBundle, case: ReplayCase) -> dict[str, Any]: + """Replay one case and return its normalized backend snapshot.""" + + user_id, session_id = _identity(case) + app_name = _app_name(case) + await bundle.session_service.delete_session( + app_name=app_name, + user_id=user_id, + session_id=session_id, + ) + bundle.summary_versions[session_id] = 0 + session = await bundle.session_service.create_session( + app_name=app_name, + user_id=user_id, + session_id=session_id, + state={"case_id": case.case_id}, + ) + + sequence = 0 + checks: dict[str, bool] = {} + summary_history: list[dict[str, Any]] = [] + previous_summary_timestamp: Optional[float] = None + for operation in case.operations: + if operation.kind == "event": + sequence += 1 + event = _event_from_operation(case, operation, sequence) + before_state = copy.deepcopy(session.state) + attempted_state_delta = copy.deepcopy(event.actions.state_delta) + await bundle.session_service.append_event(session, event) + if event.partial: + session = await _get_session(bundle, app_name, user_id, session_id) + check_prefix = f"partial_event_{sequence}" + persistent_delta = { + key: value + for key, value in attempted_state_delta.items() if not key.startswith("temp:") + } + temporary_keys = [key for key in attempted_state_delta if key.startswith("temp:")] + checks[f"{check_prefix}_not_persisted"] = all(item.id != event.id for item in session.events) + checks[f"{check_prefix}_state_clean"] = all( + session.state.get(key) == before_state.get(key) for key in persistent_delta) + checks[f"{check_prefix}_temp_state_clean"] = all(key not in session.state for key in temporary_keys) + elif operation.kind == "store_memory": + session = await _get_session(bundle, app_name, user_id, session_id) + await bundle.memory_service.store_session(session) + elif operation.kind == "summarize": + session = await _get_session(bundle, app_name, user_id, session_id) + bundle.summary_model.next_summary = operation.payload["summary_text"] + await bundle.session_service.create_session_summary(session) + public_summary = await bundle.session_service.get_session_summary(session) + stored_summary = await bundle.summary_manager.get_session_summary(session) + if stored_summary is None or public_summary is None: + checks[f"summary_{len(summary_history) + 1}_saved"] = False + else: + version = bundle.summary_versions[session_id] + 1 + bundle.summary_versions[session_id] = version + summary_history.append( + _summary_revision_snapshot( + stored_summary, + public_summary, + version, + previous_summary_timestamp, + )) + previous_summary_timestamp = stored_summary.summary_timestamp + checks[f"summary_{version}_saved"] = True + checks[f"summary_{version}_public_read"] = ( + _normalize_summary_text(public_summary) == _normalize_summary_text(stored_summary.summary_text)) + # update_session implementations may retain the supplied object. + # Re-read before subsequent writes to preserve the service boundary. + session = await _get_session(bundle, app_name, user_id, session_id) + elif operation.kind == "summarize_failure": + session = await _get_session(bundle, app_name, user_id, session_id) + before_projection = _recovery_projection(session) + before_summary = await bundle.session_service.get_session_summary(session) + failures_before = bundle.summary_model.failure_count + bundle.summary_model.fail_next = True + await bundle.session_service.create_session_summary(session) + session = await _get_session(bundle, app_name, user_id, session_id) + after_summary = await bundle.session_service.get_session_summary(session) + checks["summary_model_failed"] = bundle.summary_model.failure_count == failures_before + 1 + checks["failed_summary_preserves_session"] = _recovery_projection(session) == before_projection + checks["failed_summary_preserves_summary"] = after_summary == before_summary + else: + raise ValueError(f"Unsupported replay operation: {operation.kind}") + + session = await _get_session(bundle, app_name, user_id, session_id) + memory: dict[str, list[dict[str, Any]]] = {} + for query in case.memory_queries: + response = await bundle.memory_service.search_memory(session.save_key, query) + memory[query] = sorted( + (_memory_entry_snapshot(entry) for entry in response.memories), + key=lambda item: json.dumps(item, ensure_ascii=False, sort_keys=True), + ) + + public_summary = await bundle.session_service.get_session_summary(session) + summary = await bundle.summary_manager.get_session_summary(session) + return { + "session": _session_snapshot(session), + "memory": memory, + "summary": _summary_snapshot(summary, public_summary, bundle.summary_versions[session_id]), + "summary_history": summary_history, + "checks": checks, + } + + +async def _get_session(bundle: BackendBundle, app_name: str, user_id: str, session_id: str) -> Session: + session = await bundle.session_service.get_session( + app_name=app_name, + user_id=user_id, + session_id=session_id, + ) + if session is None: + raise AssertionError(f"{bundle.name} lost session {session_id}") + return session + + +def _session_snapshot(session: Session) -> dict[str, Any]: + return { + "id": session.id, + "app_name": session.app_name, + "user_id": session.user_id, + "save_key": session.save_key, + "state": _json_value(session.state), + "events": [_event_snapshot(event) for event in session.events], + "historical_events": [_event_snapshot(event) for event in session.historical_events], + "conversation_count": session.conversation_count, + "last_update_time": session.last_update_time, + } + + +def _event_snapshot(event: Event) -> dict[str, Any]: + is_summary = event.is_summary_event() + return { + "id": "" if is_summary else event.id, + "invocation_id": event.invocation_id, + "author": event.author, + "content": _json_value(event.content.model_dump(exclude_none=True, mode="json") if event.content else None), + "actions": _json_value(event.actions.model_dump(exclude_none=True, mode="json")), + "timestamp": "" if is_summary else event.timestamp, + "partial": event.partial, + "turn_complete": event.turn_complete, + "visible": event.visible, + "model_flags": event.model_flags, + "version": event.version, + "is_summary": is_summary, + } + + +def _memory_entry_snapshot(entry: Any) -> dict[str, Any]: + return { + "author": entry.author, + "content": _json_value(entry.content.model_dump(exclude_none=True, mode="json")), + "timestamp": entry.timestamp, + } + + +def _normalize_summary_text(summary_text: str) -> str: + return " ".join(summary_text.split()) + + +def _summary_snapshot(summary: Any, public_summary: Optional[str], version: int) -> Optional[dict[str, Any]]: + if summary is None: + return None + return { + "summary_id": f"{summary.session_id}:summary", + "session_id": summary.session_id, + "summary_text": _normalize_summary_text(public_summary or ""), + "metadata_summary_text": _normalize_summary_text(summary.summary_text), + "original_event_count": summary.original_event_count, + "compressed_event_count": summary.compressed_event_count, + "summary_timestamp": summary.summary_timestamp, + "version": version, + } + + +def _summary_revision_snapshot( + summary: Any, + public_summary: str, + version: int, + previous_timestamp: Optional[float], +) -> dict[str, Any]: + timestamp = summary.summary_timestamp + return { + "summary_id": f"{summary.session_id}:summary", + "session_id": summary.session_id, + "summary_text": _normalize_summary_text(public_summary), + "metadata_summary_text": _normalize_summary_text(summary.summary_text), + "original_event_count": summary.original_event_count, + "compressed_event_count": summary.compressed_event_count, + "timestamp_valid": isinstance(timestamp, (int, float)) and math.isfinite(timestamp) and timestamp > 0, + "updated_after_previous": previous_timestamp is None or timestamp > previous_timestamp, + "version": version, + } + + +def _recovery_projection(session: Session) -> dict[str, Any]: + """Return business fields that must survive a failed summary write.""" + + snapshot = _session_snapshot(session) + snapshot.pop("last_update_time") + return snapshot + + +def _json_value(value: Any) -> Any: + """Convert model values to stable JSON primitives without stringifying structures.""" + + return json.loads(json.dumps(value, ensure_ascii=False, sort_keys=True, default=str)) + + +def compare_snapshots( + case_id: str, + session_id: str, + left_backend: str, + right_backend: str, + left: Any, + right: Any, +) -> list[DiffEntry]: + """Recursively compare snapshots while retaining list order and field paths.""" + + differences: list[DiffEntry] = [] + + def walk(left_value: Any, right_value: Any, path: str) -> None: + if isinstance(left_value, dict) and isinstance(right_value, dict): + for key in sorted(set(left_value) | set(right_value)): + child_path = f"{path}.{key}" + if key not in left_value: + add(child_path, "", right_value[key]) + elif key not in right_value: + add(child_path, left_value[key], "") + else: + walk(left_value[key], right_value[key], child_path) + return + if isinstance(left_value, list) and isinstance(right_value, list): + for index in range(max(len(left_value), len(right_value))): + child_path = f"{path}[{index}]" + if index >= len(left_value): + add(child_path, "", right_value[index]) + elif index >= len(right_value): + add(child_path, left_value[index], "") + else: + walk(left_value[index], right_value[index], child_path) + return + if left_value != right_value: + add(path, left_value, right_value) + + def add(path: str, left_value: Any, right_value: Any) -> None: + reason = next((item.reason for item in ALLOWED_DIFFS if item.path == path), None) + event_match = re.search(r"\.(?:historical_)?events\[(\d+)\]", path) + summary_id = None + if path.startswith("$.summary"): + left_summary = left if isinstance(left, dict) else {} + right_summary = right if isinstance(right, dict) else {} + summary_id = ((left_summary.get("summary") or {}).get("summary_id") + or (right_summary.get("summary") or {}).get("summary_id")) + differences.append( + DiffEntry( + case_id=case_id, + left_backend=left_backend, + right_backend=right_backend, + session_id=session_id, + path=path, + left_value=_json_value(left_value), + right_value=_json_value(right_value), + allowed=reason is not None, + reason=reason, + event_index=int(event_match.group(1)) if event_match else None, + summary_id=summary_id, + )) + + walk(left, right, "$") + return differences + + +def _delete_first_event(snapshot: dict[str, Any]) -> None: + del snapshot["session"]["events"][0] + + +def _swap_first_events(snapshot: dict[str, Any]) -> None: + events = snapshot["session"]["events"] + events[0], events[1] = events[1], events[0] + + +def _corrupt_tool_response(snapshot: dict[str, Any]) -> None: + event = next(event for event in snapshot["session"]["events"] + if event["content"] and event["content"]["parts"][0].get("function_response")) + event["content"]["parts"][0]["function_response"]["response"]["result"] = "corrupted" + + +def _corrupt_state(snapshot: dict[str, Any]) -> None: + snapshot["session"]["state"]["phase"] = "corrupted" + + +def _corrupt_scoped_state(snapshot: dict[str, Any]) -> None: + snapshot["session"]["state"]["user:language"] = "corrupted" + + +def _corrupt_memory(snapshot: dict[str, Any]) -> None: + entry = next(iter(snapshot["memory"].values()))[0] + entry["content"]["parts"][0]["text"] = "corrupted memory" + + +def _drop_summary(snapshot: dict[str, Any]) -> None: + snapshot["summary"] = None + + +def _restore_old_summary(snapshot: dict[str, Any]) -> None: + snapshot["summary"]["summary_text"] = "summary update version one" + snapshot["summary"]["version"] = 1 + + +def _move_summary(snapshot: dict[str, Any]) -> None: + snapshot["summary"]["session_id"] = "session-wrong-owner" + + +def _duplicate_event(snapshot: dict[str, Any]) -> None: + snapshot["session"]["events"].append(copy.deepcopy(snapshot["session"]["events"][-1])) + + +def _check_failures( + case_id: str, + session_id: str, + backend_name: str, + snapshot: dict[str, Any], +) -> list[DiffEntry]: + """Turn failed replay invariants into ordinary field-level differences.""" + + return [ + DiffEntry( + case_id=case_id, + left_backend="expected", + right_backend=backend_name, + session_id=session_id, + path=f"$.checks.{check_name}", + left_value=True, + right_value=passed, + allowed=False, + ) for check_name, passed in snapshot["checks"].items() if not passed + ] + + +InjectedFault = tuple[str, str, Callable[[dict[str, Any]], None]] + +INJECTED_FAULTS: tuple[InjectedFault, ...] = ( + ("missing_event", "single_turn", _delete_first_event), + ("event_order", "multi_turn", _swap_first_events), + ("tool_response", "tool_call", _corrupt_tool_response), + ("state_overwrite", "state_overwrite", _corrupt_state), + ("state_scope", "scoped_state", _corrupt_scoped_state), + ("memory_content", "memory_roundtrip", _corrupt_memory), + ("summary_missing", "summary_create", _drop_summary), + ("summary_overwrite", "summary_update", _restore_old_summary), + ("summary_session_owner", "summary_truncation", _move_summary), + ("duplicate_retry", "partial_retry", _duplicate_event), +) + + +def _validate_replay_inputs( + backends: list[BackendBundle], + cases: tuple[ReplayCase, ...], + injected_faults: tuple[InjectedFault, ...], +) -> None: + """Validate matrix identity constraints before mutating any backend.""" + + backend_names = [backend.name for backend in backends] + if not backends: + raise ValueError("At least one replay backend is required") + if len(set(backend_names)) != len(backend_names): + raise ValueError("Replay backend names must be unique") + if not cases: + raise ValueError("At least one replay case is required") + case_id_list = [case.case_id for case in cases] + if len(set(case_id_list)) != len(case_id_list): + raise ValueError("Replay case IDs must be unique") + if not injected_faults: + raise ValueError("At least one injected fault is required") + fault_id_list = [fault_id for fault_id, _case_id, _inject in injected_faults] + if len(set(fault_id_list)) != len(fault_id_list): + raise ValueError("Injected fault IDs must be unique") + case_ids = set(case_id_list) + missing_case_ids = sorted(case_id for _fault_id, case_id, _inject in injected_faults if case_id not in case_ids) + if missing_case_ids: + raise ValueError(f"Injected faults reference unknown replay cases: {missing_case_ids}") + if not any(fault_id.startswith("summary_") for fault_id, _case_id, _inject in injected_faults): + raise ValueError("At least one summary fault is required") + + +async def run_replay_matrix( + backends: list[BackendBundle], + cases: tuple[ReplayCase, ...] = REPLAY_CASES, + injected_faults: tuple[InjectedFault, ...] = INJECTED_FAULTS, +) -> dict[str, Any]: + """Run normal comparisons and the public ten-fault detection matrix. + + The matrix takes ownership of ``backends`` and closes them even when input + validation or a replay operation fails. + """ + + started = time.perf_counter() + snapshots: dict[str, dict[str, dict[str, Any]]] = {backend.name: {} for backend in backends} + try: + _validate_replay_inputs(backends, cases, injected_faults) + for case in cases: + for backend in backends: + snapshots[backend.name][case.case_id] = await execute_case(backend, case) + + baseline = backends[0] + normal_results = [] + for case in cases: + case_diffs: list[DiffEntry] = [] + _, session_id = _identity(case) + for backend in backends: + case_diffs.extend( + _check_failures( + case.case_id, + session_id, + backend.name, + snapshots[backend.name][case.case_id], + )) + for backend in backends[1:]: + case_diffs.extend( + compare_snapshots( + case.case_id, + session_id, + baseline.name, + backend.name, + snapshots[baseline.name][case.case_id], + snapshots[backend.name][case.case_id], + )) + normal_results.append({ + "case_id": case.case_id, + "passed": not any(not diff.allowed for diff in case_diffs), + "differences": [asdict(diff) for diff in case_diffs], + }) + + injected_results = [] + for fault_id, case_id, inject in injected_faults: + expected = snapshots[baseline.name][case_id] + corrupted = copy.deepcopy(expected) + inject(corrupted) + session_id = expected["session"]["id"] + differences = compare_snapshots( + case_id, + session_id, + baseline.name, + f"injected:{fault_id}", + expected, + corrupted, + ) + detected = any(not diff.allowed for diff in differences) + injected_results.append({ + "fault_id": fault_id, + "case_id": case_id, + "detected": detected, + "differences": [asdict(diff) for diff in differences], + }) + + normal_failures = sum(not result["passed"] for result in normal_results) + detected_faults = sum(result["detected"] for result in injected_results) + summary_faults = [result for result in injected_results if result["fault_id"].startswith("summary_")] + return { + "schema_version": 1, + "generated_at": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()), + "backends": [backend.name for backend in backends], + "case_count": len(cases), + "allowed_diff": [asdict(item) for item in ALLOWED_DIFFS], + "normal_cases": normal_results, + "injected_cases": injected_results, + "metrics": { + "false_positive_rate": normal_failures / len(normal_results), + "injected_detection_rate": detected_faults / len(injected_results), + "summary_fault_detection_rate": + sum(result["detected"] for result in summary_faults) / len(summary_faults), + "duration_seconds": round(time.perf_counter() - started, 6), + }, + } + finally: + await _close_backends(backends) + + +def canonical_report(report: dict[str, Any]) -> dict[str, Any]: + """Remove runtime-only report values while retaining semantic differences.""" + + def normalize_dynamic_value(path: str, value: Any) -> Any: + if isinstance(value, dict): + if path == "$.session": + normalized = copy.deepcopy(value) + if "last_update_time" in normalized: + normalized["last_update_time"] = "" + return normalized + if path == "$.summary": + normalized = copy.deepcopy(value) + if "summary_timestamp" in normalized: + normalized["summary_timestamp"] = "" + return normalized + if isinstance(value, list): + return [normalize_dynamic_value(path, item) for item in value] + return value + + normalized = copy.deepcopy(report) + normalized.pop("generated_at", None) + normalized.get("metrics", {}).pop("duration_seconds", None) + dynamic_paths = {item.path for item in ALLOWED_DIFFS} + for section in ("normal_cases", "injected_cases"): + for result in normalized.get(section, []): + for difference in result.get("differences", []): + if difference.get("allowed") or difference.get("path") in dynamic_paths: + difference["left_value"] = "" + difference["right_value"] = "" + else: + difference["left_value"] = normalize_dynamic_value(difference.get("path", ""), + difference.get("left_value")) + difference["right_value"] = normalize_dynamic_value(difference.get("path", ""), + difference.get("right_value")) + return _json_value(normalized) + + +def stable_report_signature(report: dict[str, Any]) -> dict[str, Any]: + """Return the persisted report contract without runtime-only metadata. + + Case results and all non-allowed field values remain part of the signature; + only report timestamps, duration, and values covered by ``ALLOWED_DIFFS`` + are normalized by ``canonical_report``. + """ + + normalized = canonical_report(report) + expected_keys = { + "schema_version", + "backends", + "case_count", + "allowed_diff", + "normal_cases", + "injected_cases", + "metrics", + } + if set(normalized) != expected_keys: + raise ValueError(f"Unexpected replay report fields: {sorted(set(normalized) ^ expected_keys)}") + metrics = normalized.get("metrics", {}) + expected_metric_keys = { + "false_positive_rate", + "injected_detection_rate", + "summary_fault_detection_rate", + } + if set(metrics) != expected_metric_keys: + raise ValueError(f"Unexpected replay report metrics: {sorted(set(metrics) ^ expected_metric_keys)}") + return { + "schema_version": normalized.get("schema_version"), + "backends": normalized.get("backends"), + "case_count": normalized.get("case_count"), + "allowed_diff": normalized.get("allowed_diff"), + "normal_cases": normalized.get("normal_cases", []), + "injected_cases": normalized.get("injected_cases", []), + "metrics": { + "false_positive_rate": metrics.get("false_positive_rate"), + "injected_detection_rate": metrics.get("injected_detection_rate"), + "summary_fault_detection_rate": metrics.get("summary_fault_detection_rate"), + }, + } + + +def write_report(report: dict[str, Any], output_path: Path) -> None: + output_path.write_text(json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") + + +async def _generate_default_report( + output_path: Optional[Path] = None, + lightweight: bool = False, + include_integrations: bool = False, +) -> dict[str, Any]: + if lightweight and include_integrations: + raise ValueError("--light and --integration are mutually exclusive") + backends: list[BackendBundle] = [] + matrix_started = False + try: + backends.append(await create_in_memory_backend()) + if include_integrations: + sql_url = os.getenv("TRPC_REPLAY_SQL_URL") + redis_url = os.getenv("TRPC_REPLAY_REDIS_URL") + backends.append(await create_sql_backend(sql_url, name="sql") if sql_url else await create_sqlite_backend( + name="sql_fallback")) + backends.append(await create_redis_backend(redis_url) if redis_url else await create_mock_redis_backend()) + elif not lightweight: + backends.append(await create_sqlite_backend()) + matrix_started = True + report = await run_replay_matrix(backends) + if output_path is not None: + write_report(report, output_path) + return report + finally: + if not matrix_started: + await _close_backends(backends) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--output", + type=Path, + help="write the generated report to this path; omit to avoid changing repository files", + ) + mode = parser.add_mutually_exclusive_group() + mode.add_argument("--light", action="store_true", help="run only the dependency-free InMemory backend") + mode.add_argument( + "--integration", + action="store_true", + help="use configured SQL/Redis backends or dependency-free fallbacks", + ) + args = parser.parse_args() + report = asyncio.run( + _generate_default_report( + args.output, + lightweight=args.light, + include_integrations=args.integration, + )) + print(json.dumps(report["metrics"], ensure_ascii=False, sort_keys=True)) + + +if __name__ == "__main__": + main() diff --git a/tests/sessions/session_memory_summary_diff_report.json b/tests/sessions/session_memory_summary_diff_report.json new file mode 100644 index 000000000..9774e99c9 --- /dev/null +++ b/tests/sessions/session_memory_summary_diff_report.json @@ -0,0 +1,708 @@ +{ + "schema_version": 1, + "generated_at": "2026-07-26T05:23:06Z", + "backends": [ + "in_memory", + "sqlite" + ], + "case_count": 10, + "allowed_diff": [ + { + "path": "$.session.last_update_time", + "reason": "Persistent backends use their storage commit clock." + }, + { + "path": "$.summary.summary_timestamp", + "reason": "Summary metadata is stamped independently by each manager." + } + ], + "normal_cases": [ + { + "case_id": "single_turn", + "passed": true, + "differences": [ + { + "case_id": "single_turn", + "left_backend": "in_memory", + "right_backend": "sqlite", + "session_id": "session-single_turn", + "path": "$.session.last_update_time", + "left_value": 1785043386.1692607, + "right_value": 1785043386.0, + "allowed": true, + "reason": "Persistent backends use their storage commit clock.", + "event_index": null, + "summary_id": null + } + ] + }, + { + "case_id": "multi_turn", + "passed": true, + "differences": [ + { + "case_id": "multi_turn", + "left_backend": "in_memory", + "right_backend": "sqlite", + "session_id": "session-multi_turn", + "path": "$.session.last_update_time", + "left_value": 1785043386.2688792, + "right_value": 1785043386.0, + "allowed": true, + "reason": "Persistent backends use their storage commit clock.", + "event_index": null, + "summary_id": null + } + ] + }, + { + "case_id": "tool_call", + "passed": true, + "differences": [ + { + "case_id": "tool_call", + "left_backend": "in_memory", + "right_backend": "sqlite", + "session_id": "session-tool_call", + "path": "$.session.last_update_time", + "left_value": 1785043386.3039174, + "right_value": 1785043386.0, + "allowed": true, + "reason": "Persistent backends use their storage commit clock.", + "event_index": null, + "summary_id": null + } + ] + }, + { + "case_id": "state_overwrite", + "passed": true, + "differences": [ + { + "case_id": "state_overwrite", + "left_backend": "in_memory", + "right_backend": "sqlite", + "session_id": "session-state_overwrite", + "path": "$.session.last_update_time", + "left_value": 1785043386.3300338, + "right_value": 1785043386.0, + "allowed": true, + "reason": "Persistent backends use their storage commit clock.", + "event_index": null, + "summary_id": null + } + ] + }, + { + "case_id": "scoped_state", + "passed": true, + "differences": [ + { + "case_id": "scoped_state", + "left_backend": "in_memory", + "right_backend": "sqlite", + "session_id": "session-scoped_state", + "path": "$.session.last_update_time", + "left_value": 1785043386.3520706, + "right_value": 1785043386.0, + "allowed": true, + "reason": "Persistent backends use their storage commit clock.", + "event_index": null, + "summary_id": null + } + ] + }, + { + "case_id": "memory_roundtrip", + "passed": true, + "differences": [ + { + "case_id": "memory_roundtrip", + "left_backend": "in_memory", + "right_backend": "sqlite", + "session_id": "session-memory_roundtrip", + "path": "$.session.last_update_time", + "left_value": 1785043386.372096, + "right_value": 1785043386.0, + "allowed": true, + "reason": "Persistent backends use their storage commit clock.", + "event_index": null, + "summary_id": null + } + ] + }, + { + "case_id": "summary_create", + "passed": true, + "differences": [ + { + "case_id": "summary_create", + "left_backend": "in_memory", + "right_backend": "sqlite", + "session_id": "session-summary_create", + "path": "$.session.last_update_time", + "left_value": 1785043386.408309, + "right_value": 1785043386.0, + "allowed": true, + "reason": "Persistent backends use their storage commit clock.", + "event_index": null, + "summary_id": null + }, + { + "case_id": "summary_create", + "left_backend": "in_memory", + "right_backend": "sqlite", + "session_id": "session-summary_create", + "path": "$.summary.summary_timestamp", + "left_value": 1785043386.4093192, + "right_value": 1785043386.459698, + "allowed": true, + "reason": "Summary metadata is stamped independently by each manager.", + "event_index": null, + "summary_id": "session-summary_create:summary" + } + ] + }, + { + "case_id": "summary_update", + "passed": true, + "differences": [ + { + "case_id": "summary_update", + "left_backend": "in_memory", + "right_backend": "sqlite", + "session_id": "session-summary_update", + "path": "$.session.last_update_time", + "left_value": 1785043386.4765158, + "right_value": 1785043386.0, + "allowed": true, + "reason": "Persistent backends use their storage commit clock.", + "event_index": null, + "summary_id": null + }, + { + "case_id": "summary_update", + "left_backend": "in_memory", + "right_backend": "sqlite", + "session_id": "session-summary_update", + "path": "$.summary.summary_timestamp", + "left_value": 1785043386.481025, + "right_value": 1785043386.591315, + "allowed": true, + "reason": "Summary metadata is stamped independently by each manager.", + "event_index": null, + "summary_id": "session-summary_update:summary" + } + ] + }, + { + "case_id": "summary_truncation", + "passed": true, + "differences": [ + { + "case_id": "summary_truncation", + "left_backend": "in_memory", + "right_backend": "sqlite", + "session_id": "session-summary_truncation", + "path": "$.session.last_update_time", + "left_value": 1785043386.607425, + "right_value": 1785043386.0, + "allowed": true, + "reason": "Persistent backends use their storage commit clock.", + "event_index": null, + "summary_id": null + }, + { + "case_id": "summary_truncation", + "left_backend": "in_memory", + "right_backend": "sqlite", + "session_id": "session-summary_truncation", + "path": "$.summary.summary_timestamp", + "left_value": 1785043386.6084247, + "right_value": 1785043386.662242, + "allowed": true, + "reason": "Summary metadata is stamped independently by each manager.", + "event_index": null, + "summary_id": "session-summary_truncation:summary" + } + ] + }, + { + "case_id": "partial_retry", + "passed": true, + "differences": [ + { + "case_id": "partial_retry", + "left_backend": "in_memory", + "right_backend": "sqlite", + "session_id": "session-partial_retry", + "path": "$.session.last_update_time", + "left_value": 1785043386.6857631, + "right_value": 1785043386.0, + "allowed": true, + "reason": "Persistent backends use their storage commit clock.", + "event_index": null, + "summary_id": null + } + ] + } + ], + "injected_cases": [ + { + "fault_id": "missing_event", + "case_id": "single_turn", + "detected": true, + "differences": [ + { + "case_id": "single_turn", + "left_backend": "in_memory", + "right_backend": "injected:missing_event", + "session_id": "session-single_turn", + "path": "$.session.events[0].author", + "left_value": "user", + "right_value": "assistant", + "allowed": false, + "reason": null, + "event_index": 0, + "summary_id": null + }, + { + "case_id": "single_turn", + "left_backend": "in_memory", + "right_backend": "injected:missing_event", + "session_id": "session-single_turn", + "path": "$.session.events[0].content.parts[0].text", + "left_value": "hello replay", + "right_value": "hello user", + "allowed": false, + "reason": null, + "event_index": 0, + "summary_id": null + }, + { + "case_id": "single_turn", + "left_backend": "in_memory", + "right_backend": "injected:missing_event", + "session_id": "session-single_turn", + "path": "$.session.events[0].id", + "left_value": "single_turn-event-01", + "right_value": "single_turn-event-02", + "allowed": false, + "reason": null, + "event_index": 0, + "summary_id": null + }, + { + "case_id": "single_turn", + "left_backend": "in_memory", + "right_backend": "injected:missing_event", + "session_id": "session-single_turn", + "path": "$.session.events[0].invocation_id", + "left_value": "single_turn-invocation-01", + "right_value": "single_turn-invocation-02", + "allowed": false, + "reason": null, + "event_index": 0, + "summary_id": null + }, + { + "case_id": "single_turn", + "left_backend": "in_memory", + "right_backend": "injected:missing_event", + "session_id": "session-single_turn", + "path": "$.session.events[0].timestamp", + "left_value": 4102444801.0, + "right_value": 4102444802.0, + "allowed": false, + "reason": null, + "event_index": 0, + "summary_id": null + }, + { + "case_id": "single_turn", + "left_backend": "in_memory", + "right_backend": "injected:missing_event", + "session_id": "session-single_turn", + "path": "$.session.events[1]", + "left_value": { + "actions": { + "artifact_delta": {}, + "state_delta": {} + }, + "author": "assistant", + "content": { + "parts": [ + { + "text": "hello user" + } + ] + }, + "id": "single_turn-event-02", + "invocation_id": "single_turn-invocation-02", + "is_summary": false, + "model_flags": 1, + "partial": false, + "timestamp": 4102444802.0, + "turn_complete": true, + "version": 0, + "visible": true + }, + "right_value": "", + "allowed": false, + "reason": null, + "event_index": 1, + "summary_id": null + } + ] + }, + { + "fault_id": "event_order", + "case_id": "multi_turn", + "detected": true, + "differences": [ + { + "case_id": "multi_turn", + "left_backend": "in_memory", + "right_backend": "injected:event_order", + "session_id": "session-multi_turn", + "path": "$.session.events[0].author", + "left_value": "user", + "right_value": "assistant", + "allowed": false, + "reason": null, + "event_index": 0, + "summary_id": null + }, + { + "case_id": "multi_turn", + "left_backend": "in_memory", + "right_backend": "injected:event_order", + "session_id": "session-multi_turn", + "path": "$.session.events[0].content.parts[0].text", + "left_value": "multi user turn 1", + "right_value": "multi assistant turn 1", + "allowed": false, + "reason": null, + "event_index": 0, + "summary_id": null + }, + { + "case_id": "multi_turn", + "left_backend": "in_memory", + "right_backend": "injected:event_order", + "session_id": "session-multi_turn", + "path": "$.session.events[0].id", + "left_value": "multi_turn-event-01", + "right_value": "multi_turn-event-02", + "allowed": false, + "reason": null, + "event_index": 0, + "summary_id": null + }, + { + "case_id": "multi_turn", + "left_backend": "in_memory", + "right_backend": "injected:event_order", + "session_id": "session-multi_turn", + "path": "$.session.events[0].invocation_id", + "left_value": "multi_turn-invocation-01", + "right_value": "multi_turn-invocation-02", + "allowed": false, + "reason": null, + "event_index": 0, + "summary_id": null + }, + { + "case_id": "multi_turn", + "left_backend": "in_memory", + "right_backend": "injected:event_order", + "session_id": "session-multi_turn", + "path": "$.session.events[0].timestamp", + "left_value": 4102444801.0, + "right_value": 4102444802.0, + "allowed": false, + "reason": null, + "event_index": 0, + "summary_id": null + }, + { + "case_id": "multi_turn", + "left_backend": "in_memory", + "right_backend": "injected:event_order", + "session_id": "session-multi_turn", + "path": "$.session.events[1].author", + "left_value": "assistant", + "right_value": "user", + "allowed": false, + "reason": null, + "event_index": 1, + "summary_id": null + }, + { + "case_id": "multi_turn", + "left_backend": "in_memory", + "right_backend": "injected:event_order", + "session_id": "session-multi_turn", + "path": "$.session.events[1].content.parts[0].text", + "left_value": "multi assistant turn 1", + "right_value": "multi user turn 1", + "allowed": false, + "reason": null, + "event_index": 1, + "summary_id": null + }, + { + "case_id": "multi_turn", + "left_backend": "in_memory", + "right_backend": "injected:event_order", + "session_id": "session-multi_turn", + "path": "$.session.events[1].id", + "left_value": "multi_turn-event-02", + "right_value": "multi_turn-event-01", + "allowed": false, + "reason": null, + "event_index": 1, + "summary_id": null + }, + { + "case_id": "multi_turn", + "left_backend": "in_memory", + "right_backend": "injected:event_order", + "session_id": "session-multi_turn", + "path": "$.session.events[1].invocation_id", + "left_value": "multi_turn-invocation-02", + "right_value": "multi_turn-invocation-01", + "allowed": false, + "reason": null, + "event_index": 1, + "summary_id": null + }, + { + "case_id": "multi_turn", + "left_backend": "in_memory", + "right_backend": "injected:event_order", + "session_id": "session-multi_turn", + "path": "$.session.events[1].timestamp", + "left_value": 4102444802.0, + "right_value": 4102444801.0, + "allowed": false, + "reason": null, + "event_index": 1, + "summary_id": null + } + ] + }, + { + "fault_id": "tool_response", + "case_id": "tool_call", + "detected": true, + "differences": [ + { + "case_id": "tool_call", + "left_backend": "in_memory", + "right_backend": "injected:tool_response", + "session_id": "session-tool_call", + "path": "$.session.events[2].content.parts[0].function_response.response.result", + "left_value": "consistent", + "right_value": "corrupted", + "allowed": false, + "reason": null, + "event_index": 2, + "summary_id": null + } + ] + }, + { + "fault_id": "state_overwrite", + "case_id": "state_overwrite", + "detected": true, + "differences": [ + { + "case_id": "state_overwrite", + "left_backend": "in_memory", + "right_backend": "injected:state_overwrite", + "session_id": "session-state_overwrite", + "path": "$.session.state.phase", + "left_value": "done", + "right_value": "corrupted", + "allowed": false, + "reason": null, + "event_index": null, + "summary_id": null + } + ] + }, + { + "fault_id": "state_scope", + "case_id": "scoped_state", + "detected": true, + "differences": [ + { + "case_id": "scoped_state", + "left_backend": "in_memory", + "right_backend": "injected:state_scope", + "session_id": "session-scoped_state", + "path": "$.session.state.user:language", + "left_value": "go", + "right_value": "corrupted", + "allowed": false, + "reason": null, + "event_index": null, + "summary_id": null + } + ] + }, + { + "fault_id": "memory_content", + "case_id": "memory_roundtrip", + "detected": true, + "differences": [ + { + "case_id": "memory_roundtrip", + "left_backend": "in_memory", + "right_backend": "injected:memory_content", + "session_id": "session-memory_roundtrip", + "path": "$.memory.preference-token-oolong[0].content.parts[0].text", + "left_value": "I will remember that preference", + "right_value": "corrupted memory", + "allowed": false, + "reason": null, + "event_index": null, + "summary_id": null + } + ] + }, + { + "fault_id": "summary_missing", + "case_id": "summary_create", + "detected": true, + "differences": [ + { + "case_id": "summary_create", + "left_backend": "in_memory", + "right_backend": "injected:summary_missing", + "session_id": "session-summary_create", + "path": "$.summary", + "left_value": { + "compressed_event_count": 3, + "metadata_summary_text": "summary version one", + "original_event_count": 8, + "session_id": "session-summary_create", + "summary_id": "session-summary_create:summary", + "summary_text": "summary version one", + "summary_timestamp": 1785043386.4093192, + "version": 1 + }, + "right_value": null, + "allowed": false, + "reason": null, + "event_index": null, + "summary_id": "session-summary_create:summary" + } + ] + }, + { + "fault_id": "summary_overwrite", + "case_id": "summary_update", + "detected": true, + "differences": [ + { + "case_id": "summary_update", + "left_backend": "in_memory", + "right_backend": "injected:summary_overwrite", + "session_id": "session-summary_update", + "path": "$.summary.summary_text", + "left_value": "summary update version two", + "right_value": "summary update version one", + "allowed": false, + "reason": null, + "event_index": null, + "summary_id": "session-summary_update:summary" + }, + { + "case_id": "summary_update", + "left_backend": "in_memory", + "right_backend": "injected:summary_overwrite", + "session_id": "session-summary_update", + "path": "$.summary.version", + "left_value": 2, + "right_value": 1, + "allowed": false, + "reason": null, + "event_index": null, + "summary_id": "session-summary_update:summary" + } + ] + }, + { + "fault_id": "summary_session_owner", + "case_id": "summary_truncation", + "detected": true, + "differences": [ + { + "case_id": "summary_truncation", + "left_backend": "in_memory", + "right_backend": "injected:summary_session_owner", + "session_id": "session-summary_truncation", + "path": "$.summary.session_id", + "left_value": "session-summary_truncation", + "right_value": "session-wrong-owner", + "allowed": false, + "reason": null, + "event_index": null, + "summary_id": "session-summary_truncation:summary" + } + ] + }, + { + "fault_id": "duplicate_retry", + "case_id": "partial_retry", + "detected": true, + "differences": [ + { + "case_id": "partial_retry", + "left_backend": "in_memory", + "right_backend": "injected:duplicate_retry", + "session_id": "session-partial_retry", + "path": "$.session.events[2]", + "left_value": "", + "right_value": { + "actions": { + "artifact_delta": {}, + "state_delta": { + "recovery_status": "complete" + } + }, + "author": "assistant", + "content": { + "parts": [ + { + "text": "finished answer" + } + ] + }, + "id": "retry-event", + "invocation_id": "partial_retry-invocation-03", + "is_summary": false, + "model_flags": 1, + "partial": false, + "timestamp": 4102444803.0, + "turn_complete": true, + "version": 0, + "visible": true + }, + "allowed": false, + "reason": null, + "event_index": 2, + "summary_id": null + } + ] + } + ], + "metrics": { + "false_positive_rate": 0.0, + "injected_detection_rate": 1.0, + "summary_fault_detection_rate": 1.0, + "duration_seconds": 0.540597 + } +} diff --git a/tests/sessions/test_redis_session_service.py b/tests/sessions/test_redis_session_service.py index 8269b1862..f035f5f76 100644 --- a/tests/sessions/test_redis_session_service.py +++ b/tests/sessions/test_redis_session_service.py @@ -13,13 +13,10 @@ from __future__ import annotations -import json import time from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch -import pytest - from trpc_agent_sdk.events import Event from trpc_agent_sdk.sessions._redis_session_service import RedisSessionService from trpc_agent_sdk.sessions._session import Session @@ -113,6 +110,22 @@ def _create_service(config=None): return svc +def test_default_redis_service_decodes_text_responses(): + """Redis session state is read as text by default, not raw bytes.""" + with patch("trpc_agent_sdk.sessions._redis_session_service.RedisStorage") as mock_storage: + RedisSessionService(db_url="redis://localhost:6379") + + assert mock_storage.call_args.kwargs["decode_responses"] is True + + +def test_explicit_redis_service_bytes_responses_are_preserved(): + """Callers can opt back into raw Redis byte responses.""" + with patch("trpc_agent_sdk.sessions._redis_session_service.RedisStorage") as mock_storage: + RedisSessionService(db_url="redis://localhost:6379", decode_responses=False) + + assert mock_storage.call_args.kwargs["decode_responses"] is False + + class TestRedisCreateSession: async def test_create_basic(self): diff --git a/tests/sessions/test_replay_consistency.py b/tests/sessions/test_replay_consistency.py new file mode 100644 index 000000000..807fdc5e4 --- /dev/null +++ b/tests/sessions/test_replay_consistency.py @@ -0,0 +1,281 @@ +# 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. +"""Acceptance tests for Session/Memory/Summary replay consistency.""" + +from __future__ import annotations + +import json +import os +import time +from pathlib import Path +from unittest.mock import AsyncMock +from unittest.mock import patch + +import pytest + +from trpc_agent_sdk.storage import RedisCommand +from trpc_agent_sdk.storage import RedisCondition + +from .replay_cases import REPLAY_CASES +from .replay_harness import INJECTED_FAULTS +from .replay_harness import canonical_report +from .replay_harness import create_in_memory_backend +from .replay_harness import create_mock_redis_backend +from .replay_harness import create_redis_backend +from .replay_harness import create_sql_backend +from .replay_harness import create_sqlite_backend +from .replay_harness import execute_case +from .replay_harness import run_replay_matrix +from .replay_harness import stable_report_signature +from .replay_harness import write_report + + +async def test_replay_matrix_meets_acceptance(tmp_path): + """All dependency-free backends must agree and detect every corruption.""" + + report = await run_replay_matrix([ + await create_in_memory_backend(), + await create_sqlite_backend(), + ]) + report_path = tmp_path / "session_memory_summary_diff_report.json" + write_report(report, report_path) + persisted_report = json.loads(report_path.read_text(encoding="utf-8")) + committed_report_path = Path(__file__).with_name("session_memory_summary_diff_report.json") + committed_report = json.loads(committed_report_path.read_text(encoding="utf-8")) + + assert persisted_report["case_count"] == 10 + assert persisted_report["backends"] == ["in_memory", "sqlite"] + assert len(persisted_report["normal_cases"]) == 10 + assert len(persisted_report["injected_cases"]) == 10 + assert persisted_report["metrics"]["false_positive_rate"] <= 0.05 + assert persisted_report["metrics"]["injected_detection_rate"] == 1.0 + assert persisted_report["metrics"]["summary_fault_detection_rate"] == 1.0 + assert persisted_report["metrics"]["duration_seconds"] <= 30 + assert all(result["passed"] for result in persisted_report["normal_cases"]) + + all_injected_diffs = [ + difference for result in persisted_report["injected_cases"] for difference in result["differences"] + if not difference["allowed"] + ] + assert all(difference["session_id"] and difference["path"].startswith("$") for difference in all_injected_diffs) + assert all( + any(difference["event_index"] is not None for difference in result["differences"]) + for result in persisted_report["injected_cases"] + if result["fault_id"] in {"missing_event", "event_order", "tool_response", "duplicate_retry"}) + assert all( + any(difference["summary_id"] for difference in result["differences"]) + for result in persisted_report["injected_cases"] if result["fault_id"].startswith("summary_")) + assert stable_report_signature(persisted_report) == stable_report_signature(committed_report) + + +async def test_in_memory_lightweight_mode(): + """All traces can run without SQL, Redis, Docker, or a network model.""" + + backend = await create_in_memory_backend() + started_cases = [] + snapshots = {} + started_at = time.perf_counter() + try: + for case in REPLAY_CASES: + snapshot = await execute_case(backend, case) + assert snapshot["session"]["id"] == f"session-{case.case_id}" + assert all(snapshot["checks"].values()) + snapshots[case.case_id] = snapshot + started_cases.append(case.case_id) + finally: + await backend.close() + + final_state = snapshots["state_overwrite"]["session"]["state"] + assert final_state["phase"] == "done" + assert final_state["counter"] == 3 + assert snapshots["memory_roundtrip"]["memory"]["preference-token-oolong"] + + update_snapshot = snapshots["summary_update"] + assert [revision["version"] for revision in update_snapshot["summary_history"]] == [1, 2] + assert all(revision["updated_after_previous"] for revision in update_snapshot["summary_history"]) + assert update_snapshot["summary"]["summary_text"] == "summary update version two" + + truncation_snapshot = snapshots["summary_truncation"] + assert truncation_snapshot["summary"]["session_id"] == "session-summary_truncation" + assert truncation_snapshot["session"]["events"][0]["is_summary"] + assert truncation_snapshot["session"]["historical_events"] + assert [event["content"]["parts"][0]["text"] for event in truncation_snapshot["session"]["events"][-2:] + ] == ["follow-up after compression", "answer after compression"] + + partial_snapshot = snapshots["partial_retry"] + final_events = partial_snapshot["session"]["events"] + assert [event["id"] for event in final_events].count("retry-event") == 1 + assert final_events[-1]["content"]["parts"][0]["text"] == "finished answer" + assert partial_snapshot["session"]["state"]["recovery_status"] == "complete" + assert "temp:retry_buffer" not in partial_snapshot["session"]["state"] + + assert len(started_cases) == 10 + assert time.perf_counter() - started_at <= 30 + + +async def test_redis_integration_or_mock_fallback(): + """Use configured Redis or exercise Redis services with fakeredis.""" + + redis_url = os.getenv("TRPC_REPLAY_REDIS_URL") + if not redis_url: + pytest.importorskip("fakeredis") + redis_backend = await create_redis_backend(redis_url) if redis_url else await create_mock_redis_backend() + report = await run_replay_matrix([await create_in_memory_backend(), redis_backend]) + assert all(result["passed"] for result in report["normal_cases"]) + assert report["metrics"]["injected_detection_rate"] == 1.0 + assert report["backends"][1] == ("redis" if redis_url else "redis_mock") + + +async def test_redis_mock_uses_storage_query_and_ttl_paths(): + """The network-free Redis backend still exercises RedisStorage behavior.""" + + pytest.importorskip("fakeredis") + backend = await create_mock_redis_backend(enable_ttl=True) + memory_case = next(case for case in REPLAY_CASES if case.case_id == "memory_roundtrip") + session_storage = backend.session_service._redis_storage # pylint: disable=protected-access + memory_storage = backend.memory_service._redis_storage # pylint: disable=protected-access + try: + snapshot = await execute_case(backend, memory_case) + assert snapshot["memory"]["preference-token-oolong"] + async with memory_storage.create_db_session() as redis_session: + keys = redis_session.keys("*") + assert keys + assert any(redis_session.ttl(key) >= 0 for key in keys) + finally: + await backend.close() + assert session_storage.closed + assert memory_storage.closed + + +async def test_redis_mock_supports_storage_query_types(): + """RedisStorage query dispatch remains valid for every supported data type.""" + + pytest.importorskip("fakeredis") + backend = await create_mock_redis_backend() + storage = backend.memory_service._redis_storage # pylint: disable=protected-access + try: + async with storage.create_db_session() as redis_session: + commands = ( + RedisCommand(method="set", args=("probe:string", b'"value"')), + RedisCommand(method="hset", args=("probe:hash", "field", "value")), + RedisCommand(method="rpush", args=("probe:list", "first", "second")), + RedisCommand(method="sadd", args=("probe:set", "member")), + RedisCommand(method="zadd", args=("probe:zset", { + "member": 1.0 + })), + ) + for command in commands: + await storage.execute_command(redis_session, command) + result = dict(await storage.query(redis_session, "probe:*", RedisCondition())) + result = {key.decode() if isinstance(key, bytes) else key: value for key, value in result.items()} + result["probe:zset"] = [(member.decode() if isinstance(member, bytes) else member, score) + for member, score in result["probe:zset"]] + + assert result["probe:string"] == "value" + assert result["probe:hash"] == {"field": "value"} + assert result["probe:list"] == ["first", "second"] + assert result["probe:set"] == {"member"} + assert result["probe:zset"] == [("member", 1.0)] + finally: + await backend.close() + + +async def test_replay_matrix_rejects_empty_inputs(): + """Reusable matrix validation must fail clearly instead of dividing by zero.""" + + invalid_inputs = ( + ({ + "cases": () + }, "At least one replay case"), + ({ + "cases": (REPLAY_CASES[0], REPLAY_CASES[0]) + }, "Replay case IDs must be unique"), + ({ + "injected_faults": () + }, "At least one injected fault"), + ({ + "injected_faults": (INJECTED_FAULTS[0], INJECTED_FAULTS[0]) + }, "Injected fault IDs must be unique"), + ({ + "injected_faults": (INJECTED_FAULTS[0], ) + }, "At least one summary fault"), + ) + for kwargs, message in invalid_inputs: + backend = await create_in_memory_backend() + try: + with pytest.raises(ValueError, match=message): + await run_replay_matrix([backend], **kwargs) + finally: + await backend.close() + + +async def test_replay_matrix_closes_backends_on_validation_error(): + """Input validation must not leak already-created service resources.""" + + backend = await create_in_memory_backend() + memory_close = AsyncMock(wraps=backend.memory_service.close) + session_close = AsyncMock(wraps=backend.session_service.close) + with patch.object(backend.memory_service, "close", memory_close), patch.object(backend.session_service, "close", + session_close): + with pytest.raises(ValueError, match="At least one replay case"): + await run_replay_matrix([backend], cases=()) + memory_close.assert_awaited_once() + session_close.assert_awaited_once() + + +async def test_canonical_report_does_not_hide_dynamic_named_business_fields(): + """Only snapshot-level runtime fields may be normalized in a report.""" + + report = { + "normal_cases": [{ + "differences": [{ + "path": "$.session.state", + "allowed": False, + "left_value": { + "summary_timestamp": 1 + }, + "right_value": { + "summary_timestamp": 2 + }, + }] + }], + "injected_cases": [], + "metrics": {}, + } + + normalized = canonical_report(report) + difference = normalized["normal_cases"][0]["differences"][0] + assert difference["left_value"] != difference["right_value"] + + +async def test_sqlite_fallbacks_use_isolated_database_files(): + """Each fallback owns a database file and removes it on close.""" + + first = await create_sqlite_backend(name="first_sqlite") + second = await create_sqlite_backend(name="second_sqlite") + first_path = Path(first.session_service._sql_storage._db_engine.url.database) # pylint: disable=protected-access + second_path = Path(second.session_service._sql_storage._db_engine.url.database) # pylint: disable=protected-access + try: + assert first_path != second_path + assert first_path.is_file() + assert second_path.is_file() + finally: + await second.close() + await first.close() + assert not first_path.exists() + assert not second_path.exists() + + +async def test_sql_integration_or_sqlite_fallback(): + """Use configured SQL or exercise the generic SQL factory with SQLite.""" + + sql_url = os.getenv("TRPC_REPLAY_SQL_URL") + sql_backend = (await create_sql_backend(sql_url, name="sql") if sql_url else await create_sqlite_backend( + name="sql_fallback")) + report = await run_replay_matrix([await create_in_memory_backend(), sql_backend]) + assert all(result["passed"] for result in report["normal_cases"]) + assert report["metrics"]["injected_detection_rate"] == 1.0 + assert report["backends"][1] == ("sql" if sql_url else "sql_fallback") diff --git a/tests/storage/test_redis.py b/tests/storage/test_redis.py index 2fbcd923b..fdaa7e740 100644 --- a/tests/storage/test_redis.py +++ b/tests/storage/test_redis.py @@ -165,6 +165,11 @@ def test_deserialize_value_string(self, sync_storage): result = sync_storage._deserialize_value(b"simple string") assert result == "simple string" + def test_deserialize_value_text_response(self, sync_storage): + """Decode responses from pools configured with decode_responses=True.""" + result = sync_storage._deserialize_value('{"key": "value"}') + assert result == {"key": "value"} + def test_deserialize_value_unicode_error(self, sync_storage): """Test deserializing with unicode error.""" # Invalid UTF-8 bytes @@ -495,6 +500,57 @@ async def test_execute_command_with_method(self, async_storage): assert result == "OK" mock_conn.set.assert_called_once_with('key', 'value') + @pytest.mark.asyncio + async def test_execute_command_variadic_hset_uses_raw_command(self, sync_storage): + """Route Redis's multi-field HSET form around redis-py's single-pair helper.""" + mock_conn = MagicMock() + mock_conn.hset = MagicMock() + mock_conn.execute_command = MagicMock(return_value=2) + + command = RedisCommand(method='hset', args=('key', 'field1', 'value1', 'field2', 'value2')) + + result = await sync_storage.execute_command(mock_conn, command) + + assert result == 2 + mock_conn.hset.assert_not_called() + mock_conn.execute_command.assert_called_once_with('HSET', 'key', 'field1', 'value1', 'field2', 'value2') + + @pytest.mark.asyncio + async def test_execute_command_variadic_hset_derives_expire_key(self, sync_storage): + """Multi-field HSET still applies TTL to its first Redis key argument.""" + mock_conn = MagicMock() + mock_conn.hset = MagicMock() + mock_conn.execute_command = MagicMock(return_value=2) + mock_conn.expire = MagicMock() + expire = RedisExpire(ttl=Ttl(enable=True, ttl_seconds=30, cleanup_interval_seconds=1)) + command = RedisCommand( + method='hset', + args=('key', 'field1', 'value1', 'field2', 'value2'), + expire=expire, + ) + + result = await sync_storage.execute_command(mock_conn, command) + + assert result == 2 + assert command.expire.key == 'key' + mock_conn.execute_command.assert_called_once_with('HSET', 'key', 'field1', 'value1', 'field2', 'value2') + mock_conn.expire.assert_called_once_with('key', 30) + + @pytest.mark.asyncio + async def test_execute_command_single_field_hset_uses_helper(self, sync_storage): + """Keep Redis-py's normal single-field HSET path covered.""" + mock_conn = MagicMock() + mock_conn.hset = MagicMock(return_value=1) + mock_conn.execute_command = MagicMock() + + command = RedisCommand(method='hset', args=('key', 'field', 'value')) + + result = await sync_storage.execute_command(mock_conn, command) + + assert result == 1 + mock_conn.hset.assert_called_once_with('key', 'field', 'value') + mock_conn.execute_command.assert_not_called() + @pytest.mark.asyncio async def test_execute_command_without_method(self, async_storage): """Test executing command without existing method.""" diff --git a/trpc_agent_sdk/memory/_redis_memory_service.py b/trpc_agent_sdk/memory/_redis_memory_service.py index c9d933033..338ba004a 100644 --- a/trpc_agent_sdk/memory/_redis_memory_service.py +++ b/trpc_agent_sdk/memory/_redis_memory_service.py @@ -51,6 +51,7 @@ def _create_storage(self, db_url: str, is_async: bool, **kwargs: Any) -> RedisSt Subclasses override this factory to preserve memory behavior while selecting a deployment-specific Redis client. """ + kwargs.setdefault("decode_responses", True) return RedisStorage(is_async=is_async, redis_url=db_url, **kwargs) @override diff --git a/trpc_agent_sdk/sessions/_redis_session_service.py b/trpc_agent_sdk/sessions/_redis_session_service.py index 8bec47af1..6462fd115 100644 --- a/trpc_agent_sdk/sessions/_redis_session_service.py +++ b/trpc_agent_sdk/sessions/_redis_session_service.py @@ -94,6 +94,7 @@ def _create_storage(self, db_url: str, is_async: bool, **kwargs: Any) -> RedisSt Subclasses override this factory to retain the session semantics while selecting a different Redis deployment client, such as Redis Cluster. """ + kwargs.setdefault("decode_responses", True) return RedisStorage(is_async=is_async, redis_url=db_url, **kwargs) @override diff --git a/trpc_agent_sdk/storage/_redis.py b/trpc_agent_sdk/storage/_redis.py index 6cfbb88c7..2771ca7e9 100644 --- a/trpc_agent_sdk/storage/_redis.py +++ b/trpc_agent_sdk/storage/_redis.py @@ -134,21 +134,24 @@ def _serialize_value(self, value: Any) -> str: return str(value) return json.dumps(value, default=str) - def _deserialize_value(self, value: Optional[bytes]) -> Any: - """Deserialize value from Redis bytes.""" + def _deserialize_value(self, value: Optional[Union[bytes, str]]) -> Any: + """Deserialize a Redis byte or text response.""" if value is None: return None - try: - value_str = value.decode('utf-8') - # Try to parse as JSON first + if isinstance(value, str): + value_str = value + else: try: - return json.loads(value_str) - except json.JSONDecodeError: - # If not JSON, return as string - return value_str - except UnicodeDecodeError: - return value + value_str = value.decode('utf-8') + except UnicodeDecodeError: + return value + + # Try to parse as JSON first; plain Redis strings remain text. + try: + return json.loads(value_str) + except json.JSONDecodeError: + return value_str @override async def add(self, conn: RedisSession, data: RedisCommand) -> None: @@ -266,7 +269,11 @@ async def execute_command(self, conn: RedisSession, command: RedisCommand) -> An lower_method = command.method.lower() upper_method = command.method.upper() method = getattr(conn, lower_method, None) - if method: + # redis-py exposes HSET as a single-pair helper while Redis itself + # accepts multiple field/value pairs. The session services use the + # native variadic form, so route it through execute_command(). + use_raw_hset = lower_method == 'hset' and len(command.args) > 3 and not command.kwargs + if method and not use_raw_hset: ret = method(*command.args, **command.kwargs) else: ret = conn.execute_command(upper_method, *command.args, **command.kwargs)