From dc962b5a20c559301d056c83d0abed157bc6678d Mon Sep 17 00:00:00 2001 From: morluto <76467478+morluto@users.noreply.github.com> Date: Wed, 7 Oct 2026 03:50:11 +0800 Subject: [PATCH 1/4] perf(explorer): batch run group instance reads --- src/agent_env/explorer/routers/runs.py | 114 +++++++++++++-------- tst/unit/explorer/explorer_app_test.py | 132 ++++++++++++++++++++++++- 2 files changed, 200 insertions(+), 46 deletions(-) diff --git a/src/agent_env/explorer/routers/runs.py b/src/agent_env/explorer/routers/runs.py index 6b4b7657..4407ed4d 100644 --- a/src/agent_env/explorer/routers/runs.py +++ b/src/agent_env/explorer/routers/runs.py @@ -20,7 +20,7 @@ from agent_env.explorer.entity_ids import EntityId from agent_env.explorer.routers.common import PaginatedResponse, docs from agent_env.runner.runner import RunStatus -from agent_env.store import Filter, Sort +from agent_env.store import Filter, In, Sort from agent_env.task.store import TaskStepStatus logger = logging.getLogger(__name__) @@ -285,32 +285,28 @@ def _all_task_runs(task_id: str) -> list: offset += 500 -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.""" +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), 500): + batch = instance_ids[start:start + 500] + for instance in store.query( + TASK_INSTANCES_COLLECTION, Filter(conditions={"instance_id": [In(batch)]}) + ): + instance_id = instance.get("instance_id") + if instance_id is not None: + instances[instance_id] = instance + return instances + + +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") @@ -335,6 +331,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). @@ -349,23 +365,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, @@ -375,10 +396,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, @@ -391,9 +415,13 @@ 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({ @@ -401,16 +429,16 @@ def list_run_groups( "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)) diff --git a/tst/unit/explorer/explorer_app_test.py b/tst/unit/explorer/explorer_app_test.py index 8437518b..9de0167a 100644 --- a/tst/unit/explorer/explorer_app_test.py +++ b/tst/unit/explorer/explorer_app_test.py @@ -1,5 +1,6 @@ """The explorer over a real SQLite store: list/get/versions semantics and a real run.""" +import asyncio import dataclasses import gzip import json @@ -19,6 +20,7 @@ from agent_env.config.errors import ConfigError from agent_env.runner import store as run_store from agent_env.runner.local_runner import LocalRunner +from agent_env.runner.runner import RunRecord, RunStatus from agent_env.store.object_store.local.store import LocalFilesystemObjectStore from agent_env.store.document_store.sqlite_document_store import LocalSqliteDocumentStore from agent_env.task import Task @@ -353,8 +355,6 @@ def test_cors_default_is_loopback_only(client): def test_run_groups_are_not_truncated_at_500(client): - from agent_env.runner.runner import RunRecord, RunStatus - gid = "rg-big" for i in range(600): run_store.insert_run(RunRecord( @@ -366,9 +366,135 @@ def test_run_groups_are_not_truncated_at_500(client): assert group["total"] == 600 # paged past the 500 fetch limit, not truncated +def test_run_group_list_uses_one_run_snapshot_and_only_page_instances(client, monkeypatch): + store = get_config().get_document_store() + query_calls = [] + original_query = store.query + run_pages = [] + original_list_runs = run_store.list_runs + + def counted_query(collection, *args, **kwargs): + if collection == "task_instances": + query_calls.append(len(args[0].conditions["instance_id"][0].values)) + return original_query(collection, *args, **kwargs) + + def counted_list_runs(*args, **kwargs): + run_pages.append(kwargs.get("offset", 0)) + return original_list_runs(*args, **kwargs) + + monkeypatch.setattr(store, "query", counted_query) + monkeypatch.setattr(run_store, "list_runs", counted_list_runs) + for i in range(1000): + run_store.insert_run(RunRecord( + run_id=f"paged-{i}", runner="local", task_id="t1", task_version=1, + instance_id=f"paged-i{i}", status=RunStatus.COMPLETED, + created_at_utc=f"2026-01-{(i // 100) + 1:02d}T00:00:{i % 60:02d}Z", + overrides={"metadata": {"run_group_id": f"rg-page-{i // 10}"}}, + )) + store.insert("task_instances", {"instance_id": f"paged-i{i}", "completed_steps": []}) + + response = client.get("/api/v1/tasks/t1/run-groups?limit=20") + assert response.status_code == 200 + body = response.json() + assert body["total"] == 100 + assert len(body["items"]) == 20 + assert all(group["total"] == 10 for group in body["items"]) + assert query_calls == [200] + assert run_pages == [0, 500, 1000] + + +def test_run_group_list_batches_a_selected_group_larger_than_500(client, monkeypatch): + gid = "rg-list-big" + store = get_config().get_document_store() + query_calls = [] + original_query = store.query + + def counted_query(collection, *args, **kwargs): + if collection == "task_instances": + query_calls.append(len(args[0].conditions["instance_id"][0].values)) + return original_query(collection, *args, **kwargs) + + monkeypatch.setattr(store, "query", counted_query) + for i in range(600): + run_store.insert_run(RunRecord( + run_id=f"list-big-{i}", runner="local", task_id="t1", task_version=1, + instance_id=f"list-big-i{i}", status=RunStatus.COMPLETED, + created_at_utc=f"2026-01-01T00:{i // 60:02d}:{i % 60:02d}Z", + overrides={"metadata": {"run_group_id": gid}}, + )) + store.insert("task_instances", {"instance_id": f"list-big-i{i}", "completed_steps": []}) + + group = client.get("/api/v1/tasks/t1/run-groups?limit=1").json()["items"][0] + assert group["run_group_id"] == gid + assert group["total"] == 600 + assert len(group["instances"]) == 600 + assert query_calls == [500, 100] + + +def test_run_group_list_keeps_missing_instances_ties_and_cross_version_membership(client): + first = "rg-tied-cross-version" + for run_id, version, instance_id in (("cross-v1", 1, "cross-i1"), ("cross-v2", 2, "cross-i2")): + run_store.insert_run(RunRecord( + run_id=run_id, runner="local", task_id="t1", task_version=version, + instance_id=instance_id, status=RunStatus.COMPLETED, + created_at_utc="2026-01-04T00:00:00Z", + overrides={"metadata": {"run_group_id": first}}, + )) + get_config().get_document_store().insert("task_instances", { + "instance_id": "cross-i1", "current_step": 2, "total_steps": 3, + }) + run_store.insert_run(RunRecord( + run_id="cross-tied-other", runner="local", task_id="t1", task_version=2, + instance_id="missing-instance", status=RunStatus.COMPLETED, + created_at_utc="2026-01-04T00:00:00Z", + overrides={"metadata": {"run_group_id": "rg-tied-other"}}, + )) + + groups = client.get("/api/v1/tasks/t1/run-groups?task_version=2").json()["items"] + assert [group["run_group_id"] for group in groups] == [first, "rg-tied-other"] + selected = groups[0] + assert selected["task_version"] == 2 + assert selected["total"] == 2 + assert {instance["instance_id"] for instance in selected["instances"]} == {"cross-i1", "cross-i2"} + missing = next(instance for instance in selected["instances"] if instance["instance_id"] == "cross-i2") + assert missing["current_step"] is None and missing["total_steps"] is None + other = groups[1] + assert other["total"] == 1 + assert other["instances"][0]["current_step"] is None + + +def test_run_group_stream_reloads_status_for_each_snapshot(client, monkeypatch): + from agent_env.explorer.routers import runs as runs_router + + run_id = "stream-fresh" + run_store.insert_run(RunRecord( + run_id=run_id, runner="local", task_id="t1", task_version=1, + instance_id="stream-instance", status=RunStatus.RUNNING, + overrides={"metadata": {"run_group_id": "rg-stream"}}, + )) + poll_count = 0 + + async def finish_between_polls(_seconds): + nonlocal poll_count + poll_count += 1 + run_store.mark_terminal(run_id, RunStatus.COMPLETED) + + monkeypatch.setattr(runs_router.asyncio, "sleep", finish_between_polls) + + async def read_events(): + response = await runs_router.stream_run_group("t1", "rg-stream") + return "".join([chunk async for chunk in response.body_iterator]) + + events = asyncio.run(read_events()) + assert events.count("event: snapshot") == 2 + snapshots = [json.loads(line.removeprefix("data: ")) for line in events.splitlines() if line.startswith("data: ")] + assert snapshots[0]["running"] == 1 + assert snapshots[1]["completed"] == 1 + assert poll_count == 1 + + def test_run_groups_tally_step_progress_for_the_batch_funnel(client): """The Batch progress strip reads `step_counts`; without it every step reads 0/N.""" - from agent_env.runner.runner import RunRecord, RunStatus gid = "rg-funnel" outcomes = { From 39244205207d15a94db81c3708f2a8cb4e3823a5 Mon Sep 17 00:00:00 2001 From: morluto <76467478+morluto@users.noreply.github.com> Date: Wed, 7 Oct 2026 03:50:11 +0800 Subject: [PATCH 2/4] perf(explorer): read entity pages and totals together --- src/agent_env/explorer/routers/common.py | 9 +- .../store/document_store/document_store.py | 31 ++++ src/agent_env/store/routing.py | 38 +++++ tst/benchmarks/explorer_reads.py | 150 ++++++++++++++++++ tst/store/conformance.py | 24 +++ tst/unit/store/routing_test.py | 8 + tst/unit/store/sqlite_document_store_test.py | 63 ++++++++ 7 files changed, 316 insertions(+), 7 deletions(-) create mode 100644 tst/benchmarks/explorer_reads.py diff --git a/src/agent_env/explorer/routers/common.py b/src/agent_env/explorer/routers/common.py index 17922116..dcd144a6 100644 --- a/src/agent_env/explorer/routers/common.py +++ b/src/agent_env/explorer/routers/common.py @@ -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 @@ -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, diff --git a/src/agent_env/store/document_store/document_store.py b/src/agent_env/store/document_store/document_store.py index a45159c6..0dd3f870 100644 --- a/src/agent_env/store/document_store/document_store.py +++ b/src/agent_env/store/document_store/document_store.py @@ -273,6 +273,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. diff --git a/src/agent_env/store/routing.py b/src/agent_env/store/routing.py index 282c5622..92b8491b 100644 --- a/src/agent_env/store/routing.py +++ b/src/agent_env/store/routing.py @@ -296,6 +296,44 @@ def latest(store: DocumentStore, **page: Any) -> list[dict]: merged = self._latest_across(readers, collection, filter, id_field, version_field) return _window(_merge_sort(merged, sort, absent_last=True), limit, offset) + 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]: + """Combine latest-per-id pages and totals across the selected sources.""" + readers = self._readers(collection, filter) + if len(readers) == 1: + return self._read( + readers[0], lambda: readers[0].latest_per_id_page( + collection, filter, id_field=id_field, version_field=version_field, + sort=sort, limit=limit, offset=offset, + ) + ) + if _ids_split_by_namespace(collection, id_field): + window = (offset or 0) + limit if limit else None + pages_and_totals = [ + self._read(store, lambda: store.latest_per_id_page( + collection, filter, id_field=id_field, version_field=version_field, + sort=sort, limit=window, offset=0, + )) + for store in readers + ] + docs = [doc for page, _ in pages_and_totals for doc in page] + docs = _merge_sort(docs, sort, absent_last=True) + total = sum(total for _, total in pages_and_totals) + else: + docs = self._latest_across(readers, collection, filter, id_field, version_field) + docs = _merge_sort(docs, sort, absent_last=True) + total = len(docs) + return _window(docs, limit, offset), total + def count_distinct(self, collection: str, filter: Filter, *, id_field: str = "id") -> int: readers = self._readers(collection, filter) if len(readers) == 1 or _ids_split_by_namespace(collection, id_field): diff --git a/tst/benchmarks/explorer_reads.py b/tst/benchmarks/explorer_reads.py new file mode 100644 index 00000000..3014d876 --- /dev/null +++ b/tst/benchmarks/explorer_reads.py @@ -0,0 +1,150 @@ +"""Compare Explorer run-group and entity-page reads with a Git baseline. + +Run from a checkout with the development dependencies installed: + + PYTHONPATH=src:packages/agentenv-protocol/src python tst/benchmarks/explorer_reads.py --base origin/main + +The script makes a temporary detached worktree for ``--base``, seeds both versions +with the same 1,000 runs in 100 groups and 2,000 entities with five versions each +in temporary SQLite stores. It reports the median of three endpoint calls plus +document-store query counts. It measures Python +and SQLite endpoint work only; it makes no HTTP or network latency claim. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import statistics +import subprocess +import sys +import tempfile +import time +from pathlib import Path + +from agent_env.config import configure +from agent_env.explorer.routers.common import versioned_router +from agent_env.explorer.routers.runs import list_run_groups +from agent_env.runner import store as run_store +from agent_env.runner.runner import RunRecord, RunStatus +from agent_env.store import LocalSqliteDocumentStore + + +def _worker() -> None: + with tempfile.TemporaryDirectory(prefix="agentenv-explorer-bench-") as temp: + root = Path(temp) + config_file = root / "config.toml" + config_file.write_text("") + os.environ["AGENT_ENV_CONFIG"] = str(config_file) + store = LocalSqliteDocumentStore(str(root / "documents.db")) + configure(document_store=store) + run_store.ensure_indexes() + for i in range(1000): + group = i // 10 + run_store.insert_run(RunRecord( + run_id=f"bench-run-{i}", runner="local", task_id="bench-task", task_version=1, + instance_id=f"bench-instance-{i}", status=RunStatus.COMPLETED, + created_at_utc=f"2026-01-{group // 30 + 1:02d}T{group % 24:02d}:{i % 60:02d}:00Z", + overrides={"metadata": {"run_group_id": f"bench-group-{group}"}}, + )) + store.insert("task_instances", { + "instance_id": f"bench-instance-{i}", "completed_steps": [], + }) + for entity in range(2000): + for version in range(1, 6): + store.insert("bench_entities", { + "id": f"entity-{entity:04d}", "version": version, + "type": "benchmark", "created_at_utc": f"2026-01-{version:02d}T00:00:00Z", + }) + + entity_list = versioned_router( + prefix="/bench", tag="bench", collection="bench_entities", noun="entity", + ).routes[0].endpoint + + counts = {"query": 0, "find_one": 0} + original_query, original_find_one = store.query, store.find_one + + def counted_query(*args, **kwargs): + counts["query"] += 1 + return original_query(*args, **kwargs) + + def counted_find_one(*args, **kwargs): + counts["find_one"] += 1 + return original_find_one(*args, **kwargs) + + store.query, store.find_one = counted_query, counted_find_one + + def measure(endpoint, *args): + elapsed, query_counts, find_one_counts, digests = [], [], [], [] + for _ in range(3): + counts.update(query=0, find_one=0) + started = time.perf_counter() + result = endpoint(*args) + elapsed.append(time.perf_counter() - started) + payload = result.model_dump(mode="json") + digests.append(hashlib.sha256(json.dumps(payload, sort_keys=True).encode()).hexdigest()) + query_counts.append(counts["query"]) + find_one_counts.append(counts["find_one"]) + assert len(set(digests)) == 1 + return { + "median_seconds": statistics.median(elapsed), + "query_calls_per_call": statistics.median(query_counts), + "find_one_calls_per_call": statistics.median(find_one_counts), + "response_sha256": digests[0], + "total": result.total, + "page_items": len(result.items), + } + + run_metrics = measure(list_run_groups, "bench-task", None, 20, 0) + entity_metrics = measure(entity_list, 50, 0, None, None, "created_at_utc", True) + assert (run_metrics["total"], run_metrics["page_items"]) == (100, 20) + assert (entity_metrics["total"], entity_metrics["page_items"]) == (2000, 50) + print(json.dumps({ + "run_groups": {**run_metrics, "runs": 1000, "groups": 100, "page_groups": 20}, + "entity_list": {**entity_metrics, "entities": 2000, "versions_each": 5, "page_size": 50}, + })) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--base", default="origin/main", help="Git ref to benchmark as the baseline") + parser.add_argument("--worker", action="store_true", help=argparse.SUPPRESS) + args = parser.parse_args() + if args.worker: + _worker() + return + + repo = Path(subprocess.check_output(["git", "rev-parse", "--show-toplevel"], text=True).strip()) + script = Path(__file__).resolve() + base_sha = subprocess.check_output(["git", "rev-parse", args.base], cwd=repo, text=True).strip() + + def run(checkout: Path) -> dict: + env = os.environ.copy() + env["PYTHONPATH"] = os.pathsep.join([ + str(checkout / "src"), str(checkout / "packages/agentenv-protocol/src"), + ]) + output = subprocess.check_output( + [sys.executable, str(script), "--worker"], cwd=checkout, env=env, text=True, + ) + return json.loads(output) + + current = run(repo) + with tempfile.TemporaryDirectory(prefix="agentenv-explorer-base-") as temp: + baseline = Path(temp) / "checkout" + subprocess.run(["git", "worktree", "add", "--detach", str(baseline), base_sha], + cwd=repo, check=True, stdout=subprocess.DEVNULL) + try: + before = run(baseline) + finally: + subprocess.run(["git", "worktree", "remove", "--force", str(baseline)], + cwd=repo, check=True, stdout=subprocess.DEVNULL) + for workload in ("run_groups", "entity_list"): + if before[workload]["response_sha256"] != current[workload]["response_sha256"]: + raise SystemExit(f"{workload} response differs between {args.base} and current") + print(json.dumps({"base": args.base, "base_sha": base_sha, "baseline": before, "current": current}, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/tst/store/conformance.py b/tst/store/conformance.py index d14bab1c..780f2a47 100644 --- a/tst/store/conformance.py +++ b/tst/store/conformance.py @@ -313,6 +313,29 @@ def latest_per_id_sort_offset_limit_apply_after_grouping(store, coll): assert store.count_distinct(coll, Filter()) == 3 +def latest_per_id_page_returns_entity_total_and_window(store, coll): + for doc in [ + {"eid": "a", "rev": 1, "rank": 99}, + {"eid": "a", "rev": 2, "rank": 30}, + {"eid": "b", "rev": 1, "rank": 20}, + {"eid": "c", "rev": 1, "rank": 10}, + {"rank": 100}, + ]: + store.insert(coll, doc) + + for limit, expected in [(1, ["b"]), (0, ["b", "c"]), (None, ["b", "c"])]: + page, total = store.latest_per_id_page( + coll, Filter(), id_field="eid", version_field="rev", + sort=Sort.by("rank", descending=True), offset=1, limit=limit, + ) + assert [doc["eid"] for doc in page] == expected + assert total == 3 + + page, total = store.latest_per_id_page(coll, Filter(), id_field="eid", offset=10, limit=1) + assert page == [] + assert total == 3 + + def latest_per_id_orders_absent_last_in_both_directions(store, coll): """A missing sort field lands last for ascending and descending alike; otherwise offset/limit page the wrong entities.""" @@ -450,6 +473,7 @@ def update_one_and_get_upsert_inserts_on_a_miss(store, coll): versioned_entity_store_roundtrip, latest_per_id_reduces_to_newest_version, latest_per_id_sort_offset_limit_apply_after_grouping, + latest_per_id_page_returns_entity_total_and_window, latest_per_id_orders_absent_last_in_both_directions, latest_per_id_skips_docs_without_identity, latest_per_id_missing_version_sorts_lowest, diff --git a/tst/unit/store/routing_test.py b/tst/unit/store/routing_test.py index 7ce27a42..0078297e 100644 --- a/tst/unit/store/routing_test.py +++ b/tst/unit/store/routing_test.py @@ -253,6 +253,9 @@ def test_default_local_users_keep_the_local_namespace_in_its_own_file(cli_routin router.insert("envs", {"id": "rocket", "version": 1}) assert _ids(router.local, "envs") == {LOCAL_ENV} assert _ids(router.configured, "envs") == {"rocket"} + page, total = router.latest_per_id_page("envs", Filter(), limit=0) + assert {doc["id"] for doc in page} == {LOCAL_ENV, "rocket"} + assert total == 2 with run_scope(LOCAL_TASK): assert isinstance(config.get_object_store(), LocalFilesystemObjectStore) with pytest.raises(LocalRunWriteError): @@ -465,6 +468,11 @@ def test_an_id_recorded_in_both_stores_reduces_to_one_latest_row(stores): ] assert [d["id"] for d in router.latest_per_id("env_snapshots", Filter(), sort=by_time, limit=1, offset=1)] == ["local-only"] assert router.count_distinct("env_snapshots", Filter()) == 3 + page, total = router.latest_per_id_page( + "env_snapshots", Filter(), sort=by_time, limit=1, offset=1, + ) + assert [doc["id"] for doc in page] == ["local-only"] + assert total == 3 with run_scope(LOCAL_TASK): assert {d["id"]: d.get("from") for d in router.latest_per_id("env_snapshots", Filter())}["tied"] == "local" diff --git a/tst/unit/store/sqlite_document_store_test.py b/tst/unit/store/sqlite_document_store_test.py index eb4415cc..cbddaa71 100644 --- a/tst/unit/store/sqlite_document_store_test.py +++ b/tst/unit/store/sqlite_document_store_test.py @@ -64,6 +64,69 @@ def add(k): assert doc["rev"] == n +def test_latest_per_id_page_reads_rows_once(store_coll, monkeypatch): + store, coll = store_coll + for entity_id, version in (("a", 1), ("a", 2), ("b", 1), ("c", 1)): + store.insert(coll, {"id": entity_id, "version": version, "created_at_utc": f"{entity_id}{version}"}) + + calls = 0 + original_query = store.query + + def counted_query(*args, **kwargs): + nonlocal calls + calls += 1 + return original_query(*args, **kwargs) + + monkeypatch.setattr(store, "query", counted_query) + page, total = store.latest_per_id_page( + coll, Filter(), sort=Sort.by("created_at_utc", descending=False), limit=1, offset=1, + ) + assert [doc["id"] for doc in page] == ["b"] + assert total == 3 + assert calls == 1 + + +def test_latest_per_id_page_preserves_backend_override(store_coll, monkeypatch): + store, coll = store_coll + + class NativePaginationStore(LocalSqliteDocumentStore): + def latest_per_id(self, collection, filter, **kwargs): + self.native_page = kwargs + return [{"id": "native"}] + + native = NativePaginationStore(str(store.path)) + monkeypatch.setattr(native, "count_distinct", lambda *args, **kwargs: 17) + page, total = native.latest_per_id_page( + coll, Filter(), sort=Sort.by("id"), limit=4, offset=8, + ) + assert page == [{"id": "native"}] + assert total == 17 + assert native.native_page["limit"] == 4 + assert native.native_page["offset"] == 8 + + class NativeCountStore(LocalSqliteDocumentStore): + def count_distinct(self, collection, filter, *, id_field="id"): + self.native_count = (collection, id_field) + return 23 + + native_count = NativeCountStore(str(store.path)) + native_count.insert(coll, {"id": "counted", "version": 1}) + page, total = native_count.latest_per_id_page(coll, Filter(), limit=1) + assert len(page) == 1 + assert total == 23 + assert native_count.native_count == (coll, "id") + + +def test_latest_per_id_page_zero_limit_keeps_unlimited_semantics(store_coll): + store, coll = store_coll + for entity_id in ("a", "b", "c"): + store.insert(coll, {"id": entity_id, "version": 1}) + + page, total = store.latest_per_id_page(coll, Filter(), limit=0, offset=1) + assert len(page) == 2 + assert total == 3 + + def test_reader_sees_a_collection_created_by_another_connection(tmp_path): """A long-lived reader must see a collection another connection creates after it opened, without reopening. Covers several primitives, not one.""" From c9d8194092cedae726d5b1ab13dc4c46885bd912 Mon Sep 17 00:00:00 2001 From: morluto <76467478+morluto@users.noreply.github.com> Date: Wed, 7 Oct 2026 08:57:43 +0800 Subject: [PATCH 3/4] fix(explorer): use indexed instance batch reads --- src/agent_env/explorer/routers/runs.py | 14 ++--- .../store/document_store/document_store.py | 7 +++ .../document_store/dynamodb_document_store.py | 25 ++++++++ .../document_store/mongo_document_store.py | 6 ++ .../document_store/sqlite_document_store.py | 23 ++++++++ src/agent_env/store/routing.py | 22 +++++++ tst/benchmarks/explorer_reads.py | 17 +++++- tst/store/conformance.py | 13 +++++ tst/unit/explorer/explorer_app_test.py | 37 ++++++++---- .../store/dynamodb_document_store_test.py | 57 +++++++++++++++++++ tst/unit/store/routing_test.py | 13 +++++ tst/unit/store/sqlite_document_store_test.py | 22 +++++++ 12 files changed, 236 insertions(+), 20 deletions(-) diff --git a/src/agent_env/explorer/routers/runs.py b/src/agent_env/explorer/routers/runs.py index 4407ed4d..2425deeb 100644 --- a/src/agent_env/explorer/routers/runs.py +++ b/src/agent_env/explorer/routers/runs.py @@ -20,11 +20,13 @@ from agent_env.explorer.entity_ids import EntityId from agent_env.explorer.routers.common import PaginatedResponse, docs from agent_env.runner.runner import RunStatus -from agent_env.store import Filter, In, Sort +from agent_env.store import Filter, Sort from agent_env.task.store import TaskStepStatus logger = logging.getLogger(__name__) +INSTANCE_READ_BATCH_SIZE = 500 + router = APIRouter(prefix="/api/v1/tasks", tags=["runs"]) TASK_INSTANCES_COLLECTION = "task_instances" @@ -290,14 +292,12 @@ def _instances_for_runs(records: list) -> dict[str, dict]: 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), 500): - batch = instance_ids[start:start + 500] - for instance in store.query( - TASK_INSTANCES_COLLECTION, Filter(conditions={"instance_id": [In(batch)]}) - ): + 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[instance_id] = instance + instances.setdefault(instance_id, instance) return instances diff --git a/src/agent_env/store/document_store/document_store.py b/src/agent_env/store/document_store/document_store.py index 0dd3f870..21d065a6 100644 --- a/src/agent_env/store/document_store/document_store.py +++ b/src/agent_env/store/document_store/document_store.py @@ -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, diff --git a/src/agent_env/store/document_store/dynamodb_document_store.py b/src/agent_env/store/document_store/dynamodb_document_store.py index 412623eb..be921e60 100644 --- a/src/agent_env/store/document_store/dynamodb_document_store.py +++ b/src/agent_env/store/document_store/dynamodb_document_store.py @@ -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): @@ -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, diff --git a/src/agent_env/store/document_store/mongo_document_store.py b/src/agent_env/store/document_store/mongo_document_store.py index a20c6ce6..9c650f70 100644 --- a/src/agent_env/store/document_store/mongo_document_store.py +++ b/src/agent_env/store/document_store/mongo_document_store.py @@ -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, diff --git a/src/agent_env/store/document_store/sqlite_document_store.py b/src/agent_env/store/document_store/sqlite_document_store.py index 0f7904ce..73aebeb2 100644 --- a/src/agent_env/store/document_store/sqlite_document_store.py +++ b/src/agent_env/store/document_store/sqlite_document_store.py @@ -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: @@ -133,6 +134,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)) diff --git a/src/agent_env/store/routing.py b/src/agent_env/store/routing.py index 92b8491b..0622dbdb 100644 --- a/src/agent_env/store/routing.py +++ b/src/agent_env/store/routing.py @@ -254,6 +254,28 @@ def find_one(self, collection: str, filter: Filter, sort: Optional[Sort] = None) ] return _merge_sort(found, sort)[0] if found else None + def find_many_by_id(self, collection: str, id_field: str, ids: list[str]) -> list[dict]: + identities = list(dict.fromkeys(ids)) + readers_by_id = { + identity: self._readers(collection, Filter.of(**{id_field: identity})) + for identity in identities + } + batches: dict[int, tuple[DocumentStore, list[str]]] = {} + for identity, readers in readers_by_id.items(): + for store in readers: + _, batch = batches.setdefault(id(store), (store, [])) + batch.append(identity) + found = {} + for store, batch in batches.values(): + found[id(store)] = { + doc[id_field]: doc + for doc in self._read(store, lambda: store.find_many_by_id(collection, id_field, batch)) + } + return [ + doc for identity, readers in readers_by_id.items() + if (doc := next((found[id(store)][identity] for store in readers if identity in found[id(store)]), None)) is not None + ] + def query( self, collection: str, diff --git a/tst/benchmarks/explorer_reads.py b/tst/benchmarks/explorer_reads.py index 3014d876..e3564d76 100644 --- a/tst/benchmarks/explorer_reads.py +++ b/tst/benchmarks/explorer_reads.py @@ -41,6 +41,7 @@ def _worker() -> None: store = LocalSqliteDocumentStore(str(root / "documents.db")) configure(document_store=store) run_store.ensure_indexes() + store.ensure_index("task_instances", ["instance_id"], unique=True) for i in range(1000): group = i // 10 run_store.insert_run(RunRecord( @@ -63,7 +64,7 @@ def _worker() -> None: prefix="/bench", tag="bench", collection="bench_entities", noun="entity", ).routes[0].endpoint - counts = {"query": 0, "find_one": 0} + counts = {"query": 0, "find_one": 0, "id_lookup": 0} original_query, original_find_one = store.query, store.find_one def counted_query(*args, **kwargs): @@ -75,11 +76,19 @@ def counted_find_one(*args, **kwargs): return original_find_one(*args, **kwargs) store.query, store.find_one = counted_query, counted_find_one + if hasattr(store, "find_many_by_id"): + original_lookup = store.find_many_by_id + + def counted_lookup(*args, **kwargs): + counts["id_lookup"] += 1 + return original_lookup(*args, **kwargs) + + store.find_many_by_id = counted_lookup def measure(endpoint, *args): - elapsed, query_counts, find_one_counts, digests = [], [], [], [] + elapsed, query_counts, find_one_counts, lookup_counts, digests = [], [], [], [], [] for _ in range(3): - counts.update(query=0, find_one=0) + counts.update(query=0, find_one=0, id_lookup=0) started = time.perf_counter() result = endpoint(*args) elapsed.append(time.perf_counter() - started) @@ -87,11 +96,13 @@ def measure(endpoint, *args): digests.append(hashlib.sha256(json.dumps(payload, sort_keys=True).encode()).hexdigest()) query_counts.append(counts["query"]) find_one_counts.append(counts["find_one"]) + lookup_counts.append(counts["id_lookup"]) assert len(set(digests)) == 1 return { "median_seconds": statistics.median(elapsed), "query_calls_per_call": statistics.median(query_counts), "find_one_calls_per_call": statistics.median(find_one_counts), + "id_lookup_calls_per_call": statistics.median(lookup_counts), "response_sha256": digests[0], "total": result.total, "page_items": len(result.items), diff --git a/tst/store/conformance.py b/tst/store/conformance.py index 780f2a47..46761a66 100644 --- a/tst/store/conformance.py +++ b/tst/store/conformance.py @@ -313,6 +313,18 @@ def latest_per_id_sort_offset_limit_apply_after_grouping(store, coll): assert store.count_distinct(coll, Filter()) == 3 +def find_many_by_id_preserves_first_documents_and_request_order(store, coll): + store.insert(coll, {"identity": "a", "value": "first"}) + store.insert(coll, {"identity": "b", "value": "second"}) + store.insert(coll, {"identity": "a", "value": "later"}) + + assert store.find_many_by_id(coll, "identity", ["b", "missing", "a", "b"]) == [ + store.find_one(coll, Filter.of(identity="b")), + store.find_one(coll, Filter.of(identity="a")), + ] + assert store.find_many_by_id(coll, "identity", []) == [] + + def latest_per_id_page_returns_entity_total_and_window(store, coll): for doc in [ {"eid": "a", "rev": 1, "rank": 99}, @@ -474,6 +486,7 @@ def update_one_and_get_upsert_inserts_on_a_miss(store, coll): latest_per_id_reduces_to_newest_version, latest_per_id_sort_offset_limit_apply_after_grouping, latest_per_id_page_returns_entity_total_and_window, + find_many_by_id_preserves_first_documents_and_request_order, latest_per_id_orders_absent_last_in_both_directions, latest_per_id_skips_docs_without_identity, latest_per_id_missing_version_sorts_lowest, diff --git a/tst/unit/explorer/explorer_app_test.py b/tst/unit/explorer/explorer_app_test.py index 9de0167a..62972764 100644 --- a/tst/unit/explorer/explorer_app_test.py +++ b/tst/unit/explorer/explorer_app_test.py @@ -23,6 +23,7 @@ from agent_env.runner.runner import RunRecord, RunStatus from agent_env.store.object_store.local.store import LocalFilesystemObjectStore from agent_env.store.document_store.sqlite_document_store import LocalSqliteDocumentStore +from agent_env.store.routing import RoutingDocumentStore from agent_env.task import Task from agent_env.explorer.routers import objects as objects_router from agent_env.explorer.routers import triggers as triggers_router @@ -369,20 +370,20 @@ def test_run_groups_are_not_truncated_at_500(client): def test_run_group_list_uses_one_run_snapshot_and_only_page_instances(client, monkeypatch): store = get_config().get_document_store() query_calls = [] - original_query = store.query + original_query = store.find_many_by_id run_pages = [] original_list_runs = run_store.list_runs - def counted_query(collection, *args, **kwargs): + def counted_query(collection, id_field, ids): if collection == "task_instances": - query_calls.append(len(args[0].conditions["instance_id"][0].values)) - return original_query(collection, *args, **kwargs) + query_calls.append(len(ids)) + return original_query(collection, id_field, ids) def counted_list_runs(*args, **kwargs): run_pages.append(kwargs.get("offset", 0)) return original_list_runs(*args, **kwargs) - monkeypatch.setattr(store, "query", counted_query) + monkeypatch.setattr(store, "find_many_by_id", counted_query) monkeypatch.setattr(run_store, "list_runs", counted_list_runs) for i in range(1000): run_store.insert_run(RunRecord( @@ -403,18 +404,34 @@ def counted_list_runs(*args, **kwargs): assert run_pages == [0, 500, 1000] +def test_run_group_list_preserves_routed_instance_precedence(client, tmp_path): + configured = get_config().get_document_store() + local = LocalSqliteDocumentStore(str(tmp_path / "local.db")) + router = RoutingDocumentStore(configured, local, tmp_path / "local.db") + configured.insert("task_instances", {"instance_id": "shared", "current_step": 2, "total_steps": 4}) + local.insert("task_instances", {"instance_id": "shared", "current_step": 7, "total_steps": 9}) + run_store.insert_run(RunRecord( + run_id="shared-run", runner="local", task_id="t1", task_version=1, + instance_id="shared", status=RunStatus.RUNNING, created_at_utc="2026-01-01T00:00:00Z", + )) + set_document_store(router) + + group = client.get("/api/v1/tasks/t1/run-groups").json()["items"][0] + assert (group["instances"][0]["current_step"], group["instances"][0]["total_steps"]) == (2, 4) + + def test_run_group_list_batches_a_selected_group_larger_than_500(client, monkeypatch): gid = "rg-list-big" store = get_config().get_document_store() query_calls = [] - original_query = store.query + original_query = store.find_many_by_id - def counted_query(collection, *args, **kwargs): + def counted_query(collection, id_field, ids): if collection == "task_instances": - query_calls.append(len(args[0].conditions["instance_id"][0].values)) - return original_query(collection, *args, **kwargs) + query_calls.append(len(ids)) + return original_query(collection, id_field, ids) - monkeypatch.setattr(store, "query", counted_query) + monkeypatch.setattr(store, "find_many_by_id", counted_query) for i in range(600): run_store.insert_run(RunRecord( run_id=f"list-big-{i}", runner="local", task_id="t1", task_version=1, diff --git a/tst/unit/store/dynamodb_document_store_test.py b/tst/unit/store/dynamodb_document_store_test.py index a045a25a..e29f1be6 100644 --- a/tst/unit/store/dynamodb_document_store_test.py +++ b/tst/unit/store/dynamodb_document_store_test.py @@ -48,6 +48,63 @@ def read_then_lose_the_race(collection, filter): assert store.find_one(coll, Filter.of(id="x")) == {"id": "x", "rev": 1} +def test_batch_identity_lookup_uses_primary_keys_in_bounded_requests(store_coll, monkeypatch): + store, coll = store_coll + store.ensure_index(coll, ["instance_id"], unique=True) + for i in range(200): + store.insert(coll, {"instance_id": f"i{i}"}) + requests = [] + original_batch_get = store._client.batch_get_item + + def batch_get(**kwargs): + request = kwargs["RequestItems"][store._table_name(coll)] + assert request["ConsistentRead"] is True + requests.append(len(request["Keys"])) + return original_batch_get(**kwargs) + + def refuse_scan(*args, **kwargs): + pytest.fail("batch identity reads must not scan") + + monkeypatch.setattr(store._client, "scan", refuse_scan) + monkeypatch.setattr(store._client, "get_item", refuse_scan) + monkeypatch.setattr(store._client, "batch_get_item", batch_get) + ids = [f"i{i}" for i in range(120)] + assert store.find_many_by_id(coll, "instance_id", ids + ["missing", "i0"]) == [ + {"instance_id": identity} for identity in ids + ] + assert requests == [100, 21] + + +def test_batch_identity_lookup_retries_unprocessed_keys(store_coll, monkeypatch): + store, coll = store_coll + store.ensure_index(coll, ["instance_id"], unique=True) + store.insert(coll, {"instance_id": "i"}) + original_batch_get = store._client.batch_get_item + calls = 0 + + def batch_get(**kwargs): + nonlocal calls + calls += 1 + if calls == 1: + return {"Responses": {}, "UnprocessedKeys": kwargs["RequestItems"]} + return original_batch_get(**kwargs) + + monkeypatch.setattr(store._client, "batch_get_item", batch_get) + monkeypatch.setattr(dynamodb_mod.time, "sleep", lambda _: None) + assert store.find_many_by_id(coll, "instance_id", ["i"]) == [{"instance_id": "i"}] + assert calls == 2 + + +def test_batch_identity_lookup_does_not_hide_exhausted_retries(store_coll, monkeypatch): + store, coll = store_coll + store.ensure_index(coll, ["instance_id"], unique=True) + monkeypatch.setattr(dynamodb_mod, "_CAS_ATTEMPTS", 2) + monkeypatch.setattr(dynamodb_mod.time, "sleep", lambda _: None) + monkeypatch.setattr(store._client, "batch_get_item", lambda **kwargs: {"UnprocessedKeys": kwargs["RequestItems"]}) + with pytest.raises(TimeoutError, match="unprocessed keys"): + store.find_many_by_id(coll, "instance_id", ["i"]) + + def test_a_write_that_changes_a_unique_index_field_is_refused(store_coll): store, coll = store_coll store.ensure_index(coll, ["id", "version"], unique=True) diff --git a/tst/unit/store/routing_test.py b/tst/unit/store/routing_test.py index 0078297e..336fa896 100644 --- a/tst/unit/store/routing_test.py +++ b/tst/unit/store/routing_test.py @@ -506,6 +506,19 @@ def test_a_raw_entity_write_is_checked_like_a_versioned_one(stores, entity_id, r assert not (state_root() / "document_store" / "local.db").exists() +def test_batch_instance_lookup_keeps_the_first_routed_copy(stores): + router, configured, local = stores + configured.ensure_index("task_instances", ["instance_id"], unique=True) + local.ensure_index("task_instances", ["instance_id"], unique=True) + configured.insert("task_instances", {"instance_id": "shared", "current_step": 2}) + local.insert("task_instances", {"instance_id": "shared", "current_step": 7}) + + for scope in (nullcontext(), run_scope(LOCAL_TASK)): + with scope: + expected = router.find_one("task_instances", Filter.of(instance_id="shared")) + assert router.find_many_by_id("task_instances", "instance_id", ["shared"]) == [expected] + + def test_an_entity_store_defined_outside_core_is_routed_by_id(stores): router, configured, local = stores plugin_things = VersionedEntityStore(router, "plugin_things", serialize=dict, deserialize=dict) diff --git a/tst/unit/store/sqlite_document_store_test.py b/tst/unit/store/sqlite_document_store_test.py index cbddaa71..23c222a4 100644 --- a/tst/unit/store/sqlite_document_store_test.py +++ b/tst/unit/store/sqlite_document_store_test.py @@ -211,6 +211,28 @@ def test_other_filters_scan(filter, store_coll): assert _plans(store, lambda: store.query(coll, filter)) == ["SCAN docs_coll"] +def test_batch_identity_lookup_searches_the_index_and_decodes_only_requested_rows(store_coll, monkeypatch): + store, coll = store_coll + store.ensure_index(coll, ["instance_id"], unique=True) + for i in range(1000): + store.insert(coll, {"instance_id": f"i{i}"}) + + decoded = [] + original_loads = sqlite_document_store.json.loads + + def counted_loads(blob, *args, **kwargs): + decoded.append(blob) + return original_loads(blob, *args, **kwargs) + + monkeypatch.setattr(sqlite_document_store.json, "loads", counted_loads) + plans = _plans(store, lambda: store.find_many_by_id(coll, "instance_id", ["i500", "i2", "missing"])) + assert any("USING INDEX docs_coll_instance_id_unique" in plan for plan in plans) + assert len(decoded) == 2 + assert store.find_many_by_id(coll, "instance_id", ["i500", "i2", "missing"]) == [ + {"instance_id": "i500"}, {"instance_id": "i2"}, + ] + + @pytest.mark.parametrize("same_store", [False, True], ids=["another connection", "the same store"]) def test_a_reader_searches_an_index_created_after_its_first_read(tmp_path, same_store): path = str(tmp_path / "shared.db") From 044d9a104d55a8a789a3898de70dc223a0d5352f Mon Sep 17 00:00:00 2001 From: morluto <76467478+morluto@users.noreply.github.com> Date: Wed, 7 Oct 2026 09:01:37 +0800 Subject: [PATCH 4/4] refactor(explorer): name task run pagination size --- src/agent_env/explorer/routers/runs.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/src/agent_env/explorer/routers/runs.py b/src/agent_env/explorer/routers/runs.py index 2425deeb..a347e3fd 100644 --- a/src/agent_env/explorer/routers/runs.py +++ b/src/agent_env/explorer/routers/runs.py @@ -26,6 +26,7 @@ logger = logging.getLogger(__name__) INSTANCE_READ_BATCH_SIZE = 500 +TASK_RUN_PAGE_SIZE = 500 router = APIRouter(prefix="/api/v1/tasks", tags=["runs"]) @@ -280,11 +281,11 @@ 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 _instances_for_runs(records: list) -> dict[str, dict]: