Merge branch 'develop_feature_memory' of http://47.106.207.27:3000/AI260626/Mutual_Fund into develop_feature_memory
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
"""Logged-in client Agent domain layer."""
|
||||
|
||||
package_name = "client_agent"
|
||||
@@ -0,0 +1,7 @@
|
||||
"""Client Agent context aliases during the architecture bootstrap phase."""
|
||||
|
||||
from agent.customer_agent.context import RedisConversationContext
|
||||
|
||||
ClientConversationContext = RedisConversationContext
|
||||
|
||||
__all__ = ["ClientConversationContext"]
|
||||
@@ -0,0 +1,66 @@
|
||||
"""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
|
||||
|
||||
|
||||
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
|
||||
|
||||
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
|
||||
ttl = await _config(self.config_getter, "agent.customer.session.ttl", 1800)
|
||||
await self.redis.set(f"session:{session_id}", str(customer_id), ex=ttl)
|
||||
await self.redis.rpush(f"session:{session_id}:messages", "")
|
||||
await self.redis.expire(f"session:{session_id}:messages", ttl)
|
||||
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 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"]
|
||||
Reference in New Issue
Block a user