"""Agent 热状态、频率控制和运行事件的 Redis 适配器。""" import json from typing import Protocol, cast from redis import Redis from zbt.core.config import Settings from zbt.harness.events import RuntimeEvent from zbt.infrastructure.redis.keys import RedisKeyBuilder class AgentStateStore(Protocol): def allow_request(self, principal_id: str) -> bool: ... def touch_thread( self, thread_id: str, *, owner_type: str, owner_id: str, persona: str, ) -> None: ... def append_event(self, event: RuntimeEvent) -> None: ... def list_events(self, run_id: str) -> list[RuntimeEvent]: ... def close(self) -> None: ... class InMemoryAgentStateStore: def __init__(self, rate_limit: int = 20) -> None: self._rate_limit = rate_limit self._request_counts: dict[str, int] = {} self._threads: dict[str, dict[str, str]] = {} self._events: dict[str, list[RuntimeEvent]] = {} def allow_request(self, principal_id: str) -> bool: count = self._request_counts.get(principal_id, 0) + 1 self._request_counts[principal_id] = count return count <= self._rate_limit def touch_thread( self, thread_id: str, *, owner_type: str, owner_id: str, persona: str, ) -> None: self._threads[thread_id] = { "owner_type": owner_type, "owner_id": owner_id, "persona": persona, } def append_event(self, event: RuntimeEvent) -> None: self._events.setdefault(event.run_id, []).append(event) def list_events(self, run_id: str) -> list[RuntimeEvent]: return list(self._events.get(run_id, [])) def close(self) -> None: return None class RedisAgentStateStore: def __init__(self, settings: Settings) -> None: self._client: Redis = Redis.from_url( settings.redis_url, decode_responses=True, ) self._keys = RedisKeyBuilder(settings.redis_prefix) self._rate_limit = settings.agent_rate_limit self._rate_window_seconds = settings.agent_rate_window_seconds self._hot_ttl_seconds = settings.agent_hot_state_ttl_seconds self._event_ttl_seconds = settings.agent_event_ttl_seconds def allow_request(self, principal_id: str) -> bool: key = self._keys.agent_rate(principal_id) count = cast( int, self._client.eval( """ local current = redis.call('INCR', KEYS[1]) if current == 1 then redis.call('EXPIRE', KEYS[1], ARGV[1]) end return current """, 1, key, self._rate_window_seconds, ), ) return count <= self._rate_limit def touch_thread( self, thread_id: str, *, owner_type: str, owner_id: str, persona: str, ) -> None: payload = json.dumps( { "owner_type": owner_type, "owner_id": owner_id, "persona": persona, }, ensure_ascii=False, separators=(",", ":"), ) self._client.set( self._keys.agent_thread_hot(thread_id), payload, ex=self._hot_ttl_seconds, ) def append_event(self, event: RuntimeEvent) -> None: key = self._keys.agent_run_events(event.run_id) pipe = self._client.pipeline() pipe.rpush(key, event.model_dump_json()) pipe.expire(key, self._event_ttl_seconds) pipe.execute() def list_events(self, run_id: str) -> list[RuntimeEvent]: raw_items = cast( list[str], self._client.lrange(self._keys.agent_run_events(run_id), 0, -1), ) return [RuntimeEvent.model_validate_json(item) for item in raw_items] def close(self) -> None: self._client.close()