"""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)