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
14 changes: 8 additions & 6 deletions src/agent_env/store/document_store/document_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -286,6 +286,12 @@ def count_distinct(self, collection: str, filter: Filter, *, id_field: str = "id
keys.add(value)
return len(keys)

def latest_version(self, collection: str, entity_id: str) -> Optional[dict]:
"""Return the latest versioned entity; backends may optimize the generic sorted lookup."""
return self.find_one(
collection, Filter.of(id=entity_id), sort=Sort.by("version", descending=True)
)

@abstractmethod
def insert(self, collection: str, doc: dict) -> None:
"""Insert one document. Raises DuplicateKeyError on unique violation."""
Expand Down Expand Up @@ -389,9 +395,7 @@ def get(self, id: str, version: Optional[int] = None) -> Optional[T]:
if version is not None:
doc = self._doc_store.find_one(self._collection, Filter.of(id=id, version=version))
else:
doc = self._doc_store.find_one(
self._collection, Filter.of(id=id), sort=Sort.by("version", descending=True)
)
doc = self._doc_store.latest_version(self._collection, id)
return self._deserialize(doc) if doc is not None else None

def next_version(self, id: str) -> int:
Expand All @@ -403,9 +407,7 @@ def next_version(self, id: str) -> int:
that data orphaned with no artifact record pointing at it.
"""
self._doc_store.check_id(id)
doc = self._doc_store.find_one(
self._collection, Filter.of(id=id), sort=Sort.by("version", descending=True)
)
doc = self._doc_store.latest_version(self._collection, id)
return (doc["version"] + 1) if doc is not None else 1

def put(self, entity: T, max_retries: int = 5) -> int:
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 @@ -116,6 +116,29 @@ def find_one(
matches = evaluation.sort_docs(matches, sort)
return matches[0] if matches else None

def latest_version(self, collection: str, entity_id: str) -> Optional[dict]:
"""Read one scalar-id, integer-version entity through the compound index."""
if not isinstance(entity_id, str):
return super().latest_version(collection, entity_id)
with self._lock:
tbl = self._table(collection)
if tbl not in self._tables and not self._adopt_if_created(tbl):
return None
if not any(unique and fields == ["id", "version"] for unique, fields in self._indexes_of(tbl)):
return super().latest_version(collection, entity_id)
row = self._conn.execute(
# nosemgrep: sqlalchemy-execute-raw-query -- tbl and JSON paths are fixed/validated; id is bound
f'SELECT doc, json_type(doc, \'$.version\'), typeof({_field_expr("version")}) FROM "{tbl}" '
f'WHERE {_field_expr("id")} = json_extract(?, \'$\') '
f'ORDER BY {_field_expr("version")} DESC LIMIT 1',
(json.dumps(entity_id),),
).fetchone()
if row is None:
return None
if row[1] != "integer" or row[2] != "integer":
return super().latest_version(collection, entity_id)
return json.loads(row[0])

def query(
self,
collection: str,
Expand Down
8 changes: 8 additions & 0 deletions src/agent_env/store/routing.py
Original file line number Diff line number Diff line change
Expand Up @@ -254,6 +254,14 @@ def find_one(self, collection: str, filter: Filter, sort: Optional[Sort] = None)
]
return _merge_sort(found, sort)[0] if found else None

def latest_version(self, collection: str, entity_id: str) -> Optional[dict]:
filter = Filter.of(id=entity_id)
found = [
doc for store in self._readers(collection, filter)
if (doc := self._read(store, lambda: store.latest_version(collection, entity_id))) is not None
]
return _merge_sort(found, Sort.by("version", descending=True))[0] if found else None

def query(
self,
collection: str,
Expand Down
218 changes: 218 additions & 0 deletions tst/benchmarks/version_lookup.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,218 @@
"""Compare full VersionedEntityStore operations before and after indexed latest reads.

Run from the repository root with the shared environment:
``PYTHONPATH=src:packages/agentenv-protocol/src .venv/bin/python tst/benchmarks/version_lookup.py``
"""

from __future__ import annotations

import ast
import hashlib
import json
import platform
import statistics
import subprocess
import sys
import tempfile
import time
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Optional, TypeVar

from agent_env.store.document_store.document_store import DuplicateKeyError, Filter, Sort, VersionedEntityStore
from agent_env.store.document_store.sqlite_document_store import LocalSqliteDocumentStore

HISTORY_SIZES = (1, 100, 1000, 5000)
BASELINE_REVISION = "05b3310d5fd2c0f210b22750fd46f8c586ab6732"
UNRELATED_IDS = 1000
PAYLOAD_BYTES = 128
GET_ROUNDS = 9
PUT_ROUNDS = 7
GETS_PER_SAMPLE = 30
PUTS_PER_ROUND = 12


@dataclass(eq=True)
class Entity:
id: str
version: int
payload: str


def _serialize(entity: Entity) -> dict:
return asdict(entity)


def _deserialize(doc: dict) -> Entity:
return Entity(**doc)


class CountingDocumentStore(LocalSqliteDocumentStore):
def __init__(self, path: str) -> None:
super().__init__(path)
self.decoded_rows = 0

def _load(self, collection, filter):
rows = super()._load(collection, filter)
self.decoded_rows += len(rows)
return rows

def latest_version(self, collection: str, entity_id: str) -> dict | None:
before = self.decoded_rows
result = super().latest_version(collection, entity_id)
if self.decoded_rows == before and result is not None:
self.decoded_rows += 1
return result


def _baseline_versioned_store() -> tuple[type[VersionedEntityStore], str]:
"""Load the version-store methods directly from the pinned pre-change source."""
command = ["git", "show", f"{BASELINE_REVISION}:src/agent_env/store/document_store/document_store.py"]
try:
source = subprocess.check_output(command, text=True, stderr=subprocess.PIPE)
except subprocess.CalledProcessError:
raise SystemExit(
f"Benchmark baseline {BASELINE_REVISION} is unavailable in this checkout. "
f"Fetch it with `git fetch origin {BASELINE_REVISION}` and rerun the benchmark."
) from None
module = ast.parse(source)
original = next(
node for node in module.body
if isinstance(node, ast.ClassDef) and node.name == "VersionedEntityStore"
)
methods = [
node for node in original.body
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef))
and node.name in {"get", "next_version", "put"}
]
if {node.name for node in methods} != {"get", "next_version", "put"}:
raise RuntimeError(f"{BASELINE_REVISION} does not contain the expected version-store methods")
legacy = ast.ClassDef(
name="BaselineVersionedEntityStore",
bases=[ast.Name(id="VersionedEntityStore", ctx=ast.Load())],
keywords=[],
body=methods,
decorator_list=[],
)
namespace = {
"VersionedEntityStore": VersionedEntityStore,
"Filter": Filter,
"Sort": Sort,
"Optional": Optional,
"T": TypeVar("T"),
"DuplicateKeyError": DuplicateKeyError,
}
exec(compile(ast.fix_missing_locations(ast.Module(body=[legacy], type_ignores=[])), "baseline_versioned_store", "exec"), namespace)
return namespace["BaselineVersionedEntityStore"], hashlib.sha256(source.encode()).hexdigest()


def _seed(
path: Path,
backend: type[LocalSqliteDocumentStore],
history: int,
versioned_store: type[VersionedEntityStore],
):
docs = backend(str(path))
view = versioned_store(docs, "entities", _serialize, _deserialize)
for version in range(1, history + 1):
docs.insert("entities", {"id": "target", "version": version, "payload": "x" * PAYLOAD_BYTES})
for index in range(UNRELATED_IDS):
docs.insert("entities", {"id": f"other-{index}", "version": index + 1, "payload": "x" * PAYLOAD_BYTES})
return docs, view


def _timed(call, repetitions: int) -> float:
started = time.perf_counter_ns()
for _ in range(repetitions):
call()
return (time.perf_counter_ns() - started) / repetitions / 1000


def _plan_and_rows(store: LocalSqliteDocumentStore) -> tuple[str, int]:
tbl = store._table("entities")
plan = store._conn.execute(
f"EXPLAIN QUERY PLAN SELECT doc FROM \"{tbl}\" "
"WHERE json_extract(doc, '$.id') = json_extract(?, '$') "
"ORDER BY json_extract(doc, '$.version') DESC LIMIT 1",
(json.dumps("target"),),
).fetchall()
candidates = store._conn.execute(
f"SELECT count(*) FROM \"{tbl}\" WHERE json_extract(doc, '$.id') = json_extract(?, '$')",
(json.dumps("target"),),
).fetchone()[0]
return "; ".join(row[3] for row in plan), candidates


def main() -> None:
baseline_store, baseline_source_hash = _baseline_versioned_store()
result = {
"baseline_revision": BASELINE_REVISION,
"baseline_methods": ["get", "next_version", "put"],
"baseline_source_sha256": baseline_source_hash,
"final_worktree_head": subprocess.check_output(["git", "rev-parse", "HEAD"], text=True).strip(),
"final_worktree_dirty": bool(subprocess.check_output(["git", "status", "--porcelain"], text=True).strip()),
"python": sys.version,
"platform": platform.platform(),
"rows": [],
}
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
for history in HISTORY_SIZES:
base_docs, base = _seed(root / f"base-{history}.db", CountingDocumentStore, history, baseline_store)
fast_docs, fast = _seed(root / f"fast-{history}.db", CountingDocumentStore, history, VersionedEntityStore)
get_base, get_fast = [], []
for round_number in range(GET_ROUNDS):
order = [("base", base), ("fast", fast)]
if round_number % 2:
order.reverse()
for name, view in order:
duration = _timed(lambda: view.get("target"), GETS_PER_SAMPLE)
(get_base if name == "base" else get_fast).append(duration)
put_base, put_fast, versions_equal = [], [], True
for round_number in range(PUT_ROUNDS):
for offset in range(PUTS_PER_ROUND):
first, second = (base, fast) if (round_number + offset) % 2 == 0 else (fast, base)
start = time.perf_counter_ns()
first_version = first.put(Entity("target", 0, "p" * PAYLOAD_BYTES))
first_us = (time.perf_counter_ns() - start) / 1000
start = time.perf_counter_ns()
second_version = second.put(Entity("target", 0, "p" * PAYLOAD_BYTES))
second_us = (time.perf_counter_ns() - start) / 1000
versions_equal &= first_version == second_version
if first is base:
put_base.append(first_us)
put_fast.append(second_us)
else:
put_fast.append(first_us)
put_base.append(second_us)
latest_equal = base.get("target") == fast.get("target")
exact_equal = base.get("target", version=history) == fast.get("target", version=history)
plan, candidate_rows = _plan_and_rows(fast_docs)
base_docs.decoded_rows = fast_docs.decoded_rows = 0
base.get("target")
baseline_decode_count = base_docs.decoded_rows
fast_docs.decoded_rows = 0
fast.get("target")
final_decode_count = fast_docs.decoded_rows
result["rows"].append({
"history_seed": history,
"unrelated_ids": UNRELATED_IDS,
"get_baseline_us_median": round(statistics.median(get_base), 2),
"get_final_us_median": round(statistics.median(get_fast), 2),
"get_speedup": round(statistics.median(get_base) / statistics.median(get_fast), 1),
"put_baseline_us_median": round(statistics.median(put_base), 2),
"put_final_us_median": round(statistics.median(put_fast), 2),
"put_speedup": round(statistics.median(put_base) / statistics.median(put_fast), 1),
"put_versions_equal": versions_equal,
"latest_equal": latest_equal,
"exact_equal": exact_equal,
"indexed_candidate_rows": candidate_rows,
"latest_get_decoded_rows_baseline": baseline_decode_count,
"latest_get_decoded_rows_final": final_decode_count,
"top_one_query_plan": plan,
})
print(json.dumps(result, indent=2))


if __name__ == "__main__":
main()
11 changes: 11 additions & 0 deletions tst/unit/store/routing_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -469,6 +469,17 @@ def test_an_id_recorded_in_both_stores_reduces_to_one_latest_row(stores):
assert {d["id"]: d.get("from") for d in router.latest_per_id("env_snapshots", Filter())}["tied"] == "local"


def test_latest_version_lookup_uses_the_routed_namespace(stores):
router, _, _ = stores
versioned = VersionedEntityStore(router, "env_snapshots", dict, dict)
router.insert("env_snapshots", {"id": "shared", "version": 1})
local_id = "@local/~/bundle/env_snapshots/local-only"
with run_scope(LOCAL_TASK):
router.insert("env_snapshots", {"id": local_id, "version": 3})
assert versioned.get("shared")["version"] == 1
assert versioned.get(local_id)["version"] == 3


def test_the_local_namespace_file_cant_be_the_configured_store(tmp_path, cli_routing):
configure(document_store=LocalSqliteDocumentStore(str(state_root() / "document_store" / "local.db")))
with pytest.raises(ConfigError, match="kept for @local documents"):
Expand Down
40 changes: 39 additions & 1 deletion tst/unit/store/sqlite_document_store_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -160,6 +160,44 @@ def test_a_reader_searches_an_index_created_after_its_first_read(tmp_path, same_
assert "USING INDEX docs_coll_id_version_unique (<expr>=?)" in _plans(reader, read)[0]


def test_latest_version_uses_indexed_top_one_and_reads_existing_table(tmp_path):
store = LocalSqliteDocumentStore(str(tmp_path / "latest.db"))
store.insert("coll", {"id": "a", "version": 1})
store.insert("coll", {"id": "a", "version": 2})
store.ensure_index("coll", ["id", "version"], unique=True)
statements = []
store._conn.set_trace_callback(statements.append)
view = VersionedEntityStore(store, "coll", dict, dict)
statements.clear()
assert view.get("a")["version"] == 2
assert view.next_version("a") == 3
reads = [sql for sql in statements if sql.startswith("SELECT doc, json_type")]
assert len(reads) == 2
plan = store._conn.execute("EXPLAIN QUERY PLAN " + reads[0]).fetchall()
assert any("docs_coll_id_version_unique" in row[3] for row in plan), plan
assert all("LIMIT 1" in sql for sql in reads)


def test_latest_version_does_not_create_an_absent_table(tmp_path):
store = LocalSqliteDocumentStore(str(tmp_path / "absent.db"))
assert store.latest_version("coll", "missing") is None
assert store._conn.execute("SELECT 1 FROM sqlite_master WHERE name='docs_coll'").fetchone() is None


def test_latest_version_falls_back_for_non_integer_entity_version(tmp_path):
store = LocalSqliteDocumentStore(str(tmp_path / "malformed-version.db"))
view = VersionedEntityStore(store, "coll", dict, dict)
store.insert("coll", {"id": "a", "version": "legacy"})
assert view.get("a")["version"] == "legacy"


def test_latest_version_falls_back_for_versions_outside_sqlite_integer_range(tmp_path):
store = LocalSqliteDocumentStore(str(tmp_path / "large-version.db"))
view = VersionedEntityStore(store, "coll", dict, dict)
store.insert("coll", {"id": "a", "version": 10**30})
assert view.get("a")["version"] == 10**30


def test_indexed_reads_and_writes_match_a_full_scan(tmp_path):
"""Random documents and operations give the same results, in the same order, as the same store with no index."""
rng = random.Random(3205)
Expand Down Expand Up @@ -320,5 +358,5 @@ def _plans(store, call) -> list[str]:
call()
finally:
store._conn.set_trace_callback(None)
reads = [sql for sql in statements if sql.startswith("SELECT rowid, doc")]
reads = [sql for sql in statements if sql.startswith(("SELECT rowid, doc", "SELECT doc, json_type"))]
return [" ".join(row[3] for row in store._conn.execute(f"EXPLAIN QUERY PLAN {sql}")) for sql in reads]
20 changes: 20 additions & 0 deletions tst/unit/store/test_document_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -195,6 +195,26 @@ def test_next_version(self):
store.put({"id": "e"})
assert store.next_version("e") == 2

def test_default_latest_version_uses_custom_find_one_sort(self):
class RecordingStore(FakeDocumentStore):
def __init__(self):
super().__init__()
self.reads = []

def find_one(self, collection, filter, sort=None):
self.reads.append((collection, filter, sort))
return super().find_one(collection, filter, sort)

docs = RecordingStore()
store = _versioned(docs)
store.put({"id": "e"})
docs.reads.clear()

assert store.get("e")["version"] == 1
assert docs.reads[-1] == ("c", Filter.of(id="e"), Sort.by("version", descending=True))
assert store.next_version("e") == 2
assert docs.reads[-1] == ("c", Filter.of(id="e"), Sort.by("version", descending=True))

def test_put_retries_past_a_collision(self):
docs = FakeDocumentStore(fail_inserts=1)
store = _versioned(docs)
Expand Down
Loading
Loading