Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 2 additions & 7 deletions src/agent_env/explorer/routers/common.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,4 @@
"""Shared pieces for the explorer's read routers: every list endpoint is
``DocumentStore.latest_per_id`` + ``count_distinct`` behind one factory, so no route
depends on a backend query language.
"""
"""Shared pieces for the explorer's read routers, using backend-agnostic store calls."""

from __future__ import annotations

Expand Down Expand Up @@ -139,12 +136,10 @@ def list_items(
total = len(matched)
items = matched[offset: offset + limit] if limit else matched[offset:]
else:
items = store.latest_per_id(
items, total = store.latest_per_id_page(
collection, Filter({}), id_field=id_field,
sort=Sort.by(sort_by, descending=descending), limit=limit, offset=offset,
)
# Count distinct ids, not rows — one entity (N versions) is one result.
total = store.count_distinct(collection, Filter({}), id_field=id_field)
return PaginatedResponse(
items=items, total=total, limit=limit, offset=offset,
has_more=offset + len(items) < total,
Expand Down
115 changes: 72 additions & 43 deletions src/agent_env/explorer/routers/runs.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,9 @@

logger = logging.getLogger(__name__)

INSTANCE_READ_BATCH_SIZE = 500
TASK_RUN_PAGE_SIZE = 500

router = APIRouter(prefix="/api/v1/tasks", tags=["runs"])

TASK_INSTANCES_COLLECTION = "task_instances"
Expand Down Expand Up @@ -278,39 +281,33 @@ def _all_task_runs(task_id: str) -> list:

out, offset = [], 0
while True:
page = run_store.list_runs(task_id=task_id, limit=500, offset=offset)
page = run_store.list_runs(task_id=task_id, limit=TASK_RUN_PAGE_SIZE, offset=offset)
out.extend(page)
if len(page) < 500:
if len(page) < TASK_RUN_PAGE_SIZE:
return out
offset += 500

offset += TASK_RUN_PAGE_SIZE

def _group_id_of(record) -> str:
"""The group a run belongs to.

A batch (``POST /runs``) stamps ``metadata.run_group_id`` on every run it starts.
A single run (``POST /run``, the "Start 1 Run" button) stamps nothing, so it is its
own group, keyed by its run id. ``list_run_groups`` and ``_instances_for_group``
MUST agree on this — when only the former applied the fallback, every single run
produced a group row that the join could never match, so the Rollouts table showed
the group with "0 runs" and no instance, for every non-batch run.
"""
return (record.overrides.get("metadata") or {}).get("run_group_id") or record.run_id
def _instances_for_runs(records: list) -> dict[str, dict]:
"""Load instance documents for a selected set of runs in bounded query batches."""
instance_ids = list(dict.fromkeys(record.instance_id for record in records))
instances: dict[str, dict] = {}
store = docs()
for start in range(0, len(instance_ids), INSTANCE_READ_BATCH_SIZE):
batch = instance_ids[start:start + INSTANCE_READ_BATCH_SIZE]
for instance in store.find_many_by_id(TASK_INSTANCES_COLLECTION, "instance_id", batch):
instance_id = instance.get("instance_id")
if instance_id is not None:
instances.setdefault(instance_id, instance)
return instances


def _instances_for_group(run_group_id: str, task_id: str) -> tuple[list[dict], dict]:
"""Every run in the group, joined to its task instance for live status, plus the
per-step ``{"done", "failed"}`` tally the Batch progress strip draws."""
def _group_instances(records: list, instances: dict[str, dict]) -> tuple[list[dict], dict]:
"""Join selected runs to one instance snapshot and tally each run's final step outcomes."""
out: list[dict] = []
step_counts: dict[str, dict[str, int]] = {}
for record in _all_task_runs(task_id):
if _group_id_of(record) != run_group_id:
continue
instance = docs().find_one(
TASK_INSTANCES_COLLECTION, Filter.of(instance_id=record.instance_id)
) or {}
# One outcome per step per run. A completion callback that raises records a
# failure beside the step's success, and the failure is the final verdict.
for record in records:
instance = instances.get(record.instance_id) or {}
outcome: dict[str, str] = {}
for step in instance.get("completed_steps") or []:
step_id, status = step.get("step_id"), step.get("status")
Expand All @@ -335,6 +332,26 @@ def _instances_for_group(run_group_id: str, task_id: str) -> tuple[list[dict], d
return out, step_counts


def _group_id_of(record) -> str:
"""The group a run belongs to.

A batch (``POST /runs``) stamps ``metadata.run_group_id`` on every run it starts.
A single run (``POST /run``, the "Start 1 Run" button) stamps nothing, so it is its
own group, keyed by its run id. ``list_run_groups`` and ``_instances_for_group``
MUST agree on this — when only the former applied the fallback, every single run
produced a group row that the join could never match, so the Rollouts table showed
the group with "0 runs" and no instance, for every non-batch run.
"""
return (record.overrides.get("metadata") or {}).get("run_group_id") or record.run_id


def _instances_for_group(run_group_id: str, task_id: str) -> tuple[list[dict], dict]:
"""Every run in the group, joined to its task instance for live status, plus the
per-step ``{"done", "failed"}`` tally the Batch progress strip draws."""
records = [record for record in _all_task_runs(task_id) if _group_id_of(record) == run_group_id]
return _group_instances(records, _instances_for_runs(records))


# Run-group funnel buckets keyed by RunStatus, so the tally tracks the enum — adding a
# status is a deliberate edit here, not a silent fall-through. CANCELED folds into
# "failed"; QUEUED is "provisioning" (accepted, not yet picked up by a worker).
Expand All @@ -349,23 +366,28 @@ def _instances_for_group(run_group_id: str, task_id: str) -> tuple[list[dict], d

def _group_status(task_id: str, run_group_id: str) -> dict:
instances, step_counts = _instances_for_group(run_group_id, task_id)
tally = {"completed": 0, "failed": 0, "running": 0, "provisioning": 0}
for inst in instances:
try:
run_status: Optional[RunStatus] = RunStatus(str(inst["status"]).upper())
except ValueError:
run_status = None
tally[_FUNNEL_BUCKET.get(run_status, "provisioning")] += 1
return {
"run_group_id": run_group_id,
"task_id": task_id,
"total": len(instances),
**tally,
**_funnel_counts(instances),
"step_counts": step_counts,
"instances": instances,
}


def _funnel_counts(instances: list[dict]) -> dict[str, int]:
"""Tally current run statuses into the UI's four funnel buckets."""
tally = {"completed": 0, "failed": 0, "running": 0, "provisioning": 0}
for inst in instances:
try:
run_status: Optional[RunStatus] = RunStatus(str(inst["status"]).upper())
except ValueError:
run_status = None
tally[_FUNNEL_BUCKET.get(run_status, "provisioning")] += 1
return tally


@router.get("/{task_id}/run-groups", response_model=PaginatedResponse)
def list_run_groups(
task_id: EntityId,
Expand All @@ -375,10 +397,13 @@ def list_run_groups(
) -> PaginatedResponse:
"""Run groups for a task, newest first — the Rollouts table."""
groups: dict[str, dict] = {}
for record in _all_task_runs(task_id):
records = _all_task_runs(task_id)
group_records: dict[str, list] = {}
for record in records:
gid = _group_id_of(record)
group_records.setdefault(gid, []).append(record)
if task_version is not None and record.task_version != task_version:
continue
gid = _group_id_of(record)
group = groups.setdefault(gid, {
"run_group_id": gid,
"task_id": task_id,
Expand All @@ -391,26 +416,30 @@ def list_run_groups(

ordered = sorted(groups.values(), key=lambda g: g["created_at_utc"] or "", reverse=True)
page = ordered[offset: offset + limit]
selected_records = [record for group in page for record in group_records[group["run_group_id"]]]
instance_snapshot = _instances_for_runs(selected_records)
items = []
for group in page:
status = _group_status(task_id, group["run_group_id"])
selected = group_records[group["run_group_id"]]
instances, step_counts = _group_instances(selected, instance_snapshot)
tally = _funnel_counts(instances)
# The list row nests the tally under `counts` and names its timestamp
# `earliest_created_at_utc` — unlike the flat shape /run-groups/{id} returns.
items.append({
"run_group_id": group["run_group_id"],
"task_id": task_id,
"task_version": group["task_version"],
"earliest_created_at_utc": group["created_at_utc"],
"total": status["total"],
"total": len(instances),
"counts": {
"completed": status["completed"],
"failed": status["failed"],
"running": status["running"],
"provisioning": status["provisioning"],
"completed": tally["completed"],
"failed": tally["failed"],
"running": tally["running"],
"provisioning": tally["provisioning"],
"waiting": 0,
},
"step_counts": status["step_counts"],
"instances": status["instances"],
"step_counts": step_counts,
"instances": instances,
})
return PaginatedResponse(items=items, total=len(ordered), limit=limit, offset=offset,
has_more=offset + len(items) < len(ordered))
Expand Down
38 changes: 38 additions & 0 deletions src/agent_env/store/document_store/document_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -245,6 +245,13 @@ def query(
def count(self, collection: str, filter: Filter) -> int:
"""Return the number of matching documents."""

def find_many_by_id(self, collection: str, id_field: str, ids: list[str]) -> list[dict]:
"""First document for each distinct top-level string identity, in requested order; missing IDs are omitted."""
return [
doc for identity in dict.fromkeys(ids)
if (doc := self.find_one(collection, Filter.of(**{id_field: identity}))) is not None
]

def latest_per_id(
self,
collection: str,
Expand Down Expand Up @@ -273,6 +280,37 @@ def latest_per_id(
docs = docs[:limit]
return docs

def latest_per_id_page(
self,
collection: str,
filter: Filter,
*,
id_field: str = "id",
version_field: str = "version",
sort: Optional[Sort] = None,
limit: Optional[int] = None,
offset: int = 0,
) -> tuple[list[dict], int]:
"""Return a latest-per-id page and its total, sharing generic query work.

Backends with custom pagination or counts retain those implementations
unless they override this combined operation.
"""
if (
type(self).latest_per_id is not DocumentStore.latest_per_id
or type(self).count_distinct is not DocumentStore.count_distinct
):
page = self.latest_per_id(
collection, filter, id_field=id_field, version_field=version_field,
sort=sort, limit=limit, offset=offset,
)
return page, self.count_distinct(collection, filter, id_field=id_field)
docs = _reduce_to_latest(self.query(collection, filter), id_field, version_field)
docs = _apply_sort(docs, sort)
total = len(docs)
end = offset + limit if limit else None
return docs[offset:end], total

def count_distinct(self, collection: str, filter: Filter, *, id_field: str = "id") -> int:
"""Number of distinct ``id_field`` values among matching documents.

Expand Down
25 changes: 25 additions & 0 deletions src/agent_env/store/document_store/dynamodb_document_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@
_TABLE_POLL_SECONDS = 1
_TABLE_POLL_ATTEMPTS = 120
_UNINDEXED_SORT_KEY = "sk"
_BATCH_GET_MAX_KEYS = 100


def _json_default(o):
Expand Down Expand Up @@ -209,6 +210,30 @@ def find_one(self, collection: str, filter: Filter, sort: Optional[Sort] = None)
docs = evaluation.sort_docs(self._matching(collection, filter), sort)
return docs[0] if docs else None

def find_many_by_id(self, collection: str, id_field: str, ids: list[str]) -> list[dict]:
identities = list(dict.fromkeys(ids))
if not identities or not self._table_exists(collection):
return []
if self._key_fields(collection) != [id_field]:
return super().find_many_by_id(collection, id_field, identities)
name = self._table_name(collection)
found = {}
for start in range(0, len(identities), _BATCH_GET_MAX_KEYS):
keys = [self._key(collection, {id_field: identity}) for identity in identities[start:start + _BATCH_GET_MAX_KEYS]]
pending = {name: {"Keys": keys, "ConsistentRead": True}}
for attempt in range(_CAS_ATTEMPTS):
result = self._client.batch_get_item(RequestItems=pending)
for item in result.get("Responses", {}).get(name, []):
doc = json.loads(item["doc"]["S"])
found.setdefault(doc[id_field], doc)
pending = result.get("UnprocessedKeys", {})
if not pending:
break
time.sleep(random.uniform(0, min(_CAS_MAX_BACKOFF_SECONDS, _CAS_BACKOFF_SECONDS * 2**attempt)))
else:
raise TimeoutError(f"{collection!r}: DynamoDB batch read left unprocessed keys after {_CAS_ATTEMPTS} attempts")
return [found[identity] for identity in identities if identity in found]

def query(
self,
collection: str,
Expand Down
6 changes: 6 additions & 0 deletions src/agent_env/store/document_store/mongo_document_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -168,6 +168,12 @@ def find_one(
)
return _strip_id(doc)

def find_many_by_id(self, collection: str, id_field: str, ids: list[str]) -> list[dict]:
found = {}
for doc in self.query(collection, Filter().where(id_field, In(list(dict.fromkeys(ids))))):
found.setdefault(doc[id_field], doc)
return [found[identity] for identity in dict.fromkeys(ids) if identity in found]

def query(
self,
collection: str,
Expand Down
23 changes: 23 additions & 0 deletions src/agent_env/store/document_store/sqlite_document_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@

_BUSY_TIMEOUT_SECONDS = 5.0
_WAL_SWITCH_RETRY_DELAY_SECONDS = 0.01
_ID_LOOKUP_BATCH_SIZE = 500


def _enable_wal(conn: sqlite3.Connection) -> None:
Expand Down Expand Up @@ -156,6 +157,28 @@ def query(
docs = docs[:limit]
return docs

def find_many_by_id(self, collection: str, id_field: str, ids: list[str]) -> list[dict]:
identities = list(dict.fromkeys(ids))
if not identities:
return []
with self._lock:
tbl = self._table(collection)
if tbl not in self._tables and not self._adopt_if_created(tbl):
return []
field = self._safe_path(id_field)
found = {}
for start in range(0, len(identities), _ID_LOOKUP_BATCH_SIZE):
batch = identities[start:start + _ID_LOOKUP_BATCH_SIZE]
placeholders = ", ".join(["?"] * len(batch))
rows = self._conn.execute(
# nosemgrep: sqlalchemy-execute-raw-query -- identifiers are validated; values are bound
f'SELECT rowid, doc FROM "{tbl}" WHERE {_field_expr(field)} IN ({placeholders}) ORDER BY rowid', batch,
)
for _, blob in rows:
doc = json.loads(blob)
found.setdefault(doc[id_field], doc)
return [found[identity] for identity in identities if identity in found]

def count(self, collection: str, filter: Filter) -> int:
with self._lock:
return sum(1 for _, doc in self._load(collection, filter) if evaluation.matches(doc, filter))
Expand Down
Loading
Loading