|
@@ -2,8 +2,9 @@ from __future__ import annotations
|
|
|
|
|
|
|
|
import asyncio
|
|
import asyncio
|
|
|
import json
|
|
import json
|
|
|
|
|
+from collections import OrderedDict
|
|
|
from collections.abc import Hashable, Iterable, Sequence
|
|
from collections.abc import Hashable, Iterable, Sequence
|
|
|
-from dataclasses import dataclass
|
|
|
|
|
|
|
+from dataclasses import dataclass, replace
|
|
|
from typing import Any
|
|
from typing import Any
|
|
|
|
|
|
|
|
from agent_lab.application.events.kernel import EventKernel
|
|
from agent_lab.application.events.kernel import EventKernel
|
|
@@ -36,17 +37,31 @@ class EventBatchExecutor:
|
|
|
*,
|
|
*,
|
|
|
max_parallel_events: int = 4,
|
|
max_parallel_events: int = 4,
|
|
|
batch_timeout_seconds: float = 15.0,
|
|
batch_timeout_seconds: float = 15.0,
|
|
|
|
|
+ terminal_grace_seconds: float = 2.0,
|
|
|
|
|
+ max_replay_entries: int = 100,
|
|
|
) -> None:
|
|
) -> None:
|
|
|
if max_parallel_events < 1:
|
|
if max_parallel_events < 1:
|
|
|
raise ValueError("max_parallel_events must be positive")
|
|
raise ValueError("max_parallel_events must be positive")
|
|
|
if batch_timeout_seconds <= 0:
|
|
if batch_timeout_seconds <= 0:
|
|
|
raise ValueError("batch_timeout_seconds must be positive")
|
|
raise ValueError("batch_timeout_seconds must be positive")
|
|
|
|
|
+ if terminal_grace_seconds <= 0:
|
|
|
|
|
+ raise ValueError("terminal_grace_seconds must be positive")
|
|
|
|
|
+ if max_replay_entries < 1:
|
|
|
|
|
+ raise ValueError("max_replay_entries must be positive")
|
|
|
self.kernel = kernel
|
|
self.kernel = kernel
|
|
|
self.max_parallel_events = max_parallel_events
|
|
self.max_parallel_events = max_parallel_events
|
|
|
self.batch_timeout_seconds = batch_timeout_seconds
|
|
self.batch_timeout_seconds = batch_timeout_seconds
|
|
|
|
|
+ self.terminal_grace_seconds = terminal_grace_seconds
|
|
|
|
|
+ self.max_replay_entries = max_replay_entries
|
|
|
self._semaphore = asyncio.Semaphore(max_parallel_events)
|
|
self._semaphore = asyncio.Semaphore(max_parallel_events)
|
|
|
self._conflict_locks: dict[str, asyncio.Lock] = {}
|
|
self._conflict_locks: dict[str, asyncio.Lock] = {}
|
|
|
- self._replay_results: dict[tuple[Hashable, str, str], EventResult] = {}
|
|
|
|
|
|
|
+ self._replay_results: OrderedDict[
|
|
|
|
|
+ tuple[Hashable, str, str], EventResult
|
|
|
|
|
+ ] = OrderedDict()
|
|
|
|
|
+ self._inflight: dict[
|
|
|
|
|
+ tuple[Hashable, str, str], asyncio.Task[_ExecutionOutcome]
|
|
|
|
|
+ ] = {}
|
|
|
|
|
+ self._state_lock = asyncio.Lock()
|
|
|
|
|
|
|
|
async def execute(
|
|
async def execute(
|
|
|
self,
|
|
self,
|
|
@@ -63,7 +78,8 @@ class EventBatchExecutor:
|
|
|
indexes_by_key: dict[str, list[int]] = {}
|
|
indexes_by_key: dict[str, list[int]] = {}
|
|
|
coalesced_count = 0
|
|
coalesced_count = 0
|
|
|
for index, request in enumerate(requests):
|
|
for index, request in enumerate(requests):
|
|
|
- key = self._request_key(request)
|
|
|
|
|
|
|
+ definition = self.kernel.registry.definition(request.name)
|
|
|
|
|
+ key = self._request_key(request, definition)
|
|
|
indexes = indexes_by_key.get(key)
|
|
indexes = indexes_by_key.get(key)
|
|
|
if indexes is not None:
|
|
if indexes is not None:
|
|
|
indexes.append(index)
|
|
indexes.append(index)
|
|
@@ -72,53 +88,63 @@ class EventBatchExecutor:
|
|
|
indexes_by_key[key] = [index]
|
|
indexes_by_key[key] = [index]
|
|
|
items.append(_BatchItem(key=key, request=request))
|
|
items.append(_BatchItem(key=key, request=request))
|
|
|
|
|
|
|
|
- results_by_key: dict[str, EventResult] = {}
|
|
|
|
|
|
|
+ outcomes_by_key: dict[str, _ExecutionOutcome] = {}
|
|
|
replayed_count = 0
|
|
replayed_count = 0
|
|
|
- pending: list[_BatchItem] = []
|
|
|
|
|
- for item in items:
|
|
|
|
|
- replay_key = (scope, item.request.id, item.key)
|
|
|
|
|
- replayed = self._replay_results.get(replay_key)
|
|
|
|
|
- if replayed is None:
|
|
|
|
|
- pending.append(item)
|
|
|
|
|
- continue
|
|
|
|
|
- results_by_key[item.key] = replayed
|
|
|
|
|
- replayed_count += 1
|
|
|
|
|
-
|
|
|
|
|
nonterminal: list[_BatchItem] = []
|
|
nonterminal: list[_BatchItem] = []
|
|
|
terminal: list[_BatchItem] = []
|
|
terminal: list[_BatchItem] = []
|
|
|
- for item in pending:
|
|
|
|
|
|
|
+ for item in items:
|
|
|
definition = self.kernel.registry.definition(item.request.name)
|
|
definition = self.kernel.registry.definition(item.request.name)
|
|
|
if definition is not None and definition.terminal:
|
|
if definition is not None and definition.terminal:
|
|
|
terminal.append(item)
|
|
terminal.append(item)
|
|
|
else:
|
|
else:
|
|
|
nonterminal.append(item)
|
|
nonterminal.append(item)
|
|
|
|
|
|
|
|
- deadline = asyncio.get_running_loop().time() + self.batch_timeout_seconds
|
|
|
|
|
|
|
+ loop = asyncio.get_running_loop()
|
|
|
|
|
+ batch_deadline = loop.time() + self.batch_timeout_seconds
|
|
|
deadline_exceeded = False
|
|
deadline_exceeded = False
|
|
|
- for phase in (nonterminal, terminal):
|
|
|
|
|
- outcomes = await asyncio.gather(
|
|
|
|
|
|
|
+ phases = (
|
|
|
|
|
+ (nonterminal, batch_deadline, True),
|
|
|
|
|
+ (terminal, None, False),
|
|
|
|
|
+ )
|
|
|
|
|
+ for phase, deadline, count_deadline_exceeded in phases:
|
|
|
|
|
+ if deadline is None:
|
|
|
|
|
+ deadline = loop.time() + self.terminal_grace_seconds
|
|
|
|
|
+ phase_results = await asyncio.gather(
|
|
|
*[
|
|
*[
|
|
|
- self._execute_one(
|
|
|
|
|
- item.request,
|
|
|
|
|
|
|
+ self._execute_replayable(
|
|
|
|
|
+ item,
|
|
|
|
|
+ scope=scope,
|
|
|
enabled_names=enabled_names,
|
|
enabled_names=enabled_names,
|
|
|
context=context,
|
|
context=context,
|
|
|
deadline=deadline,
|
|
deadline=deadline,
|
|
|
|
|
+ count_deadline_exceeded=count_deadline_exceeded,
|
|
|
)
|
|
)
|
|
|
for item in phase
|
|
for item in phase
|
|
|
]
|
|
]
|
|
|
)
|
|
)
|
|
|
- for item, outcome in zip(phase, outcomes, strict=True):
|
|
|
|
|
- results_by_key[item.key] = outcome.result
|
|
|
|
|
- self._replay_results[(scope, item.request.id, item.key)] = (
|
|
|
|
|
- outcome.result
|
|
|
|
|
- )
|
|
|
|
|
|
|
+ for item, (outcome, replayed) in zip(
|
|
|
|
|
+ phase,
|
|
|
|
|
+ phase_results,
|
|
|
|
|
+ strict=True,
|
|
|
|
|
+ ):
|
|
|
|
|
+ outcomes_by_key[item.key] = outcome
|
|
|
|
|
+ replayed_count += int(replayed)
|
|
|
deadline_exceeded = deadline_exceeded or outcome.batch_timed_out
|
|
deadline_exceeded = deadline_exceeded or outcome.batch_timed_out
|
|
|
|
|
|
|
|
ordered: list[EventResult | None] = [None] * len(requests)
|
|
ordered: list[EventResult | None] = [None] * len(requests)
|
|
|
for key, indexes in indexes_by_key.items():
|
|
for key, indexes in indexes_by_key.items():
|
|
|
- result = results_by_key[key]
|
|
|
|
|
|
|
+ result = outcomes_by_key[key].result
|
|
|
|
|
+ primary_id = requests[indexes[0]].id
|
|
|
for index in indexes:
|
|
for index in indexes:
|
|
|
- ordered[index] = result
|
|
|
|
|
|
|
+ ordered[index] = self._clone_result(
|
|
|
|
|
+ result,
|
|
|
|
|
+ requests[index],
|
|
|
|
|
+ deduplicated_from=(
|
|
|
|
|
+ primary_id if requests[index].id != primary_id else None
|
|
|
|
|
+ ),
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+ await self._cache_ordered_results(scope, requests, indexes_by_key, ordered)
|
|
|
|
|
|
|
|
return EventBatchResult(
|
|
return EventBatchResult(
|
|
|
results=tuple(result for result in ordered if result is not None),
|
|
results=tuple(result for result in ordered if result is not None),
|
|
@@ -127,6 +153,124 @@ class EventBatchExecutor:
|
|
|
deadline_exceeded=deadline_exceeded,
|
|
deadline_exceeded=deadline_exceeded,
|
|
|
)
|
|
)
|
|
|
|
|
|
|
|
|
|
+ def release_scope(self, scope: Hashable) -> None:
|
|
|
|
|
+ for key in tuple(self._replay_results):
|
|
|
|
|
+ if key[0] == scope:
|
|
|
|
|
+ self._replay_results.pop(key, None)
|
|
|
|
|
+
|
|
|
|
|
+ async def _execute_replayable(
|
|
|
|
|
+ self,
|
|
|
|
|
+ item: _BatchItem,
|
|
|
|
|
+ *,
|
|
|
|
|
+ scope: Hashable,
|
|
|
|
|
+ enabled_names: Iterable[str] | None,
|
|
|
|
|
+ context: EventExecutionContext | None,
|
|
|
|
|
+ deadline: float,
|
|
|
|
|
+ count_deadline_exceeded: bool,
|
|
|
|
|
+ ) -> tuple[_ExecutionOutcome, bool]:
|
|
|
|
|
+ replay_key = (scope, item.request.id, item.key)
|
|
|
|
|
+ async with self._state_lock:
|
|
|
|
|
+ replayed = self._replay_results.get(replay_key)
|
|
|
|
|
+ if replayed is not None:
|
|
|
|
|
+ self._replay_results.move_to_end(replay_key)
|
|
|
|
|
+ return self._outcome_from_cached(replayed), True
|
|
|
|
|
+ task = self._inflight.get(replay_key)
|
|
|
|
|
+ reused = task is not None
|
|
|
|
|
+ if task is None:
|
|
|
|
|
+ task = asyncio.create_task(
|
|
|
|
|
+ self._run_and_store(
|
|
|
|
|
+ replay_key,
|
|
|
|
|
+ item.request,
|
|
|
|
|
+ enabled_names=enabled_names,
|
|
|
|
|
+ context=context,
|
|
|
|
|
+ deadline=deadline,
|
|
|
|
|
+ count_deadline_exceeded=count_deadline_exceeded,
|
|
|
|
|
+ )
|
|
|
|
|
+ )
|
|
|
|
|
+ self._inflight[replay_key] = task
|
|
|
|
|
+ return await asyncio.shield(task), reused
|
|
|
|
|
+
|
|
|
|
|
+ async def _run_and_store(
|
|
|
|
|
+ self,
|
|
|
|
|
+ replay_key: tuple[Hashable, str, str],
|
|
|
|
|
+ request: EventRequest,
|
|
|
|
|
+ *,
|
|
|
|
|
+ enabled_names: Iterable[str] | None,
|
|
|
|
|
+ context: EventExecutionContext | None,
|
|
|
|
|
+ deadline: float,
|
|
|
|
|
+ count_deadline_exceeded: bool,
|
|
|
|
|
+ ) -> _ExecutionOutcome:
|
|
|
|
|
+ outcome: _ExecutionOutcome | None = None
|
|
|
|
|
+ try:
|
|
|
|
|
+ outcome = await self._execute_one(
|
|
|
|
|
+ request,
|
|
|
|
|
+ enabled_names=enabled_names,
|
|
|
|
|
+ context=context,
|
|
|
|
|
+ deadline=deadline,
|
|
|
|
|
+ count_deadline_exceeded=count_deadline_exceeded,
|
|
|
|
|
+ )
|
|
|
|
|
+ finally:
|
|
|
|
|
+ async with self._state_lock:
|
|
|
|
|
+ if self._inflight.get(replay_key) is asyncio.current_task():
|
|
|
|
|
+ self._inflight.pop(replay_key, None)
|
|
|
|
|
+ if outcome is not None:
|
|
|
|
|
+ self._store_replay_locked(replay_key, outcome.result)
|
|
|
|
|
+ assert outcome is not None
|
|
|
|
|
+ return outcome
|
|
|
|
|
+
|
|
|
|
|
+ async def _cache_ordered_results(
|
|
|
|
|
+ self,
|
|
|
|
|
+ scope: Hashable,
|
|
|
|
|
+ requests: Sequence[EventRequest],
|
|
|
|
|
+ indexes_by_key: dict[str, list[int]],
|
|
|
|
|
+ ordered: list[EventResult | None],
|
|
|
|
|
+ ) -> None:
|
|
|
|
|
+ async with self._state_lock:
|
|
|
|
|
+ for key, indexes in indexes_by_key.items():
|
|
|
|
|
+ for index in indexes:
|
|
|
|
|
+ result = ordered[index]
|
|
|
|
|
+ assert result is not None
|
|
|
|
|
+ self._store_replay_locked(
|
|
|
|
|
+ (scope, requests[index].id, key),
|
|
|
|
|
+ result,
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+ def _store_replay_locked(
|
|
|
|
|
+ self,
|
|
|
|
|
+ replay_key: tuple[Hashable, str, str],
|
|
|
|
|
+ result: EventResult,
|
|
|
|
|
+ ) -> None:
|
|
|
|
|
+ self._replay_results[replay_key] = result
|
|
|
|
|
+ self._replay_results.move_to_end(replay_key)
|
|
|
|
|
+ while len(self._replay_results) > self.max_replay_entries:
|
|
|
|
|
+ self._replay_results.popitem(last=False)
|
|
|
|
|
+
|
|
|
|
|
+ def _outcome_from_cached(self, result: EventResult) -> _ExecutionOutcome:
|
|
|
|
|
+ return _ExecutionOutcome(
|
|
|
|
|
+ result,
|
|
|
|
|
+ batch_timed_out=(
|
|
|
|
|
+ result.status is EventStatus.TIMEOUT
|
|
|
|
|
+ and result.error == "batch deadline exceeded"
|
|
|
|
|
+ ),
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
|
|
+ def _clone_result(
|
|
|
|
|
+ self,
|
|
|
|
|
+ result: EventResult,
|
|
|
|
|
+ request: EventRequest,
|
|
|
|
|
+ *,
|
|
|
|
|
+ deduplicated_from: str | None,
|
|
|
|
|
+ ) -> EventResult:
|
|
|
|
|
+ payload = dict(result.payload)
|
|
|
|
|
+ if deduplicated_from is not None:
|
|
|
|
|
+ payload["deduplicated_from"] = deduplicated_from
|
|
|
|
|
+ return replace(
|
|
|
|
|
+ result,
|
|
|
|
|
+ event_id=request.id,
|
|
|
|
|
+ raw_arguments=request.raw_arguments,
|
|
|
|
|
+ payload=payload,
|
|
|
|
|
+ )
|
|
|
|
|
+
|
|
|
async def _execute_one(
|
|
async def _execute_one(
|
|
|
self,
|
|
self,
|
|
|
request: EventRequest,
|
|
request: EventRequest,
|
|
@@ -134,6 +278,7 @@ class EventBatchExecutor:
|
|
|
enabled_names: Iterable[str] | None,
|
|
enabled_names: Iterable[str] | None,
|
|
|
context: EventExecutionContext | None,
|
|
context: EventExecutionContext | None,
|
|
|
deadline: float,
|
|
deadline: float,
|
|
|
|
|
+ count_deadline_exceeded: bool,
|
|
|
) -> _ExecutionOutcome:
|
|
) -> _ExecutionOutcome:
|
|
|
definition = self.kernel.registry.definition(request.name)
|
|
definition = self.kernel.registry.definition(request.name)
|
|
|
remaining = deadline - asyncio.get_running_loop().time()
|
|
remaining = deadline - asyncio.get_running_loop().time()
|
|
@@ -146,7 +291,7 @@ class EventBatchExecutor:
|
|
|
EventStatus.TIMEOUT,
|
|
EventStatus.TIMEOUT,
|
|
|
"batch deadline exceeded",
|
|
"batch deadline exceeded",
|
|
|
),
|
|
),
|
|
|
- batch_timed_out=True,
|
|
|
|
|
|
|
+ batch_timed_out=count_deadline_exceeded,
|
|
|
)
|
|
)
|
|
|
batch_limited = event_timeout is None or remaining <= event_timeout
|
|
batch_limited = event_timeout is None or remaining <= event_timeout
|
|
|
timeout = remaining if event_timeout is None else min(remaining, event_timeout)
|
|
timeout = remaining if event_timeout is None else min(remaining, event_timeout)
|
|
@@ -179,7 +324,7 @@ class EventBatchExecutor:
|
|
|
EventStatus.TIMEOUT,
|
|
EventStatus.TIMEOUT,
|
|
|
error,
|
|
error,
|
|
|
),
|
|
),
|
|
|
- batch_timed_out=batch_limited,
|
|
|
|
|
|
|
+ batch_timed_out=(batch_limited and count_deadline_exceeded),
|
|
|
)
|
|
)
|
|
|
except Exception as exc:
|
|
except Exception as exc:
|
|
|
return _ExecutionOutcome(
|
|
return _ExecutionOutcome(
|
|
@@ -226,10 +371,25 @@ class EventBatchExecutor:
|
|
|
**metadata,
|
|
**metadata,
|
|
|
)
|
|
)
|
|
|
|
|
|
|
|
- def _request_key(self, request: EventRequest) -> str:
|
|
|
|
|
|
|
+ def _request_key(
|
|
|
|
|
+ self,
|
|
|
|
|
+ request: EventRequest,
|
|
|
|
|
+ definition: EventDefinition | None,
|
|
|
|
|
+ ) -> str:
|
|
|
|
|
+ arguments_value = self._safe_arguments(request.arguments)
|
|
|
|
|
+ key_kind = "original"
|
|
|
|
|
+ if definition is not None and definition.normalizer is not None:
|
|
|
|
|
+ try:
|
|
|
|
|
+ normalized = definition.normalizer(arguments_value)
|
|
|
|
|
+ if not isinstance(normalized, dict):
|
|
|
|
|
+ raise TypeError("normalizer returned non-object")
|
|
|
|
|
+ arguments_value = self._safe_arguments(normalized)
|
|
|
|
|
+ key_kind = "normalized"
|
|
|
|
|
+ except Exception:
|
|
|
|
|
+ arguments_value = self._safe_arguments(request.arguments)
|
|
|
try:
|
|
try:
|
|
|
arguments = json.dumps(
|
|
arguments = json.dumps(
|
|
|
- request.arguments,
|
|
|
|
|
|
|
+ arguments_value,
|
|
|
ensure_ascii=False,
|
|
ensure_ascii=False,
|
|
|
sort_keys=True,
|
|
sort_keys=True,
|
|
|
separators=(",", ":"),
|
|
separators=(",", ":"),
|
|
@@ -238,7 +398,7 @@ class EventBatchExecutor:
|
|
|
except (TypeError, ValueError):
|
|
except (TypeError, ValueError):
|
|
|
arguments = request.raw_arguments
|
|
arguments = request.raw_arguments
|
|
|
return json.dumps(
|
|
return json.dumps(
|
|
|
- [request.id, request.name, request.source.value, arguments],
|
|
|
|
|
|
|
+ [request.name, request.source.value, key_kind, arguments],
|
|
|
ensure_ascii=False,
|
|
ensure_ascii=False,
|
|
|
separators=(",", ":"),
|
|
separators=(",", ":"),
|
|
|
)
|
|
)
|