perf(agent): wake on state change instead of 500ms polling (#305)

Co-authored-by: 0xallam <ahmed39652003@gmail.com>
This commit is contained in:
Bala
2026-05-04 03:30:41 +01:00
committed by GitHub
parent 6942ecb33e
commit e1f38f8339
2 changed files with 15 additions and 2 deletions
+1 -1
View File
@@ -282,7 +282,7 @@ class BaseAgent(metaclass=AgentMeta):
return return
await asyncio.sleep(0.5) await self.state.wait_for_wake(timeout=0.5)
async def _enter_waiting_state( async def _enter_waiting_state(
self, self,
+14 -1
View File
@@ -1,8 +1,9 @@
import asyncio
import uuid import uuid
from datetime import UTC, datetime from datetime import UTC, datetime
from typing import Any from typing import Any
from pydantic import BaseModel, Field from pydantic import BaseModel, Field, PrivateAttr
def _generate_agent_id() -> str: def _generate_agent_id() -> str:
@@ -40,6 +41,8 @@ class AgentState(BaseModel):
errors: list[str] = Field(default_factory=list) errors: list[str] = Field(default_factory=list)
_wake_event: asyncio.Event = PrivateAttr(default_factory=asyncio.Event)
def increment_iteration(self) -> None: def increment_iteration(self) -> None:
self.iteration += 1 self.iteration += 1
self.last_updated = datetime.now(UTC).isoformat() self.last_updated = datetime.now(UTC).isoformat()
@@ -52,6 +55,8 @@ class AgentState(BaseModel):
message["thinking_blocks"] = thinking_blocks message["thinking_blocks"] = thinking_blocks
self.messages.append(message) self.messages.append(message)
self.last_updated = datetime.now(UTC).isoformat() self.last_updated = datetime.now(UTC).isoformat()
if self.waiting_for_input:
self._wake_event.set()
def add_action(self, action: dict[str, Any]) -> None: def add_action(self, action: dict[str, Any]) -> None:
self.actions_taken.append( self.actions_taken.append(
@@ -109,6 +114,14 @@ class AgentState(BaseModel):
if new_task: if new_task:
self.task = new_task self.task = new_task
self.last_updated = datetime.now(UTC).isoformat() self.last_updated = datetime.now(UTC).isoformat()
self._wake_event.set()
async def wait_for_wake(self, timeout: float = 0.5) -> None:
try:
await asyncio.wait_for(self._wake_event.wait(), timeout=timeout)
self._wake_event.clear()
except TimeoutError:
pass
def has_reached_max_iterations(self) -> bool: def has_reached_max_iterations(self) -> bool:
return self.iteration >= self.max_iterations return self.iteration >= self.max_iterations