52 lines
1.9 KiB
Python
52 lines
1.9 KiB
Python
"""Redis short-term conversation context for the customer-service Agent."""
|
|||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import json
|
||
|
|
from inspect import isawaitable
|
||
|
|
|
||
|
|
|
||
|
|
async def _config(config_getter, key: str, default):
|
||
|
|
value = config_getter(key, str(default))
|
||
|
|
if isawaitable(value):
|
||
|
|
value = await value
|
||
|
|
return type(default)(value)
|
||
|
|
|
||
|
|
|
||
|
|
class RedisConversationContext:
|
||
|
|
def __init__(self, redis, *, config_getter, token_counter=None):
|
||
|
|
self.redis = redis
|
||
|
|
self.config_getter = config_getter
|
||
|
|
self.token_counter = token_counter or (lambda text: max(1, len(text) // 4))
|
||
|
|
|
||
|
|
async def append(self, session_id: str, role: str, content: str) -> None:
|
||
|
|
key = f"session:{session_id}:messages"
|
||
|
|
payload = json.dumps(
|
||
|
|
{"role": role, "content": content}, ensure_ascii=False
|
||
|
|
)
|
||
|
|
await self.redis.rpush(key, payload)
|
||
|
|
ttl = await _config(self.config_getter, "agent.customer.session.ttl", 1800)
|
||
|
|
await self.redis.expire(key, ttl)
|
||
|
|
await self._trim(key)
|
||
|
|
|
||
|
|
async def get(self, session_id: str) -> list[dict]:
|
||
|
|
raw_messages = await self.redis.lrange(
|
||
|
|
f"session:{session_id}:messages", 0, -1
|
||
|
|
)
|
||
|
|
return [json.loads(raw) for raw in raw_messages if raw]
|
||
|
|
|
||
|
|
async def _trim(self, key: str) -> None:
|
||
|
|
limit = await _config(
|
||
|
|
self.config_getter, "agent.customer.session_max_token", 4096
|
||
|
|
)
|
||
|
|
raw_messages = [raw for raw in await self.redis.lrange(key, 0, -1) if raw]
|
||
|
|
total = 0
|
||
|
|
keep_from = len(raw_messages)
|
||
|
|
for index in range(len(raw_messages) - 1, -1, -1):
|
||
|
|
message = json.loads(raw_messages[index])
|
||
|
|
total += self.token_counter(message["content"])
|
||
|
|
if total > limit:
|
||
|
|
keep_from = index + 1
|
||
|
|
break
|
||
|
|
keep_from = index
|
||
|
|
await self.redis.ltrim(key, keep_from, -1)
|