|
|
@@ -6,6 +6,7 @@ import json
|
|
|
from collections import OrderedDict
|
|
|
from collections.abc import Hashable, Iterable, Sequence
|
|
|
from dataclasses import dataclass, replace
|
|
|
+from itertools import count
|
|
|
from typing import Any
|
|
|
|
|
|
from agent_lab.application.events.kernel import EventKernel
|
|
|
@@ -20,6 +21,7 @@ from agent_lab.application.events.models import (
|
|
|
|
|
|
|
|
|
_ReplayKey = tuple[Hashable, int, str, str]
|
|
|
+_SCOPE_TOKEN_COUNTER = count(1)
|
|
|
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
|
@@ -63,7 +65,7 @@ class EventBatchExecutor:
|
|
|
OrderedDict()
|
|
|
)
|
|
|
self._inflight: dict[_ReplayKey, asyncio.Task[_ExecutionOutcome]] = {}
|
|
|
- self._scope_generations: dict[Hashable, int] = {}
|
|
|
+ self._scope_tokens: dict[Hashable, int] = {}
|
|
|
self._state_lock = asyncio.Lock()
|
|
|
|
|
|
async def execute(
|
|
|
@@ -76,7 +78,10 @@ class EventBatchExecutor:
|
|
|
) -> EventBatchResult:
|
|
|
if not requests:
|
|
|
return EventBatchResult()
|
|
|
- generation = self._scope_generations.get(scope, 0)
|
|
|
+ scope_token = self._scope_tokens.get(scope)
|
|
|
+ if scope_token is None:
|
|
|
+ scope_token = next(_SCOPE_TOKEN_COUNTER)
|
|
|
+ self._scope_tokens[scope] = scope_token
|
|
|
|
|
|
items: list[_BatchItem] = []
|
|
|
indexes_by_key: dict[str, list[int]] = {}
|
|
|
@@ -118,7 +123,7 @@ class EventBatchExecutor:
|
|
|
self._execute_replayable(
|
|
|
item,
|
|
|
scope=scope,
|
|
|
- generation=generation,
|
|
|
+ scope_token=scope_token,
|
|
|
enabled_names=enabled_names,
|
|
|
context=context,
|
|
|
deadline=deadline,
|
|
|
@@ -151,7 +156,7 @@ class EventBatchExecutor:
|
|
|
|
|
|
await self._cache_ordered_results(
|
|
|
scope,
|
|
|
- generation,
|
|
|
+ scope_token,
|
|
|
requests,
|
|
|
indexes_by_key,
|
|
|
outcomes_by_key,
|
|
|
@@ -166,13 +171,14 @@ class EventBatchExecutor:
|
|
|
)
|
|
|
|
|
|
def release_scope(self, scope: Hashable) -> None:
|
|
|
- generation = self._scope_generations.get(scope, 0)
|
|
|
- self._scope_generations[scope] = generation + 1
|
|
|
+ scope_token = self._scope_tokens.pop(scope, None)
|
|
|
+ if scope_token is None:
|
|
|
+ return
|
|
|
for key in tuple(self._replay_results):
|
|
|
- if key[0] == scope and key[1] == generation:
|
|
|
+ if key[0] == scope and key[1] == scope_token:
|
|
|
self._replay_results.pop(key, None)
|
|
|
for key in tuple(self._inflight):
|
|
|
- if key[0] == scope and key[1] == generation:
|
|
|
+ if key[0] == scope and key[1] == scope_token:
|
|
|
self._inflight.pop(key, None)
|
|
|
|
|
|
async def _execute_replayable(
|
|
|
@@ -180,13 +186,13 @@ class EventBatchExecutor:
|
|
|
item: _BatchItem,
|
|
|
*,
|
|
|
scope: Hashable,
|
|
|
- generation: int,
|
|
|
+ scope_token: int,
|
|
|
enabled_names: Iterable[str] | None,
|
|
|
context: EventExecutionContext | None,
|
|
|
deadline: float,
|
|
|
count_deadline_exceeded: bool,
|
|
|
) -> tuple[_ExecutionOutcome, bool]:
|
|
|
- replay_key = (scope, generation, item.request.id, item.key)
|
|
|
+ replay_key = (scope, scope_token, item.request.id, item.key)
|
|
|
async with self._state_lock:
|
|
|
replayed = self._replay_results.get(replay_key)
|
|
|
if replayed is not None:
|
|
|
@@ -233,7 +239,7 @@ class EventBatchExecutor:
|
|
|
self._inflight.pop(replay_key, None)
|
|
|
if (
|
|
|
outcome is not None
|
|
|
- and self._scope_generations.get(replay_key[0], 0)
|
|
|
+ and self._scope_tokens.get(replay_key[0])
|
|
|
== replay_key[1]
|
|
|
):
|
|
|
self._store_replay_locked(replay_key, outcome)
|
|
|
@@ -243,14 +249,14 @@ class EventBatchExecutor:
|
|
|
async def _cache_ordered_results(
|
|
|
self,
|
|
|
scope: Hashable,
|
|
|
- generation: int,
|
|
|
+ scope_token: int,
|
|
|
requests: Sequence[EventRequest],
|
|
|
indexes_by_key: dict[str, list[int]],
|
|
|
outcomes_by_key: dict[str, _ExecutionOutcome],
|
|
|
ordered: list[EventResult | None],
|
|
|
) -> None:
|
|
|
async with self._state_lock:
|
|
|
- if self._scope_generations.get(scope, 0) != generation:
|
|
|
+ if self._scope_tokens.get(scope) != scope_token:
|
|
|
return
|
|
|
for key, indexes in indexes_by_key.items():
|
|
|
outcome = outcomes_by_key[key]
|
|
|
@@ -258,7 +264,7 @@ class EventBatchExecutor:
|
|
|
result = ordered[index]
|
|
|
assert result is not None
|
|
|
self._store_replay_locked(
|
|
|
- (scope, generation, requests[index].id, key),
|
|
|
+ (scope, scope_token, requests[index].id, key),
|
|
|
_ExecutionOutcome(
|
|
|
result,
|
|
|
batch_timed_out=outcome.batch_timed_out,
|
|
|
@@ -466,7 +472,10 @@ class EventBatchExecutor:
|
|
|
copied = json.loads(serialized)
|
|
|
except (TypeError, ValueError):
|
|
|
return None
|
|
|
- if not isinstance(copied, dict) or copied != arguments:
|
|
|
+ if not isinstance(copied, dict) or not _json_values_equal(
|
|
|
+ copied,
|
|
|
+ arguments,
|
|
|
+ ):
|
|
|
return None
|
|
|
return serialized
|
|
|
|
|
|
@@ -478,3 +487,18 @@ class EventBatchExecutor:
|
|
|
except (TypeError, ValueError):
|
|
|
return {}
|
|
|
return value if isinstance(value, dict) else {}
|
|
|
+
|
|
|
+
|
|
|
+def _json_values_equal(left: Any, right: Any) -> bool:
|
|
|
+ if type(left) is not type(right):
|
|
|
+ return False
|
|
|
+ if isinstance(left, dict):
|
|
|
+ return left.keys() == right.keys() and all(
|
|
|
+ _json_values_equal(left[key], right[key]) for key in left
|
|
|
+ )
|
|
|
+ if isinstance(left, list):
|
|
|
+ return len(left) == len(right) and all(
|
|
|
+ _json_values_equal(left_item, right_item)
|
|
|
+ for left_item, right_item in zip(left, right, strict=True)
|
|
|
+ )
|
|
|
+ return left == right
|