"""Redis sessions owned by authenticated client Agent customers.""" from __future__ import annotations import time import uuid from inspect import isawaitable from agent.customer_agent.session import SessionOwnershipError from service.client_agent.archive_schedule import SessionArchiveSchedule 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 ClientSessionService: """Create and verify Redis sessions bound to a customer ID.""" def __init__(self, redis, *, config_getter, clock=time.time): self.redis = redis self.config_getter = config_getter self.clock = clock self.archive_schedule = SessionArchiveSchedule(redis) async def _archive_grace(self) -> int: """读取 Redis 消息归档保护缓冲时间,默认 60 秒。""" return await _config( self.config_getter, "agent.customer.session.archive_grace", 60 ) async def create_session(self, customer_id: int) -> str: """Create a session whose Redis owner value is the current customer.""" session_id = uuid.uuid4().hex now = self.clock() ttl = await _config(self.config_getter, "agent.customer.session.ttl", 1800) max_lifetime = await _config( self.config_getter, "agent.customer.session.max_lifetime", 86400 ) grace = await self._archive_grace() effective_ttl = min(ttl, max_lifetime) await self.redis.set( f"session:{session_id}", str(customer_id), ex=effective_ttl + grace ) await self.redis.rpush(f"session:{session_id}:messages", "") await self.redis.expire( f"session:{session_id}:messages", effective_ttl + grace ) await self.archive_schedule.register( session_id=session_id, created_at=now, last_activity_at=now, archive_due_at=now + ttl, absolute_expire_at=now + max_lifetime, ) return session_id async def verify_session_ownership(self, session_id: str, *, customer_id: int) -> None: """Reject missing sessions and sessions owned by another customer.""" owner = await self.redis.get(f"session:{session_id}") if owner is None: raise SessionOwnershipError(404, "会话不存在或已过期") if isinstance(owner, bytes): owner = owner.decode() if str(owner) != str(customer_id): raise SessionOwnershipError(403, "无权访问该会话") async def touch_session(self, session_id: str, *, customer_id: int) -> None: """刷新所有权 TTL 和自动归档截止时间,但不突破最长生命周期。""" await self.verify_session_ownership(session_id, customer_id=customer_id) now = self.clock() meta = await self.redis.hgetall( self.archive_schedule.meta_key(session_id) ) absolute_expire_at = float(meta.get("absolute_expire_at", now)) remaining = int(absolute_expire_at - now) if remaining <= 0: raise SessionOwnershipError(404, "会话不存在或已过期") ttl = await _config(self.config_getter, "agent.customer.session.ttl", 1800) effective_ttl = min(ttl, remaining) grace = await self._archive_grace() await self.redis.expire(f"session:{session_id}", effective_ttl + grace) await self.archive_schedule.touch( session_id=session_id, last_activity_at=now, archive_due_at=now + effective_ttl, ) async def forget_session(self, session_id: str) -> None: """移除会话的自动归档调度记录。""" await self.archive_schedule.remove(session_id) async def consume_chat_quota(self, session_id: str, *, customer_id: int) -> int | None: """Apply the existing per-session rate limit with a customer-scoped key.""" window = await _config( self.config_getter, "agent.customer.rate_limit.window_sec", 60 ) maximum = await _config( self.config_getter, "agent.customer.rate_limit.max_requests", 20 ) key = f"rate:limit:client:{customer_id}:{session_id}:chat" now = self.clock() await self.redis.zremrangebyscore(key, 0, (now - window) * 1000) count = await self.redis.zcard(key) if count >= maximum: entries = getattr(self.redis, "sorted_sets", {}).get(key, []) oldest = min((score for score, _ in entries), default=now * 1000) return max(1, int((oldest / 1000 + window) - now)) await self.redis.zadd(key, {uuid.uuid4().hex: now * 1000}) await self.redis.expire(key, window) return None __all__ = ["ClientSessionService", "SessionOwnershipError"]