import json import unittest from agent.customer_agent.context import RedisConversationContext class ContextRedis: def __init__(self): self.values = {} async def rpush(self, key, value): self.values.setdefault(key, []).append(value) async def lrange(self, key, start, end): values = self.values.get(key, [])[start:] return values if end == -1 else values[: end - start + 1] async def ltrim(self, key, start, end): values = self.values.get(key, [])[start:] self.values[key] = values if end == -1 else values[: end - start + 1] async def expire(self, key, seconds): pass class ContextTests(unittest.IsolatedAsyncioTestCase): async def test_appends_messages_and_trims_oldest_when_token_limit_is_exceeded(self): redis = ContextRedis() context = RedisConversationContext( redis, config_getter={ "agent.customer.session_max_token": "5", "agent.customer.session.ttl": "120", }.get, token_counter=lambda text: len(text), ) await context.append("s1", "user", "123") await context.append("s1", "assistant", "456") await context.append("s1", "user", "789") messages = await context.get("s1") self.assertEqual([item["content"] for item in messages], ["789"]) async def test_message_payload_is_structured_json(self): redis = ContextRedis() context = RedisConversationContext(redis, config_getter={}.get) await context.append("s1", "user", "基金是什么") raw = redis.values["session:s1:messages"][0] self.assertEqual(json.loads(raw)["role"], "user") if __name__ == "__main__": unittest.main()