agent_state.py 4.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137
  1. """Agent 热状态、频率控制和运行事件的 Redis 适配器。"""
  2. import json
  3. from typing import Protocol, cast
  4. from redis import Redis
  5. from zbt.core.config import Settings
  6. from zbt.harness.events import RuntimeEvent
  7. from zbt.infrastructure.redis.keys import RedisKeyBuilder
  8. class AgentStateStore(Protocol):
  9. def allow_request(self, principal_id: str) -> bool: ...
  10. def touch_thread(
  11. self,
  12. thread_id: str,
  13. *,
  14. owner_type: str,
  15. owner_id: str,
  16. persona: str,
  17. ) -> None: ...
  18. def append_event(self, event: RuntimeEvent) -> None: ...
  19. def list_events(self, run_id: str) -> list[RuntimeEvent]: ...
  20. def close(self) -> None: ...
  21. class InMemoryAgentStateStore:
  22. def __init__(self, rate_limit: int = 20) -> None:
  23. self._rate_limit = rate_limit
  24. self._request_counts: dict[str, int] = {}
  25. self._threads: dict[str, dict[str, str]] = {}
  26. self._events: dict[str, list[RuntimeEvent]] = {}
  27. def allow_request(self, principal_id: str) -> bool:
  28. count = self._request_counts.get(principal_id, 0) + 1
  29. self._request_counts[principal_id] = count
  30. return count <= self._rate_limit
  31. def touch_thread(
  32. self,
  33. thread_id: str,
  34. *,
  35. owner_type: str,
  36. owner_id: str,
  37. persona: str,
  38. ) -> None:
  39. self._threads[thread_id] = {
  40. "owner_type": owner_type,
  41. "owner_id": owner_id,
  42. "persona": persona,
  43. }
  44. def append_event(self, event: RuntimeEvent) -> None:
  45. self._events.setdefault(event.run_id, []).append(event)
  46. def list_events(self, run_id: str) -> list[RuntimeEvent]:
  47. return list(self._events.get(run_id, []))
  48. def close(self) -> None:
  49. return None
  50. class RedisAgentStateStore:
  51. def __init__(self, settings: Settings) -> None:
  52. self._client: Redis = Redis.from_url(
  53. settings.redis_url,
  54. decode_responses=True,
  55. )
  56. self._keys = RedisKeyBuilder(settings.redis_prefix)
  57. self._rate_limit = settings.agent_rate_limit
  58. self._rate_window_seconds = settings.agent_rate_window_seconds
  59. self._hot_ttl_seconds = settings.agent_hot_state_ttl_seconds
  60. self._event_ttl_seconds = settings.agent_event_ttl_seconds
  61. def allow_request(self, principal_id: str) -> bool:
  62. key = self._keys.agent_rate(principal_id)
  63. count = cast(
  64. int,
  65. self._client.eval(
  66. """
  67. local current = redis.call('INCR', KEYS[1])
  68. if current == 1 then
  69. redis.call('EXPIRE', KEYS[1], ARGV[1])
  70. end
  71. return current
  72. """,
  73. 1,
  74. key,
  75. self._rate_window_seconds,
  76. ),
  77. )
  78. return count <= self._rate_limit
  79. def touch_thread(
  80. self,
  81. thread_id: str,
  82. *,
  83. owner_type: str,
  84. owner_id: str,
  85. persona: str,
  86. ) -> None:
  87. payload = json.dumps(
  88. {
  89. "owner_type": owner_type,
  90. "owner_id": owner_id,
  91. "persona": persona,
  92. },
  93. ensure_ascii=False,
  94. separators=(",", ":"),
  95. )
  96. self._client.set(
  97. self._keys.agent_thread_hot(thread_id),
  98. payload,
  99. ex=self._hot_ttl_seconds,
  100. )
  101. def append_event(self, event: RuntimeEvent) -> None:
  102. key = self._keys.agent_run_events(event.run_id)
  103. pipe = self._client.pipeline()
  104. pipe.rpush(key, event.model_dump_json())
  105. pipe.expire(key, self._event_ttl_seconds)
  106. pipe.execute()
  107. def list_events(self, run_id: str) -> list[RuntimeEvent]:
  108. raw_items = cast(
  109. list[str],
  110. self._client.lrange(self._keys.agent_run_events(run_id), 0, -1),
  111. )
  112. return [RuntimeEvent.model_validate_json(item) for item in raw_items]
  113. def close(self) -> None:
  114. self._client.close()