57 lines
1.7 KiB
Python
57 lines
1.7 KiB
Python
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()
|