diff --git a/tests/sessions/test_redis_session_service.py b/tests/sessions/test_redis_session_service.py index 8269b1862..b768d93d7 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 @@ -84,10 +91,20 @@ async def execute_command(self, session, command): elif method == 'hset': key = args[0] pairs = args[1:] + if command.kwargs.get("mapping"): + 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 + # 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 i in range(0, len(pairs), 2): - self._hash_store[key][pairs[i]] = pairs[i + 1] + for pair_key, pair_value in encoded_pairs: + self._hash_store[key][pair_key] = pair_value return True elif method == 'hgetall': key = args[0] @@ -273,6 +290,38 @@ 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) + + 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): svc = _create_service() session = await svc.create_session(app_name="app", user_id="user", session_id="s1") 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/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) 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