mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-09 21:49:37 +00:00
test(checkpoint): measure postgres storage growth (#5051)
This commit is contained in:
parent
cd35363a05
commit
c56c7293f8
@ -152,10 +152,18 @@ the number of required IDs, whichever is larger; missing exact runs use targeted
|
||||
**Checkpoint channel benchmark**: `scripts/benchmark/checkpoint/bench_channels.py`
|
||||
runs paired `full`/`delta` message-only StateGraphs in a fresh child process per
|
||||
case, using sync `InMemorySaver` or `SqliteSaver` so reducer, serialization, and
|
||||
saver costs stay separate from Gateway/async scheduling. It reports deterministic
|
||||
saver costs stay separate from Gateway/async scheduling. Optional
|
||||
`AsyncPostgresSaver` cases are enabled only when `TEST_POSTGRES_URI` is set.
|
||||
Postgres cases use a unique thread and remove only that benchmark thread through
|
||||
the saver's public `adelete_thread` API after measurement. It reports deterministic
|
||||
correctness digests, write windows/percentiles, warm and graph-rebuilt cold reads,
|
||||
logical checkpoint/write bytes, SQLite DB/WAL/SHM footprint, reducer replay time,
|
||||
and peak RSS as versioned JSONL. The controller alternates mode order and rejects
|
||||
backend-neutral checkpoint/blob/write row and byte fields, aggregate logical
|
||||
checkpoint/write bytes, SQLite DB/WAL/SHM footprint, reducer replay time, and
|
||||
peak RSS as versioned JSONL. SQLite embeds channel blobs in its checkpoint
|
||||
payload, so its separate blob metrics are zero; Postgres reports its
|
||||
`checkpoint_blobs` table separately. Byte fields describe each saver's serialized
|
||||
representation and should not be treated as identical encodings across backends.
|
||||
The controller alternates mode order and rejects
|
||||
performance data when paired modes materialize different state. Its default 1 GiB
|
||||
estimated cumulative full-payload cap skips both modes of an oversized pair when
|
||||
`full` is selected, including every delta cadence in a `--snapshot-frequencies`
|
||||
@ -177,6 +185,10 @@ cd backend
|
||||
PYTHONPATH=. uv run python scripts/benchmark/checkpoint/bench_channels.py \
|
||||
--backends sqlite --updates 100,500,999,1000,1001 --payload-bytes 128 \
|
||||
--repetitions 7 --output /tmp/checkpoint-bench.jsonl
|
||||
TEST_POSTGRES_URI=postgresql://... \
|
||||
PYTHONPATH=. uv run python scripts/benchmark/checkpoint/bench_channels.py \
|
||||
--backends sqlite,postgres --updates 100 --payload-bytes 128 \
|
||||
--output /tmp/checkpoint-cross-backend.jsonl
|
||||
PYTHONPATH=. uv run python scripts/benchmark/checkpoint/summarize_channels.py \
|
||||
/tmp/checkpoint-bench.jsonl
|
||||
```
|
||||
|
||||
@ -17,17 +17,27 @@ Examples::
|
||||
--backends sqlite --updates 1000 --payload-bytes 128 \
|
||||
--repetitions 7 --output snapshot-boundary.jsonl
|
||||
|
||||
TEST_POSTGRES_URI=postgresql://... \
|
||||
PYTHONPATH=. uv run python scripts/benchmark/checkpoint/bench_channels.py \
|
||||
--backends sqlite,postgres --updates 100 --payload-bytes 128 \
|
||||
--output cross-backend.jsonl
|
||||
|
||||
The controller suppresses matrix cells whose estimated cumulative full-mode
|
||||
message payload exceeds ``--max-estimated-full-bytes``. Full mode and every
|
||||
swept delta cadence are skipped together so every emitted result remains
|
||||
comparable and subject to the same safety cap. Use
|
||||
``--allow-large-cases`` only on a machine provisioned for the resulting disk
|
||||
and memory use.
|
||||
|
||||
Postgres is opt-in through ``TEST_POSTGRES_URI``. Each case uses a unique
|
||||
benchmark thread and removes only that thread through the saver's public
|
||||
``adelete_thread`` API after collecting the measurements.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import gc
|
||||
import importlib.metadata
|
||||
import json
|
||||
@ -36,11 +46,13 @@ import platform
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Awaitable, Callable
|
||||
from contextlib import suppress
|
||||
from dataclasses import asdict, dataclass
|
||||
from functools import cache
|
||||
from pathlib import Path
|
||||
from typing import Annotated, Any, Literal, TypedDict
|
||||
from uuid import uuid4
|
||||
|
||||
from langchain_core.messages import AIMessage, AnyMessage, BaseMessage, HumanMessage
|
||||
from langgraph.channels import DeltaChannel
|
||||
@ -69,20 +81,25 @@ _file_size = _common.file_size
|
||||
_peak_rss_bytes = _common.peak_rss_bytes
|
||||
|
||||
Mode = Literal["full", "delta"]
|
||||
Backend = Literal["memory", "sqlite"]
|
||||
Backend = Literal["memory", "sqlite", "postgres"]
|
||||
|
||||
SCHEMA_VERSION = 1
|
||||
BENCHMARK_VERSION = 1
|
||||
SCHEMA_VERSION = 2
|
||||
BENCHMARK_VERSION = 2
|
||||
PRODUCTION_SNAPSHOT_FREQUENCY = DEFAULT_CHECKPOINT_SNAPSHOT_FREQUENCY
|
||||
DEFAULT_MAX_ESTIMATED_FULL_BYTES = 1024**3
|
||||
_MODES: tuple[Mode, ...] = ("full", "delta")
|
||||
_BACKENDS: tuple[Backend, ...] = ("memory", "sqlite")
|
||||
_BACKENDS: tuple[Backend, ...] = ("memory", "sqlite", "postgres")
|
||||
_STORAGE_STAT_FIELDS = (
|
||||
"logical_checkpoint_bytes",
|
||||
"logical_write_bytes",
|
||||
"checkpoint_rows",
|
||||
"checkpoint_bytes",
|
||||
"blob_rows",
|
||||
"blob_bytes",
|
||||
"write_rows",
|
||||
"write_bytes",
|
||||
)
|
||||
_PROCESS_RUN_ID = uuid4().hex[:12]
|
||||
|
||||
|
||||
class _FullBenchmarkState(TypedDict):
|
||||
@ -210,7 +227,7 @@ def _build_graph(mode: Mode, saver: Any, snapshot_frequency: int) -> Any:
|
||||
def _config(case: BenchmarkCase) -> dict[str, Any]:
|
||||
config: dict[str, Any] = {
|
||||
"configurable": {
|
||||
"thread_id": f"checkpoint-bench-{case.seed}-{case.repetition}",
|
||||
"thread_id": (f"checkpoint-bench-{case.seed}-{case.repetition}-{case.backend}-{case.mode}-{case.update_count}-{case.payload_bytes}-{_PROCESS_RUN_ID}"),
|
||||
}
|
||||
}
|
||||
inject_checkpoint_mode(config, case.mode)
|
||||
@ -244,6 +261,12 @@ def _base_row(case: BenchmarkCase) -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _safe_benchmark_error(error: BaseException | str, *, work_dir: Path | None = None) -> str:
|
||||
message = _safe_error(error, work_dir=work_dir)
|
||||
postgres_uri = os.environ.get("TEST_POSTGRES_URI")
|
||||
return message.replace(postgres_uri, "<TEST_POSTGRES_URI>") if postgres_uri else message
|
||||
|
||||
|
||||
def _collect_storage_stats(collector: Callable[[], dict[str, int]]) -> dict[str, Any]:
|
||||
"""Keep timing data usable when a saver's diagnostic layout changes."""
|
||||
try:
|
||||
@ -251,10 +274,43 @@ def _collect_storage_stats(collector: Callable[[], dict[str, int]]) -> dict[str,
|
||||
except Exception as exc:
|
||||
return {
|
||||
**dict.fromkeys(_STORAGE_STAT_FIELDS),
|
||||
"storage_stats_error": _safe_error(exc),
|
||||
"storage_stats_error": _safe_benchmark_error(exc),
|
||||
}
|
||||
|
||||
|
||||
async def _collect_storage_stats_async(collector: Callable[[], Awaitable[dict[str, int]]]) -> dict[str, Any]:
|
||||
"""Async counterpart for saver diagnostics backed by an async cursor."""
|
||||
try:
|
||||
return {**(await collector()), "storage_stats_error": None}
|
||||
except Exception as exc:
|
||||
return {
|
||||
**dict.fromkeys(_STORAGE_STAT_FIELDS),
|
||||
"storage_stats_error": _safe_benchmark_error(exc),
|
||||
}
|
||||
|
||||
|
||||
def _normalized_storage_stats(
|
||||
*,
|
||||
checkpoint_rows: int,
|
||||
checkpoint_bytes: int,
|
||||
blob_rows: int,
|
||||
blob_bytes: int,
|
||||
write_rows: int,
|
||||
write_bytes: int,
|
||||
) -> dict[str, int]:
|
||||
"""Return one backend-neutral checkpoint/blob/write measurement shape."""
|
||||
return {
|
||||
"logical_checkpoint_bytes": checkpoint_bytes + blob_bytes,
|
||||
"logical_write_bytes": write_bytes,
|
||||
"checkpoint_rows": checkpoint_rows,
|
||||
"checkpoint_bytes": checkpoint_bytes,
|
||||
"blob_rows": blob_rows,
|
||||
"blob_bytes": blob_bytes,
|
||||
"write_rows": write_rows,
|
||||
"write_bytes": write_bytes,
|
||||
}
|
||||
|
||||
|
||||
def _memory_storage_stats(saver: InMemorySaver, thread_id: str) -> dict[str, int]:
|
||||
checkpoint_rows = 0
|
||||
checkpoint_bytes = 0
|
||||
@ -262,9 +318,12 @@ def _memory_storage_stats(saver: InMemorySaver, thread_id: str) -> dict[str, int
|
||||
for checkpoint, metadata, _parent_id in namespace.values():
|
||||
checkpoint_rows += 1
|
||||
checkpoint_bytes += len(checkpoint[1]) + len(metadata[1])
|
||||
blob_rows = 0
|
||||
blob_bytes = 0
|
||||
for (stored_thread_id, _namespace, _channel, _version), (_type_tag, blob) in saver.blobs.items():
|
||||
if stored_thread_id == thread_id:
|
||||
checkpoint_bytes += len(blob)
|
||||
blob_rows += 1
|
||||
blob_bytes += len(blob)
|
||||
|
||||
write_rows = 0
|
||||
write_bytes = 0
|
||||
@ -274,12 +333,14 @@ def _memory_storage_stats(saver: InMemorySaver, thread_id: str) -> dict[str, int
|
||||
for _task_id, _channel, (_type_tag, blob), _task_path in writes.values():
|
||||
write_rows += 1
|
||||
write_bytes += len(blob)
|
||||
return {
|
||||
"logical_checkpoint_bytes": checkpoint_bytes,
|
||||
"logical_write_bytes": write_bytes,
|
||||
"checkpoint_rows": checkpoint_rows,
|
||||
"write_rows": write_rows,
|
||||
}
|
||||
return _normalized_storage_stats(
|
||||
checkpoint_rows=checkpoint_rows,
|
||||
checkpoint_bytes=checkpoint_bytes,
|
||||
blob_rows=blob_rows,
|
||||
blob_bytes=blob_bytes,
|
||||
write_rows=write_rows,
|
||||
write_bytes=write_bytes,
|
||||
)
|
||||
|
||||
|
||||
def _sqlite_storage_stats(saver: SqliteSaver, thread_id: str) -> dict[str, int]:
|
||||
@ -292,12 +353,42 @@ def _sqlite_storage_stats(saver: SqliteSaver, thread_id: str) -> dict[str, int]:
|
||||
"SELECT COUNT(*), COALESCE(SUM(length(value)), 0) FROM writes WHERE thread_id = ?",
|
||||
(thread_id,),
|
||||
).fetchone()
|
||||
return {
|
||||
"logical_checkpoint_bytes": int(checkpoint_bytes),
|
||||
"logical_write_bytes": int(write_bytes),
|
||||
"checkpoint_rows": int(checkpoint_rows),
|
||||
"write_rows": int(write_rows),
|
||||
}
|
||||
return _normalized_storage_stats(
|
||||
checkpoint_rows=int(checkpoint_rows),
|
||||
checkpoint_bytes=int(checkpoint_bytes),
|
||||
# SqliteSaver stores channel values inside the serialized checkpoint
|
||||
# payload; unlike PostgresSaver, it has no separate blob table.
|
||||
blob_rows=0,
|
||||
blob_bytes=0,
|
||||
write_rows=int(write_rows),
|
||||
write_bytes=int(write_bytes),
|
||||
)
|
||||
|
||||
|
||||
async def _postgres_storage_stats(saver: Any, thread_id: str) -> dict[str, int]:
|
||||
"""Measure the LangGraph-owned Postgres tables for one benchmark thread."""
|
||||
queries = (
|
||||
"SELECT COUNT(*) AS rows, COALESCE(SUM(pg_column_size(checkpoint) + pg_column_size(metadata)), 0) AS bytes FROM checkpoints WHERE thread_id = %s",
|
||||
"SELECT COUNT(*) AS rows, COALESCE(SUM(octet_length(blob)), 0) AS bytes FROM checkpoint_blobs WHERE thread_id = %s",
|
||||
"SELECT COUNT(*) AS rows, COALESCE(SUM(octet_length(blob)), 0) AS bytes FROM checkpoint_writes WHERE thread_id = %s",
|
||||
)
|
||||
measured: list[tuple[int, int]] = []
|
||||
async with saver._cursor() as cursor:
|
||||
for query in queries:
|
||||
await cursor.execute(query, (thread_id,))
|
||||
row = await cursor.fetchone()
|
||||
if row is None:
|
||||
raise RuntimeError("Postgres storage diagnostic returned no row")
|
||||
measured.append((int(row["rows"]), int(row["bytes"])))
|
||||
(checkpoint_rows, checkpoint_bytes), (blob_rows, blob_bytes), (write_rows, write_bytes) = measured
|
||||
return _normalized_storage_stats(
|
||||
checkpoint_rows=checkpoint_rows,
|
||||
checkpoint_bytes=checkpoint_bytes,
|
||||
blob_rows=blob_rows,
|
||||
blob_bytes=blob_bytes,
|
||||
write_rows=write_rows,
|
||||
write_bytes=write_bytes,
|
||||
)
|
||||
|
||||
|
||||
def _write_and_read(case: BenchmarkCase, saver: Any, messages: list[BaseMessage]) -> tuple[dict[str, Any], list[AnyMessage]]:
|
||||
@ -339,6 +430,42 @@ def _cold_read(case: BenchmarkCase, saver: Any) -> tuple[float, list[AnyMessage]
|
||||
return elapsed_ms, list(snapshot.values.get("messages", []))
|
||||
|
||||
|
||||
async def _awrite_and_read(case: BenchmarkCase, saver: Any, messages: list[BaseMessage]) -> tuple[dict[str, Any], list[AnyMessage]]:
|
||||
graph = _build_graph(case.mode, saver, case.snapshot_frequency)
|
||||
accessor = CheckpointStateAccessor.bind(graph, saver, mode=case.mode)
|
||||
config = _config(case)
|
||||
update_latencies: list[float] = []
|
||||
write_start = time.perf_counter()
|
||||
for message in messages:
|
||||
update_start = time.perf_counter()
|
||||
await graph.ainvoke({"messages": [message]}, config)
|
||||
update_latencies.append((time.perf_counter() - update_start) * 1000)
|
||||
write_total_ms = (time.perf_counter() - write_start) * 1000
|
||||
|
||||
warm_start = time.perf_counter()
|
||||
snapshot = await accessor.aget(config)
|
||||
warm_read_ms = (time.perf_counter() - warm_start) * 1000
|
||||
return {
|
||||
"write_total_ms": write_total_ms,
|
||||
"write_p50_ms": _percentile(update_latencies, 50),
|
||||
"write_p95_ms": _percentile(update_latencies, 95),
|
||||
"write_p99_ms": _percentile(update_latencies, 99),
|
||||
"write_first_window_ms": _window_median(update_latencies, "first"),
|
||||
"write_middle_window_ms": _window_median(update_latencies, "middle"),
|
||||
"write_last_window_ms": _window_median(update_latencies, "last"),
|
||||
"warm_read_ms": warm_read_ms,
|
||||
}, list(snapshot.values.get("messages", []))
|
||||
|
||||
|
||||
async def _acold_read(case: BenchmarkCase, saver: Any) -> tuple[float, list[AnyMessage]]:
|
||||
graph = _build_graph(case.mode, saver, case.snapshot_frequency)
|
||||
accessor = CheckpointStateAccessor.bind(graph, saver, mode=case.mode)
|
||||
gc.collect()
|
||||
start = time.perf_counter()
|
||||
snapshot = await accessor.aget(_config(case))
|
||||
return (time.perf_counter() - start) * 1000, list(snapshot.values.get("messages", []))
|
||||
|
||||
|
||||
def _validate_materialized(case: BenchmarkCase, expected: list[BaseMessage], warm: list[AnyMessage], cold: list[AnyMessage]) -> tuple[int, str]:
|
||||
expected_digest = _canonical_messages_digest(expected)
|
||||
warm_digest = _canonical_messages_digest(warm)
|
||||
@ -446,14 +573,75 @@ def _run_sqlite_case(case: BenchmarkCase, messages: list[BaseMessage], db_path:
|
||||
return result
|
||||
|
||||
|
||||
async def _delete_postgres_benchmark_thread(uri: str, thread_id: str) -> None:
|
||||
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
|
||||
|
||||
async with AsyncPostgresSaver.from_conn_string(uri) as saver:
|
||||
await saver.setup()
|
||||
await saver.adelete_thread(thread_id)
|
||||
|
||||
|
||||
async def _run_postgres_case(case: BenchmarkCase, messages: list[BaseMessage], uri: str) -> dict[str, Any]:
|
||||
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
|
||||
|
||||
cache_opt_in = os.environ.get(_HISTORY_CACHE_ENV) == "1"
|
||||
thread_id = _config(case)["configurable"]["thread_id"]
|
||||
completed = False
|
||||
try:
|
||||
async with AsyncPostgresSaver.from_conn_string(uri) as saver:
|
||||
await saver.setup()
|
||||
write_saver = _wrap_history_cache(saver) if cache_opt_in else saver
|
||||
metrics, warm = await _awrite_and_read(case, write_saver, messages)
|
||||
stats = await _collect_storage_stats_async(lambda: _postgres_storage_stats(saver, thread_id))
|
||||
|
||||
reopen_start = time.perf_counter()
|
||||
async with AsyncPostgresSaver.from_conn_string(uri) as reopened:
|
||||
await reopened.setup()
|
||||
saver_reopen_ms = (time.perf_counter() - reopen_start) * 1000
|
||||
cold_saver = _wrap_history_cache(reopened) if cache_opt_in else reopened
|
||||
cold_read_ms, cold = await _acold_read(case, cold_saver)
|
||||
|
||||
actual_count, digest = _validate_materialized(case, messages, warm, cold)
|
||||
result = {
|
||||
**metrics,
|
||||
**stats,
|
||||
"cold_read_ms": cold_read_ms,
|
||||
"saver_reopen_ms": saver_reopen_ms,
|
||||
"db_bytes": None,
|
||||
"wal_bytes": None,
|
||||
"shm_bytes": None,
|
||||
"durable_db_bytes": None,
|
||||
"expected_message_count": case.update_count,
|
||||
"actual_message_count": actual_count,
|
||||
"content_sha256": digest,
|
||||
}
|
||||
if cache_opt_in:
|
||||
result["history_cache_enabled"] = True
|
||||
result.update(_history_cache_stats(write_saver, "history_cache_write_"))
|
||||
result.update(_history_cache_stats(cold_saver, "history_cache_cold_"))
|
||||
completed = True
|
||||
return result
|
||||
finally:
|
||||
if completed:
|
||||
await _delete_postgres_benchmark_thread(uri, thread_id)
|
||||
else:
|
||||
with suppress(Exception):
|
||||
await _delete_postgres_benchmark_thread(uri, thread_id)
|
||||
|
||||
|
||||
def _run_case(case: BenchmarkCase, *, work_dir: Path) -> dict[str, Any]:
|
||||
row = _base_row(case)
|
||||
messages = [_message_for_update(index, case.payload_bytes) for index in range(case.update_count)]
|
||||
try:
|
||||
if case.backend == "memory":
|
||||
measured = _run_memory_case(case, messages)
|
||||
else:
|
||||
elif case.backend == "sqlite":
|
||||
measured = _run_sqlite_case(case, messages, work_dir / "checkpoint-benchmark.sqlite")
|
||||
else:
|
||||
uri = os.environ.get("TEST_POSTGRES_URI")
|
||||
if not uri:
|
||||
raise RuntimeError("postgres benchmark requires TEST_POSTGRES_URI")
|
||||
measured = asyncio.run(_run_postgres_case(case, messages, uri))
|
||||
|
||||
reducer_writes = [[message] for message in messages]
|
||||
reducer_start = time.perf_counter()
|
||||
@ -467,7 +655,7 @@ def _run_case(case: BenchmarkCase, *, work_dir: Path) -> dict[str, Any]:
|
||||
row["peak_rss_bytes"] = _peak_rss_bytes()
|
||||
except Exception as exc:
|
||||
row["success"] = False
|
||||
row["error"] = _safe_error(exc, work_dir=work_dir)
|
||||
row["error"] = _safe_benchmark_error(exc, work_dir=work_dir)
|
||||
return row
|
||||
|
||||
|
||||
@ -533,7 +721,11 @@ def _run_child_case(case: BenchmarkCase, *, timeout_seconds: float, git_sha: str
|
||||
def _build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(description="Benchmark full and delta checkpoint message channels")
|
||||
parser.add_argument("--modes", default="full,delta", help="Comma-separated modes (default: full,delta)")
|
||||
parser.add_argument("--backends", default="memory,sqlite", help="Comma-separated backends (default: memory,sqlite)")
|
||||
parser.add_argument(
|
||||
"--backends",
|
||||
default="memory,sqlite",
|
||||
help="Comma-separated backends: memory, sqlite, postgres (default: memory,sqlite; postgres requires TEST_POSTGRES_URI)",
|
||||
)
|
||||
parser.add_argument("--updates", default="10,100", help="Comma-separated message update counts (default: 10,100)")
|
||||
parser.add_argument("--payload-bytes", default="128", help="Comma-separated exact message content sizes (default: 128)")
|
||||
parser.add_argument(
|
||||
@ -588,6 +780,8 @@ def main(argv: list[str] | None = None) -> int:
|
||||
try:
|
||||
modes = _parse_choice_csv(args.modes, option="--modes", choices=_MODES)
|
||||
backends = _parse_choice_csv(args.backends, option="--backends", choices=_BACKENDS)
|
||||
if "postgres" in backends and not os.environ.get("TEST_POSTGRES_URI"):
|
||||
raise ValueError("--backends postgres requires TEST_POSTGRES_URI")
|
||||
updates = _parse_positive_int_csv(args.updates, option="--updates")
|
||||
payload_bytes = _parse_positive_int_csv(args.payload_bytes, option="--payload-bytes")
|
||||
snapshot_frequencies = _parse_positive_int_csv(args.snapshot_frequencies, option="--snapshot-frequencies")
|
||||
|
||||
@ -32,6 +32,12 @@ DEFAULT_METRICS = [
|
||||
"write_last_window_ms",
|
||||
"warm_read_ms",
|
||||
"cold_read_ms",
|
||||
"checkpoint_rows",
|
||||
"checkpoint_bytes",
|
||||
"blob_rows",
|
||||
"blob_bytes",
|
||||
"write_rows",
|
||||
"write_bytes",
|
||||
"logical_checkpoint_bytes",
|
||||
"logical_write_bytes",
|
||||
"durable_db_bytes",
|
||||
|
||||
@ -2,8 +2,11 @@ from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
@ -159,6 +162,17 @@ def test_case_rejects_non_positive_snapshot_frequency() -> None:
|
||||
)
|
||||
|
||||
|
||||
def test_shared_backend_thread_ids_are_case_scoped_and_stable() -> None:
|
||||
full = bench.BenchmarkCase(mode="full", backend="postgres", update_count=4, payload_bytes=64, repetition=0, seed=1)
|
||||
delta = bench.BenchmarkCase(mode="delta", backend="postgres", update_count=4, payload_bytes=64, repetition=0, seed=1)
|
||||
|
||||
full_thread_id = bench._config(full)["configurable"]["thread_id"]
|
||||
|
||||
assert full_thread_id == bench._config(full)["configurable"]["thread_id"]
|
||||
assert full_thread_id != bench._config(delta)["configurable"]["thread_id"]
|
||||
assert full_thread_id.startswith("checkpoint-bench-")
|
||||
|
||||
|
||||
def test_memory_smoke_case_materializes_at_low_snapshot_frequency(tmp_path: Path) -> None:
|
||||
case = bench.BenchmarkCase(
|
||||
mode="delta",
|
||||
@ -233,6 +247,9 @@ def test_memory_smoke_case_materializes_expected_state(mode: str, tmp_path: Path
|
||||
assert len(row["content_sha256"]) == 64
|
||||
assert row["db_bytes"] is None
|
||||
assert row["logical_checkpoint_bytes"] is not None
|
||||
assert row["logical_checkpoint_bytes"] == row["checkpoint_bytes"] + row["blob_bytes"]
|
||||
assert row["logical_write_bytes"] == row["write_bytes"]
|
||||
assert row["blob_rows"] > 0
|
||||
|
||||
|
||||
def test_memory_case_keeps_timings_when_private_storage_stats_change(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None:
|
||||
@ -255,8 +272,12 @@ def test_memory_case_keeps_timings_when_private_storage_stats_change(monkeypatch
|
||||
assert row["success"] is True
|
||||
assert row["write_total_ms"] >= 0
|
||||
assert row["logical_checkpoint_bytes"] is None
|
||||
assert row["checkpoint_bytes"] is None
|
||||
assert row["blob_bytes"] is None
|
||||
assert row["logical_write_bytes"] is None
|
||||
assert row["write_bytes"] is None
|
||||
assert row["checkpoint_rows"] is None
|
||||
assert row["blob_rows"] is None
|
||||
assert row["write_rows"] is None
|
||||
assert "private saver layout changed" in row["storage_stats_error"]
|
||||
|
||||
@ -280,7 +301,113 @@ def test_sqlite_smoke_case_reports_durable_and_logical_storage(mode: str, tmp_pa
|
||||
assert row["logical_checkpoint_bytes"] > 0
|
||||
assert row["logical_write_bytes"] > 0
|
||||
assert row["checkpoint_rows"] > 0
|
||||
assert row["checkpoint_bytes"] > 0
|
||||
assert row["blob_rows"] == 0
|
||||
assert row["blob_bytes"] == 0
|
||||
assert row["write_rows"] > 0
|
||||
assert row["write_bytes"] > 0
|
||||
assert row["logical_checkpoint_bytes"] == row["checkpoint_bytes"] + row["blob_bytes"]
|
||||
assert row["logical_write_bytes"] == row["write_bytes"]
|
||||
|
||||
|
||||
class _FakePostgresCursor:
|
||||
def __init__(self, rows: list[dict[str, int]]) -> None:
|
||||
self._rows = iter(rows)
|
||||
self.executed: list[tuple[str, tuple[str, ...]]] = []
|
||||
|
||||
async def execute(self, query: str, params: tuple[str, ...]) -> None:
|
||||
self.executed.append((query, params))
|
||||
|
||||
async def fetchone(self) -> dict[str, int]:
|
||||
return next(self._rows)
|
||||
|
||||
|
||||
class _FakePostgresSaver:
|
||||
def __init__(self, cursor: _FakePostgresCursor) -> None:
|
||||
self.cursor = cursor
|
||||
|
||||
@asynccontextmanager
|
||||
async def _cursor(self) -> Any:
|
||||
yield self.cursor
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_postgres_storage_stats_normalize_checkpoint_tables() -> None:
|
||||
cursor = _FakePostgresCursor(
|
||||
[
|
||||
{"rows": 4, "bytes": 100},
|
||||
{"rows": 7, "bytes": 250},
|
||||
{"rows": 9, "bytes": 400},
|
||||
]
|
||||
)
|
||||
|
||||
stats = await bench._postgres_storage_stats(_FakePostgresSaver(cursor), "thread-1")
|
||||
|
||||
assert stats == {
|
||||
"checkpoint_rows": 4,
|
||||
"checkpoint_bytes": 100,
|
||||
"blob_rows": 7,
|
||||
"blob_bytes": 250,
|
||||
"write_rows": 9,
|
||||
"write_bytes": 400,
|
||||
"logical_checkpoint_bytes": 350,
|
||||
"logical_write_bytes": 400,
|
||||
}
|
||||
assert [params for _query, params in cursor.executed] == [("thread-1",)] * 3
|
||||
assert "FROM checkpoints" in cursor.executed[0][0]
|
||||
assert "FROM checkpoint_blobs" in cursor.executed[1][0]
|
||||
assert "FROM checkpoint_writes" in cursor.executed[2][0]
|
||||
|
||||
|
||||
def test_postgres_backend_requires_explicit_connection_uri(monkeypatch: pytest.MonkeyPatch, tmp_path: Path, capsys: pytest.CaptureFixture[str]) -> None:
|
||||
monkeypatch.delenv("TEST_POSTGRES_URI", raising=False)
|
||||
|
||||
with pytest.raises(SystemExit, match="2"):
|
||||
bench.main(["--backends", "postgres", "--output", str(tmp_path / "results.jsonl")])
|
||||
|
||||
assert "TEST_POSTGRES_URI" in capsys.readouterr().err
|
||||
|
||||
|
||||
def test_postgres_failure_does_not_expose_connection_uri(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None:
|
||||
uri = "postgresql://benchmark-user:secret@example.invalid/checkpoints"
|
||||
monkeypatch.setenv("TEST_POSTGRES_URI", uri)
|
||||
|
||||
async def fail_case(*_args: Any, **_kwargs: Any) -> dict[str, Any]:
|
||||
raise RuntimeError(f"could not connect to {uri}")
|
||||
|
||||
monkeypatch.setattr(bench, "_run_postgres_case", fail_case)
|
||||
case = bench.BenchmarkCase(mode="delta", backend="postgres", update_count=1, payload_bytes=32, repetition=0, seed=1)
|
||||
|
||||
row = bench._run_case(case, work_dir=tmp_path)
|
||||
|
||||
assert row["success"] is False
|
||||
assert uri not in row["error"]
|
||||
assert "<TEST_POSTGRES_URI>" in row["error"]
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_postgres_smoke_case_reports_normalized_storage() -> None:
|
||||
uri = os.environ.get("TEST_POSTGRES_URI")
|
||||
if not uri:
|
||||
pytest.skip("TEST_POSTGRES_URI is not set")
|
||||
|
||||
case = bench.BenchmarkCase(
|
||||
mode="delta",
|
||||
backend="postgres",
|
||||
update_count=3,
|
||||
payload_bytes=64,
|
||||
repetition=0,
|
||||
seed=4,
|
||||
)
|
||||
row = await bench._run_postgres_case(case, [bench._message_for_update(index, case.payload_bytes) for index in range(case.update_count)], uri)
|
||||
|
||||
assert row["actual_message_count"] == 3
|
||||
assert row["checkpoint_rows"] > 0
|
||||
assert row["checkpoint_bytes"] > 0
|
||||
assert row["blob_rows"] > 0
|
||||
assert row["blob_bytes"] > 0
|
||||
assert row["write_rows"] > 0
|
||||
assert row["write_bytes"] > 0
|
||||
|
||||
|
||||
def test_controller_writes_versioned_jsonl_without_sensitive_case_paths(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None:
|
||||
@ -290,8 +417,8 @@ def test_controller_writes_versioned_jsonl_without_sensitive_case_paths(monkeypa
|
||||
def fake_run_child(case, *, timeout_seconds, git_sha, profile_dir=None):
|
||||
child_git_shas.append(git_sha)
|
||||
return {
|
||||
"schema_version": 1,
|
||||
"benchmark_version": 1,
|
||||
"schema_version": bench.SCHEMA_VERSION,
|
||||
"benchmark_version": bench.BENCHMARK_VERSION,
|
||||
"success": True,
|
||||
"error": None,
|
||||
"mode": case.mode,
|
||||
@ -328,7 +455,7 @@ def test_controller_writes_versioned_jsonl_without_sensitive_case_paths(monkeypa
|
||||
assert rc == 0
|
||||
rows = [json.loads(line) for line in output.read_text(encoding="utf-8").splitlines()]
|
||||
assert [row["mode"] for row in rows] == ["full", "delta"]
|
||||
assert all(row["schema_version"] == 1 for row in rows)
|
||||
assert all(row["schema_version"] == bench.SCHEMA_VERSION for row in rows)
|
||||
assert all("work_dir" not in row and "database_path" not in row for row in rows)
|
||||
assert child_git_shas == ["controller-sha", "controller-sha"]
|
||||
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user