Files
Mutual_Fund/agent/client_agent/session.py
T

118 lines
4.8 KiB
Python
Raw Normal View History

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