From 385ae3626eb164d0df6c8447921ab1915d0a4590 Mon Sep 17 00:00:00 2001 From: MelsonLi <3010632904@qq.com> Date: Thu, 30 Jul 2026 18:54:36 +0800 Subject: [PATCH 1/3] fix: use redis hset mapping for scoped state --- tests/sessions/replay/backends.py | 8 ++++++-- tests/sessions/test_redis_session_service.py | 3 +++ .../sessions/_redis_session_service.py | 16 ++++------------ 3 files changed, 13 insertions(+), 14 deletions(-) diff --git a/tests/sessions/replay/backends.py b/tests/sessions/replay/backends.py index 5c638a0fe..7abe1dc62 100644 --- a/tests/sessions/replay/backends.py +++ b/tests/sessions/replay/backends.py @@ -75,8 +75,12 @@ def sqlite_backend(db_url: str = "sqlite:///:memory:") -> ReplayBackend: def redis_backend(url: str) -> ReplayBackend: - svc = RedisSessionService(db_url=url, summarizer_manager=_manager(), session_config=_session_config(), is_async=True) - mem = RedisMemoryService(db_url=url, enabled=True, is_async=True) + svc = RedisSessionService(db_url=url, + summarizer_manager=_manager(), + session_config=_session_config(), + is_async=True, + decode_responses=True) + mem = RedisMemoryService(db_url=url, enabled=True, is_async=True, decode_responses=True) return ReplayBackend("redis", svc, mem) diff --git a/tests/sessions/test_redis_session_service.py b/tests/sessions/test_redis_session_service.py index 8269b1862..3ffc2fb7f 100644 --- a/tests/sessions/test_redis_session_service.py +++ b/tests/sessions/test_redis_session_service.py @@ -86,6 +86,9 @@ async def execute_command(self, session, command): pairs = args[1:] if key not in self._hash_store: self._hash_store[key] = {} + if command.kwargs.get("mapping"): + self._hash_store[key].update(command.kwargs["mapping"]) + return True for i in range(0, len(pairs), 2): self._hash_store[key][pairs[i]] = pairs[i + 1] return True diff --git a/trpc_agent_sdk/sessions/_redis_session_service.py b/trpc_agent_sdk/sessions/_redis_session_service.py index 8bec47af1..c1ae1eb3e 100644 --- a/trpc_agent_sdk/sessions/_redis_session_service.py +++ b/trpc_agent_sdk/sessions/_redis_session_service.py @@ -285,13 +285,9 @@ async def _update_app_state(self, redis_session: RedisSession, app_name: str, await self._refresh_ttl(redis_session, key) return app_state - # Use HSET with TTL if TTL is configured, otherwise use HSET - args = [key] - for k, v in app_state.items(): - args.extend([k, v]) - command = RedisCommand(method='hset', - args=tuple(args), + args=(key, ), + kwargs={"mapping": app_state}, expire=RedisExpire(key=key, ttl=self._session_config.ttl)) await self._redis_storage.execute_command(redis_session, command) @@ -325,13 +321,9 @@ async def _update_user_state(self, redis_session: RedisSession, app_name: str, u await self._refresh_ttl(redis_session, key) return user_state - # Use HSET with TTL if TTL is configured, otherwise use HSET - args = [key] - for k, v in user_state.items(): - args.extend([k, v]) - command = RedisCommand(method='hset', - args=tuple(args), + args=(key, ), + kwargs={"mapping": user_state}, expire=RedisExpire(key=key, ttl=self._session_config.ttl)) await self._redis_storage.execute_command(redis_session, command) From c9d1b4c46171543b3e89849d1b965f8e20caeb5c Mon Sep 17 00:00:00 2001 From: MelsonLi <3010632904@qq.com> Date: Thu, 30 Jul 2026 19:43:43 +0800 Subject: [PATCH 2/3] test: mirror redis hash value encoding in mock --- tests/sessions/test_redis_session_service.py | 40 +++++++++++++++++++- 1 file changed, 38 insertions(+), 2 deletions(-) diff --git a/tests/sessions/test_redis_session_service.py b/tests/sessions/test_redis_session_service.py index 3ffc2fb7f..a76bb6125 100644 --- a/tests/sessions/test_redis_session_service.py +++ b/tests/sessions/test_redis_session_service.py @@ -19,6 +19,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from redis.exceptions import DataError from trpc_agent_sdk.events import Event from trpc_agent_sdk.sessions._redis_session_service import RedisSessionService @@ -66,6 +67,12 @@ def __init__(self): async def create_db_session(self): yield MagicMock() + @staticmethod + def _encode_hash_value(value): + if isinstance(value, bool): + raise DataError("Invalid input of type: 'bool'. Convert to a bytes, string, int or float first.") + return str(value) + async def execute_command(self, session, command): method = command.method args = command.args @@ -87,10 +94,12 @@ async def execute_command(self, session, command): if key not in self._hash_store: self._hash_store[key] = {} if command.kwargs.get("mapping"): - self._hash_store[key].update(command.kwargs["mapping"]) + self._hash_store[key].update({ + k: self._encode_hash_value(v) for k, v in command.kwargs["mapping"].items() + }) return True for i in range(0, len(pairs), 2): - self._hash_store[key][pairs[i]] = pairs[i + 1] + self._hash_store[key][pairs[i]] = self._encode_hash_value(pairs[i + 1]) return True elif method == 'hgetall': key = args[0] @@ -276,6 +285,33 @@ async def test_append_with_state_delta(self): assert stored.state[f"{State.USER_PREFIX}user_key"] == "uv" await svc.close() + async def test_append_with_numeric_scoped_state_matches_redis_encoding(self): + svc = _create_service() + session = await svc.create_session(app_name="app", user_id="user", session_id="s1") + event = _make_event(state_delta={ + "session_count": 3, + f"{State.APP_PREFIX}app_count": 7, + f"{State.USER_PREFIX}user_ratio": 0.5, + }) + + await svc.append_event(session, event) + + stored = await svc.get_session(app_name="app", user_id="user", session_id="s1") + assert stored.state["session_count"] == 3 + assert stored.state[f"{State.APP_PREFIX}app_count"] == "7" + assert stored.state[f"{State.USER_PREFIX}user_ratio"] == "0.5" + await svc.close() + + async def test_append_with_bool_scoped_state_matches_redis_rejection(self): + svc = _create_service() + session = await svc.create_session(app_name="app", user_id="user", session_id="s1") + event = _make_event(state_delta={f"{State.APP_PREFIX}app_enabled": True}) + + with pytest.raises(DataError): + await svc.append_event(session, event) + + await svc.close() + async def test_append_does_not_persist_merged_or_temp_state_in_session_json(self): svc = _create_service() session = await svc.create_session(app_name="app", user_id="user", session_id="s1") From 92b851f98d73256f7b5a14dee5c2646cf6503f84 Mon Sep 17 00:00:00 2001 From: MelsonLi <3010632904@qq.com> Date: Thu, 30 Jul 2026 23:09:15 +0800 Subject: [PATCH 3/3] fix: decode redis responses by default --- tests/sessions/replay/backends.py | 8 ++----- tests/sessions/test_redis_session_service.py | 22 ++++++++++++++------ tests/storage/test_redis.py | 16 ++++++++++++-- trpc_agent_sdk/storage/_redis.py | 1 + 4 files changed, 33 insertions(+), 14 deletions(-) diff --git a/tests/sessions/replay/backends.py b/tests/sessions/replay/backends.py index 7abe1dc62..5c638a0fe 100644 --- a/tests/sessions/replay/backends.py +++ b/tests/sessions/replay/backends.py @@ -75,12 +75,8 @@ def sqlite_backend(db_url: str = "sqlite:///:memory:") -> ReplayBackend: def redis_backend(url: str) -> ReplayBackend: - svc = RedisSessionService(db_url=url, - summarizer_manager=_manager(), - session_config=_session_config(), - is_async=True, - decode_responses=True) - mem = RedisMemoryService(db_url=url, enabled=True, is_async=True, decode_responses=True) + svc = RedisSessionService(db_url=url, summarizer_manager=_manager(), session_config=_session_config(), is_async=True) + mem = RedisMemoryService(db_url=url, enabled=True, is_async=True) return ReplayBackend("redis", svc, mem) diff --git a/tests/sessions/test_redis_session_service.py b/tests/sessions/test_redis_session_service.py index a76bb6125..b768d93d7 100644 --- a/tests/sessions/test_redis_session_service.py +++ b/tests/sessions/test_redis_session_service.py @@ -91,15 +91,20 @@ async def execute_command(self, session, command): elif method == 'hset': key = args[0] pairs = args[1:] - if key not in self._hash_store: - self._hash_store[key] = {} if command.kwargs.get("mapping"): - self._hash_store[key].update({ + encoded_mapping = { k: self._encode_hash_value(v) for k, v in command.kwargs["mapping"].items() - }) + } + if key not in self._hash_store: + self._hash_store[key] = {} + self._hash_store[key].update(encoded_mapping) return True - for i in range(0, len(pairs), 2): - self._hash_store[key][pairs[i]] = self._encode_hash_value(pairs[i + 1]) + # Kept for compatibility with older mock callers; production scoped state uses mapping. + encoded_pairs = [(pairs[i], self._encode_hash_value(pairs[i + 1])) for i in range(0, len(pairs), 2)] + if key not in self._hash_store: + self._hash_store[key] = {} + for pair_key, pair_value in encoded_pairs: + self._hash_store[key][pair_key] = pair_value return True elif method == 'hgetall': key = args[0] @@ -310,6 +315,11 @@ async def test_append_with_bool_scoped_state_matches_redis_rejection(self): with pytest.raises(DataError): await svc.append_event(session, event) + stored = await svc.get_session(app_name="app", user_id="user", session_id="s1") + assert stored is not None + assert stored.state == {} + assert stored.events == [] + assert svc._redis_storage._hash_store == {} await svc.close() async def test_append_does_not_persist_merged_or_temp_state_in_session_json(self): diff --git a/tests/storage/test_redis.py b/tests/storage/test_redis.py index 2fbcd923b..87fdb12a7 100644 --- a/tests/storage/test_redis.py +++ b/tests/storage/test_redis.py @@ -46,6 +46,18 @@ def test_init(self, redis_url): assert storage._kwargs == {"max_connections": 10, "decode_responses": True} assert storage._redis_pool is None + def test_init_defaults_to_decoded_responses(self, redis_url): + """Test RedisStorage decodes responses by default.""" + storage = RedisStorage(redis_url=redis_url, is_async=True) + + assert storage._kwargs == {"decode_responses": True} + + def test_init_allows_raw_responses(self, redis_url): + """Test RedisStorage can still opt into raw Redis bytes.""" + storage = RedisStorage(redis_url=redis_url, is_async=True, decode_responses=False) + + assert storage._kwargs == {"decode_responses": False} + @pytest.mark.asyncio async def test_create_redis_engine_async(self, async_storage): """Test creating async Redis connection pool.""" @@ -54,7 +66,7 @@ async def test_create_redis_engine_async(self, async_storage): await async_storage.create_redis_engine() - mock_pool.from_url.assert_called_once_with(async_storage._redis_url) + mock_pool.from_url.assert_called_once_with(async_storage._redis_url, decode_responses=True) assert async_storage._redis_pool is not None @pytest.mark.asyncio @@ -65,7 +77,7 @@ async def test_create_redis_engine_sync(self, sync_storage): await sync_storage.create_redis_engine() - mock_pool.from_url.assert_called_once_with(sync_storage._redis_url) + mock_pool.from_url.assert_called_once_with(sync_storage._redis_url, decode_responses=True) assert sync_storage._redis_pool is not None @pytest.mark.asyncio diff --git a/trpc_agent_sdk/storage/_redis.py b/trpc_agent_sdk/storage/_redis.py index 6cfbb88c7..47e811467 100644 --- a/trpc_agent_sdk/storage/_redis.py +++ b/trpc_agent_sdk/storage/_redis.py @@ -98,6 +98,7 @@ def __init__(self, redis_url: str, is_async: bool = False, **kwargs: Any) -> Non super().__init__() self._redis_url = redis_url self._is_async = is_async + kwargs.setdefault("decode_responses", True) self._kwargs = kwargs self._redis_pool: Optional[RedisConnectionPool] = None