| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137 |
- """Agent 热状态、频率控制和运行事件的 Redis 适配器。"""
- import json
- from typing import Protocol, cast
- from redis import Redis
- from app.core.config import Settings
- from app.harness.events import RuntimeEvent
- from app.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()
|