From 8fc513879b80e7c93ffc0966f51d340e84a919b7 Mon Sep 17 00:00:00 2001 From: Eldon Marks Date: Wed, 12 Aug 2026 15:48:36 -0400 Subject: [PATCH 1/9] fix: honor JVSPATIAL_WEBHOOK_API_KEY_REQUIRE_HTTPS in env adapter Allowlist and map the query-param HTTPS gate (plus WEBHOOK_HTTPS_REQUIRED) into ServerConfig so local HTTP tunnels can disable it via env. Co-authored-by: Cursor --- CHANGELOG.md | 11 ++++++++ docs/md/environment-keys-reference.md | 1 + jvspatial/env_adapter.py | 13 +++++++++ tests/api/test_env_allowlist_audit.py | 2 ++ tests/api/test_webhook_env_overrides.py | 35 +++++++++++++++++++++++++ 5 files changed, 62 insertions(+) create mode 100644 tests/api/test_webhook_env_overrides.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 2dc35d6..ddd153c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,17 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Fixed + +- **`JVSPATIAL_WEBHOOK_API_KEY_REQUIRE_HTTPS` was allowlist-rejected** + (`jvspatial/env_adapter.py`). The key was never in `ALLOWED_ENV_KEYS`, so + startup warned that it was ignored and + `server_config_overrides_from_env()` never mapped it onto + `WebhookConfig.webhook_api_key_require_https`. Local HTTP callback tunnels + (ngrok → plain `http://127.0.0.1`) could not disable the query-param HTTPS + gate via env alone. Also maps `JVSPATIAL_WEBHOOK_HTTPS_REQUIRED` into + `WebhookConfig.webhook_https_required` for the same ServerConfig path. + ### Changed - **`uvicorn` is capped below 1.0** (`pyproject.toml`). It was floor-only diff --git a/docs/md/environment-keys-reference.md b/docs/md/environment-keys-reference.md index 932db0f..c579ca8 100644 --- a/docs/md/environment-keys-reference.md +++ b/docs/md/environment-keys-reference.md @@ -125,6 +125,7 @@ For full examples and default values, see: - `JVSPATIAL_WEBHOOK_MAX_PAYLOAD_SIZE` - Max payload bytes. - `JVSPATIAL_WEBHOOK_IDEMPOTENCY_TTL` - Idempotency TTL. - `JVSPATIAL_WEBHOOK_HTTPS_REQUIRED` - HTTPS-only webhook policy. +- `JVSPATIAL_WEBHOOK_API_KEY_REQUIRE_HTTPS` - Require HTTPS when the webhook API key is supplied via query parameter (default `true`). Set `false` for local plain-HTTP tunnels (e.g. ngrok → `http://127.0.0.1`). ### Walker safety and deferred execution - `JVSPATIAL_WALKER_PROTECTION_ENABLED` - Enables walker safety guards. diff --git a/jvspatial/env_adapter.py b/jvspatial/env_adapter.py index c65523d..0a9d384 100644 --- a/jvspatial/env_adapter.py +++ b/jvspatial/env_adapter.py @@ -217,6 +217,18 @@ def server_config_overrides_from_env() -> Dict[str, Any]: if proxy: o["proxy"] = proxy + webhook: Dict[str, Any] = {} + if "JVSPATIAL_WEBHOOK_API_KEY_REQUIRE_HTTPS" in os.environ: + webhook["webhook_api_key_require_https"] = _parse_bool( + os.environ["JVSPATIAL_WEBHOOK_API_KEY_REQUIRE_HTTPS"] + ) + if "JVSPATIAL_WEBHOOK_HTTPS_REQUIRED" in os.environ: + webhook["webhook_https_required"] = _parse_bool( + os.environ["JVSPATIAL_WEBHOOK_HTTPS_REQUIRED"] + ) + if webhook: + o["webhook"] = webhook + return o @@ -322,6 +334,7 @@ def server_config_overrides_from_env() -> Dict[str, Any]: "JVSPATIAL_EVENTBRIDGE_SCHEDULER_GROUP", "JVSPATIAL_LWA_ENV_DEFAULTS", # Webhooks + "JVSPATIAL_WEBHOOK_API_KEY_REQUIRE_HTTPS", "JVSPATIAL_WEBHOOK_HMAC_ALGORITHM", "JVSPATIAL_WEBHOOK_HMAC_SECRET", "JVSPATIAL_WEBHOOK_HTTPS_REQUIRED", diff --git a/tests/api/test_env_allowlist_audit.py b/tests/api/test_env_allowlist_audit.py index 519d334..bd6339c 100644 --- a/tests/api/test_env_allowlist_audit.py +++ b/tests/api/test_env_allowlist_audit.py @@ -69,5 +69,7 @@ def test_allowlist_contains_canonical_keys(): "JVSPATIAL_DOCS_DISABLED", "JVSPATIAL_WALKER_MAX_STEPS", "JVSPATIAL_CORS_ORIGINS", + "JVSPATIAL_WEBHOOK_API_KEY_REQUIRE_HTTPS", + "JVSPATIAL_WEBHOOK_HTTPS_REQUIRED", } assert expected.issubset(ALLOWED_ENV_KEYS) diff --git a/tests/api/test_webhook_env_overrides.py b/tests/api/test_webhook_env_overrides.py new file mode 100644 index 0000000..302c61d --- /dev/null +++ b/tests/api/test_webhook_env_overrides.py @@ -0,0 +1,35 @@ +"""Webhook HTTPS env keys → ServerConfig.webhook via env adapter.""" + +import os +from unittest.mock import patch + +from jvspatial.api.config import ServerConfig +from jvspatial.env_adapter import deep_merge, server_config_overrides_from_env + + +def test_webhook_api_key_require_https_from_env(): + with patch.dict( + os.environ, + {"JVSPATIAL_WEBHOOK_API_KEY_REQUIRE_HTTPS": "false"}, + clear=False, + ): + overrides = server_config_overrides_from_env() + assert overrides["webhook"]["webhook_api_key_require_https"] is False + + merged = deep_merge(ServerConfig().model_dump(), overrides) + config = ServerConfig(**merged) + assert config.webhook.webhook_api_key_require_https is False + + +def test_webhook_https_required_from_env(): + with patch.dict( + os.environ, + {"JVSPATIAL_WEBHOOK_HTTPS_REQUIRED": "false"}, + clear=False, + ): + overrides = server_config_overrides_from_env() + assert overrides["webhook"]["webhook_https_required"] is False + + merged = deep_merge(ServerConfig().model_dump(), overrides) + config = ServerConfig(**merged) + assert config.webhook.webhook_https_required is False From 23af4c058b36a4cb9933d3cffd6222db61e05cb3 Mon Sep 17 00:00:00 2001 From: Eldon Marks Date: Fri, 11 Sep 2026 13:59:21 -0400 Subject: [PATCH 2/9] bench: add hub-node scale recorder and 0.0.17 baseline Phase 0 of the hub-node scale remediation. Seeds a Postgres hub with 1k/10k/100k edges (COPY) and records p50/p95 latency, DB round trips (db_op_counter) and on-disk sizes for connect(), save(), neighbour listings, the len(nodes()) count pattern and 32-way concurrent connect(). Opt-in via `-m bench` (`bench_slow` for 100k). The baseline confirms the expected shape on 0.0.17: connect()/save() grow with hub degree (save 11 ms -> 1.15 s from 1k to 100k), concurrent connects serialise on the hub row lock (11.7 s wall for 32 at 100k), and list-form nodes(limit=20) costs 1 + ceil(degree/500) + 1 round trips (201 at 100k) while the class form stays at 1. Co-Authored-By: Claude Opus 5 --- CHANGELOG.md | 10 + CONTRIBUTING.md | 4 + docs/bench/2026-09-hub-node-baseline.md | 114 ++++++ docs/bench/2026-09-hub-node-phase0.jsonl | 3 + docs/md/benchmarks.md | 16 + pyproject.toml | 2 + tests/benchmarks/hub_bench_report.py | 112 ++++++ tests/benchmarks/test_hub_node_bench.py | 427 +++++++++++++++++++++++ 8 files changed, 688 insertions(+) create mode 100644 docs/bench/2026-09-hub-node-baseline.md create mode 100644 docs/bench/2026-09-hub-node-phase0.jsonl create mode 100644 tests/benchmarks/hub_bench_report.py create mode 100644 tests/benchmarks/test_hub_node_bench.py diff --git a/CHANGELOG.md b/CHANGELOG.md index ddd153c..ddde33f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,16 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- **Hub-node scale recorder** (`tests/benchmarks/test_hub_node_bench.py`, + `tests/benchmarks/hub_bench_report.py`). Seeds a Postgres hub with 1k / 10k / + 100k edges and records p50/p95 latency, DB round trips and on-disk sizes + for `connect()`, `save()`, neighbour listings, the `len(nodes())` count + pattern and 32-way concurrent `connect()`. Opt-in via `-m bench` + (`bench_slow` for 100k). Baseline and per-change results: + `docs/bench/2026-09-hub-node-baseline.md`. + ### Fixed - **`JVSPATIAL_WEBHOOK_API_KEY_REQUIRE_HTTPS` was allowlist-rejected** diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 01e5833..1b3b90d 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -142,6 +142,10 @@ flake8 jvspatial/ tests/ # performance benchmarks (only run when asked; not part of -q above) pytest tests/benchmarks --benchmark-only + +# hub-node scale recorder (Postgres; 100k tier opt-in via bench_slow) +JVSPATIAL_POSTGRES_TEST_DSN=postgresql://... \ + pytest tests/benchmarks/test_hub_node_bench.py -m bench -s ``` For details on the benchmark suite (what's in it, what CI does with diff --git a/docs/bench/2026-09-hub-node-baseline.md b/docs/bench/2026-09-hub-node-baseline.md new file mode 100644 index 0000000..2863773 --- /dev/null +++ b/docs/bench/2026-09-hub-node-baseline.md @@ -0,0 +1,114 @@ +# Hub-node scale baseline (Postgres object-spatial layer) + +**Date:** 2026-09-11 +**Purpose:** fix the "before" numbers for the hub-node scale remediation, then +append an "after" table per phase. Every phase cites this file. + +## How the numbers are produced + +Harness: [`tests/benchmarks/test_hub_node_bench.py`](../../tests/benchmarks/test_hub_node_bench.py) +(recorder, not a pytest-benchmark bench). Render tables with +[`tests/benchmarks/hub_bench_report.py`](../../tests/benchmarks/hub_bench_report.py). + +```bash +docker run -d --name jvspatial-bench-pg -p 55432:5432 \ + -e POSTGRES_USER=jvspatial -e POSTGRES_PASSWORD=jvspatial -e POSTGRES_DB=jvspatial \ + pgvector/pgvector:pg16 -c shared_buffers=512MB -c max_connections=200 + +JVSPATIAL_POSTGRES_TEST_DSN=postgresql://jvspatial:jvspatial@localhost:55432/jvspatial \ +JVSPATIAL_BENCH_RESULTS=docs/bench/.jsonl \ + pytest tests/benchmarks/test_hub_node_bench.py -m "bench or bench_slow" -s -p no:randomly +python tests/benchmarks/hub_bench_report.py docs/bench/.jsonl +``` + +Per tier (degree ∈ {1k, 10k, 100k}), on a fresh schema with the default +`Node` / `Edge` / bench-class indexes ensured: + +- **Seed** (COPY via `bulk_save_detailed`): one hub with `degree` outgoing + `BenchContains` edges to `BenchLeaf` nodes, one sink with `degree` incoming + `BenchContains` edges from the same leaves, plus 82 unconnected spare leaves. + Records mirror the active adjacency mode — with persisted edge ids the hub + and sink rows each carry `degree` ids in `edges`. +- **Reads** — 20 samples (5 at 100k), context cache cleared before each so + every sample is cold; round trips counted with `db_op_counter` through + `create_database(..., observe=True)`. +- **`save()`** — 50 samples of `hub.counter = i; await hub.save()`. +- **`connect()`** — 50 sequential `hub.connect(spare_leaf, edge=BenchContains)`. +- **Concurrency** — 32 `hub.connect()` calls to distinct spare leaves under + `asyncio.gather` (pool `max_size=40`, so the pool is not the bottleneck). +- **Sizes** — `pg_column_size` of the hub row, `node_data_gin`, and total + relation sizes, after seeding and again after the 82 hub writes. + +Raw records: `2026-09-hub-node-phase0.jsonl` (one JSON object per tier). + +**Machine.** Apple M1 Pro (10 cores, 32 GB), macOS 15.4. Postgres 16.14 +(`pgvector/pgvector:pg16` in Docker Desktop, `shared_buffers=512MB`). Python +3.11.10, asyncpg 0.31.0. Client and server share the host (loopback), so a +round trip costs ~0.1 ms here. Over a real network every extra round trip +multiplies, so the round-trip counts matter as much as the latencies. + +## Before — jvspatial 0.0.17 (`8fc5138`, `edge_ids_mode=persist`) + +| operation | 1k p50 / p95 ms (trips) | 10k p50 / p95 ms (trips) | 100k p50 / p95 ms (trips) | +|---|---|---|---| +| `ctx.get(Hub)` (hydrate hub) | 1.0 / 3.6 | 4.2 / 9.2 | 34.1 / 43.6 | +| `hub.connect(leaf, edge=E)` | 7.7 / 11.7 (5) | 21.5 / 48.5 (5) | 162.3 / 558.5 (5) | +| `hub.save()` after scalar change | 11.3 / 13.7 | 107.0 / 123.6 | 1151.3 / 1535.1 | +| `hub.nodes(edge=[E], node=['Leaf'], limit=20)` | 39.5 / 64.3 (3) | 448.2 / 488.6 (21) | 4821.4 / 4835.3 (201) | +| `hub.nodes(edge=E, limit=20)` | 1.3 / 3.7 (1) | 1.1 / 4.0 (1) | 2.0 / 3.9 (1) | +| `sink.nodes(edge=[E], node=['Leaf'], direction='in', limit=20)` | 46.7 / 70.8 (3) | 469.9 / 513.6 (21) | 5050.2 / 5104.4 (201) | +| `sink.nodes(edge=[E], node=['Leaf'], direction='in')` | 39.2 / 72.3 (3) | 485.2 / 513.4 (21) | 4972.2 / 5101.1 (201) | +| `len(await hub.nodes(edge=[E]))` | 42.8 / 70.8 (3) | 467.8 / 508.6 (21) | 4682.5 / 4927.7 (201) | +| 32× concurrent `connect()` (ms) | 573 wall / 568 max | 1131 wall / 1121 max | 11665 wall / 11657 max | + +| size (after seed → after 82 hub writes) | 1k | 10k | 100k | +|---|---|---|---| +| hub row `pg_column_size(data)` | 25.9 KB → 27.9 KB | 257.2 KB → 251.0 KB | 2.5 MB → 2.4 MB | +| `node_data_gin` size | 192.0 KB → 3.3 MB | 1.6 MB → 8.6 MB | 23.4 MB → 47.3 MB | +| `node` table total | 784.0 KB → 7.6 MB | 6.4 MB → 47.1 MB | 71.0 MB → 337.3 MB | +| `edge` table total | 1.7 MB → 1.8 MB | 16.3 MB → 18.6 MB | 163.8 MB → 163.9 MB | + +(`sink.nodes(..., direction='in')` stands in for the brief's +`hub.nodes(..., direction="in")`: the hub has no incoming edges, so the sink +carries the same fan-out inbound.) + +### Interpretation + +The expected shape holds. Every cost that should scale with the request +instead scales with the hub's degree. + +- **Writes.** `connect()` stays at 5 round trips but gets about 21× slower + from 1k to 100k. `save()` gets about 100× slower (11 ms → 1.15 s). + - Each `connect()` makes two `atomic_add_edge_id` calls. Each call takes + `SELECT … FOR UPDATE` on the node row, decodes the whole `edges` array, + applies `$addToSet` in Python, and rewrites the row. + - `save()` unions the array in SQL (`save_with_edge_merge`) and rewrites a + 2.5 MB row. +- **Row-lock serialisation.** 32 concurrent `connect()`s to one hub queue + behind each other on the row lock. The slowest call waits for the other 31 + (568 ms at 1k, 11.7 s at 100k), so wall time equals the worst single call. +- **List-form reads.** `nodes(edge=[E], node=[…], limit=20)` skips the join + fast path. It loads every outgoing edge, hydrates every neighbour in + 500-id `find_many` chunks, then filters and slices in Python. + - Round trips are `1 + ⌈degree/500⌉ + 1`: 201 at 100k, which is 4.8 s to + return 20 rows. + - The class form (`edge=E`) takes the join fast path and stays at 1 round + trip and about 1–2 ms at every tier. + - The `len(nodes())` count anti-pattern costs the same as a full listing. +- **Index / heap bloat.** Only 82 rewrites of the hub row grow `node_data_gin` + 17× at 1k (192 KB → 3.3 MB) and double it at 100k. The `node` table grows + from 71 MB to 337 MB, all dead TOASTed copies of the hub row. The + whole-document `jsonb_path_ops` GIN re-indexes every array element on every + rewrite. + +## After Phase 1 + +_Pending._ + +## After Phase 2 + +_Pending._ + +## After Phase 3 + +_Pending._ diff --git a/docs/bench/2026-09-hub-node-phase0.jsonl b/docs/bench/2026-09-hub-node-phase0.jsonl new file mode 100644 index 0000000..998ea69 --- /dev/null +++ b/docs/bench/2026-09-hub-node-phase0.jsonl @@ -0,0 +1,3 @@ +{"degree": 1000, "jvspatial": "0.0.17", "git_sha": "8fc5138", "edge_ids_mode": "persist", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 2.45, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 26559, "node_gin_bytes": 196608, "node_total_bytes": 802816, "edge_total_bytes": 1802240}, "hub_get": {"p50_ms": 1.023, "p95_ms": 3.572, "max_ms": 5.647, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 39.496, "p95_ms": 64.317, "max_ms": 65.316, "n": 20, "round_trips": 3}, "nodes_class_out_limit20": {"p50_ms": 1.323, "p95_ms": 3.678, "max_ms": 6.555, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 46.712, "p95_ms": 70.827, "max_ms": 76.017, "n": 20, "round_trips": 3}, "nodes_list_in_unlimited": {"p50_ms": 39.198, "p95_ms": 72.297, "max_ms": 73.201, "n": 20, "round_trips": 3}, "count_via_len_nodes": {"p50_ms": 42.823, "p95_ms": 70.776, "max_ms": 80.152, "n": 20, "round_trips": 3}, "hub_save": {"p50_ms": 11.318, "p95_ms": 13.688, "max_ms": 29.778, "n": 50}, "hub_connect": {"p50_ms": 7.651, "p95_ms": 11.7, "max_ms": 15.058, "n": 50, "round_trips": 5}, "concurrent_connect_32": {"wall_ms": 573.063, "max_call_ms": 568.168, "p50_call_ms": 519.322}, "sizes_after_writes": {"hub_row_bytes": 28579, "node_gin_bytes": 3440640, "node_total_bytes": 7929856, "edge_total_bytes": 1900544}} +{"degree": 10000, "jvspatial": "0.0.17", "git_sha": "8fc5138", "edge_ids_mode": "persist", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 23.99, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 263377, "node_gin_bytes": 1638400, "node_total_bytes": 6684672, "edge_total_bytes": 17072128}, "hub_get": {"p50_ms": 4.169, "p95_ms": 9.209, "max_ms": 11.369, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 448.219, "p95_ms": 488.609, "max_ms": 495.432, "n": 20, "round_trips": 21}, "nodes_class_out_limit20": {"p50_ms": 1.144, "p95_ms": 3.987, "max_ms": 5.476, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 469.865, "p95_ms": 513.627, "max_ms": 550.486, "n": 20, "round_trips": 21}, "nodes_list_in_unlimited": {"p50_ms": 485.197, "p95_ms": 513.422, "max_ms": 515.798, "n": 20, "round_trips": 21}, "count_via_len_nodes": {"p50_ms": 467.759, "p95_ms": 508.647, "max_ms": 521.43, "n": 20, "round_trips": 21}, "hub_save": {"p50_ms": 106.992, "p95_ms": 123.644, "max_ms": 137.609, "n": 50}, "hub_connect": {"p50_ms": 21.548, "p95_ms": 48.476, "max_ms": 52.092, "n": 50, "round_trips": 5}, "concurrent_connect_32": {"wall_ms": 1131.011, "max_call_ms": 1121.222, "p50_call_ms": 780.366}, "sizes_after_writes": {"hub_row_bytes": 257067, "node_gin_bytes": 8994816, "node_total_bytes": 49348608, "edge_total_bytes": 19529728}} +{"degree": 100000, "jvspatial": "0.0.17", "git_sha": "8fc5138", "edge_ids_mode": "persist", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 251.84, "read_iterations": 5, "sizes_after_seed": {"hub_row_bytes": 2634438, "node_gin_bytes": 24551424, "node_total_bytes": 74498048, "edge_total_bytes": 171761664}, "hub_get": {"p50_ms": 34.101, "p95_ms": 43.617, "max_ms": 43.617, "n": 5}, "nodes_list_out_limit20": {"p50_ms": 4821.444, "p95_ms": 4835.256, "max_ms": 4835.256, "n": 5, "round_trips": 201}, "nodes_class_out_limit20": {"p50_ms": 2.05, "p95_ms": 3.891, "max_ms": 3.891, "n": 5, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 5050.231, "p95_ms": 5104.35, "max_ms": 5104.35, "n": 5, "round_trips": 201}, "nodes_list_in_unlimited": {"p50_ms": 4972.176, "p95_ms": 5101.09, "max_ms": 5101.09, "n": 5, "round_trips": 201}, "count_via_len_nodes": {"p50_ms": 4682.547, "p95_ms": 4927.673, "max_ms": 4927.673, "n": 5, "round_trips": 201}, "hub_save": {"p50_ms": 1151.34, "p95_ms": 1535.099, "max_ms": 1589.655, "n": 50}, "hub_connect": {"p50_ms": 162.314, "p95_ms": 558.513, "max_ms": 582.77, "n": 50, "round_trips": 5}, "concurrent_connect_32": {"wall_ms": 11665.166, "max_call_ms": 11657.035, "p50_call_ms": 6012.509}, "sizes_after_writes": {"hub_row_bytes": 2465290, "node_gin_bytes": 49577984, "node_total_bytes": 353665024, "edge_total_bytes": 171835392}} diff --git a/docs/md/benchmarks.md b/docs/md/benchmarks.md index bd1b2c5..f4f0f76 100644 --- a/docs/md/benchmarks.md +++ b/docs/md/benchmarks.md @@ -59,6 +59,22 @@ The current benches guard the IO hot paths landed in Phases A1 and A2: seeded chain graph. * `test_bench_postgres_find_many_bulk` -- bulk fetch by id list. +* **Hub-node scale recorder** (`tests/benchmarks/test_hub_node_bench.py`) — + not a pytest-benchmark bench. Seeds a hub with 1k / 10k / 100k edges and + records p50/p95 latency, DB round trips (`db_op_counter`) and on-disk + sizes for `connect()`, `save()`, neighbour listings, the `len(nodes())` + count pattern and 32-way concurrent `connect()`. Marked `bench` (100k tier + `bench_slow`), so it runs only when selected: + + ```bash + JVSPATIAL_POSTGRES_TEST_DSN=postgresql://... \ + JVSPATIAL_BENCH_RESULTS=run.jsonl \ + pytest tests/benchmarks/test_hub_node_bench.py -m "bench or bench_slow" -s + python tests/benchmarks/hub_bench_report.py run.jsonl # markdown tables + ``` + + Recorded results live in `docs/bench/2026-09-hub-node-baseline.md`. + The `_fallback_*` benches deliberately exercise the *slow* path so that future contributors who refactor the translator can see whether the legacy in-Python filter path got faster or slower. diff --git a/pyproject.toml b/pyproject.toml index 44715a5..5094089 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -142,6 +142,8 @@ markers = [ "asyncio: mark test as an asyncio test", "s3: mark test as requiring boto3 for S3 storage", "benchmark: performance regression benchmark (run with --benchmark-only)", + "bench: hub-node scale recorder (run with pytest tests/benchmarks/test_hub_node_bench.py -m bench -s)", + "bench_slow: 100k-degree tier of the hub-node recorder (opt-in: -m 'bench or bench_slow')", ] filterwarnings = [ "ignore::DeprecationWarning:pydantic.*", diff --git a/tests/benchmarks/hub_bench_report.py b/tests/benchmarks/hub_bench_report.py new file mode 100644 index 0000000..c0b93dd --- /dev/null +++ b/tests/benchmarks/hub_bench_report.py @@ -0,0 +1,112 @@ +"""Render hub-node bench JSONL (``$JVSPATIAL_BENCH_RESULTS``) as markdown tables. + +Usage:: + + python tests/benchmarks/hub_bench_report.py docs/bench/2026-09-hub-node-phase0.jsonl + +Prints one latency table (p50 / p95 ms, round trips) and one size table, +columns ordered by degree. Paste the output into the bench document. +""" + +from __future__ import annotations + +import json +import sys +from typing import Any, Dict, List + +_LATENCY_ROWS = [ + ("hub_get", "`ctx.get(Hub)` (hydrate hub)"), + ("hub_connect", "`hub.connect(leaf, edge=E)`"), + ("hub_save", "`hub.save()` after scalar change"), + ("nodes_list_out_limit20", "`hub.nodes(edge=[E], node=['Leaf'], limit=20)`"), + ("nodes_class_out_limit20", "`hub.nodes(edge=E, limit=20)`"), + ( + "nodes_list_in_limit20", + "`sink.nodes(edge=[E], node=['Leaf'], direction='in', limit=20)`", + ), + ( + "nodes_list_in_unlimited", + "`sink.nodes(edge=[E], node=['Leaf'], direction='in')`", + ), + ("count_via_len_nodes", "`len(await hub.nodes(edge=[E]))`"), + ("count_nodes", "`hub.count_nodes(edge=[E])`"), +] + +_SIZE_ROWS = [ + ("hub_row_bytes", "hub row `pg_column_size(data)`"), + ("node_gin_bytes", "`node_data_gin` size"), + ("node_total_bytes", "`node` table total"), + ("edge_total_bytes", "`edge` table total"), +] + + +def _fmt_bytes(n: float) -> str: + for unit in ("B", "KB", "MB", "GB"): + if n < 1024 or unit == "GB": + return f"{n:.0f} {unit}" if unit == "B" else f"{n:.1f} {unit}" + n /= 1024 + return str(n) + + +def _label(degree: int) -> str: + return f"{degree // 1000}k" if degree % 1000 == 0 else str(degree) + + +def render(records: List[Dict[str, Any]]) -> str: + records = sorted(records, key=lambda r: r["degree"]) + heads = [_label(r["degree"]) for r in records] + out: List[str] = [] + meta = records[0] + out.append( + f"jvspatial {meta['jvspatial']} @ `{meta['git_sha']}`, " + f"edge_ids_mode=`{meta['edge_ids_mode']}`, Postgres {meta['postgres']}" + ) + out.append("") + out.append( + "| operation | " + " | ".join(f"{h} p50 / p95 ms (trips)" for h in heads) + " |" + ) + out.append("|---|" + "---|" * len(heads)) + for key, label in _LATENCY_ROWS: + if not any(key in r for r in records): + continue + cells = [] + for r in records: + m = r.get(key) + if not m: + cells.append("—") + continue + trips = f" ({m['round_trips']})" if "round_trips" in m else "" + cells.append(f"{m['p50_ms']:.1f} / {m['p95_ms']:.1f}{trips}") + out.append(f"| {label} | " + " | ".join(cells) + " |") + cells = [] + for r in records: + c = r["concurrent_connect_32"] + cells.append(f"{c['wall_ms']:.0f} wall / {c['max_call_ms']:.0f} max") + out.append("| 32× concurrent `connect()` (ms) | " + " | ".join(cells) + " |") + out.append("") + out.append( + "| size (after seed → after 82 hub writes) | " + " | ".join(heads) + " |" + ) + out.append("|---|" + "---|" * len(heads)) + for key, label in _SIZE_ROWS: + cells = [ + f"{_fmt_bytes(r['sizes_after_seed'][key])} → " + f"{_fmt_bytes(r['sizes_after_writes'][key])}" + for r in records + ] + out.append(f"| {label} | " + " | ".join(cells) + " |") + return "\n".join(out) + + +def main(argv: List[str]) -> int: + if len(argv) != 2: + print(__doc__) + return 2 + with open(argv[1], encoding="utf-8") as fh: + records = [json.loads(line) for line in fh if line.strip()] + print(render(records)) + return 0 + + +if __name__ == "__main__": + sys.exit(main(sys.argv)) diff --git a/tests/benchmarks/test_hub_node_bench.py b/tests/benchmarks/test_hub_node_bench.py new file mode 100644 index 0000000..3d1d03e --- /dev/null +++ b/tests/benchmarks/test_hub_node_bench.py @@ -0,0 +1,427 @@ +"""Hub-node scale benchmark for the Postgres object-spatial layer. + +Records the cost of the operations whose price grows with the *fan-out* of a +node rather than with the size of the request: ``connect()`` / ``save()`` on a +hub, neighbour listings with and without type filters, the +``len(await n.nodes())`` count anti-pattern, the on-disk size of the hub row, +and 32-way concurrent ``connect()`` against one hub (row-lock serialisation). + +This is a *recorder*, not a pytest-benchmark regression bench: each tier +prints p50/p95 latencies, DB round trips (``db_op_counter``) and sizes, and +optionally appends a JSON record to ``$JVSPATIAL_BENCH_RESULTS``. The numbers +feed ``docs/bench/2026-09-hub-node-baseline.md``. + +Run (Postgres reachable at ``JVSPATIAL_POSTGRES_TEST_DSN``):: + + pytest tests/benchmarks/test_hub_node_bench.py -m bench -s + # include the 100k-degree tier + pytest tests/benchmarks/test_hub_node_bench.py -m "bench or bench_slow" -s + +Skips cleanly when asyncpg is missing or the DSN is unreachable. Each tier +runs in a throwaway schema that is dropped afterwards. +""" + +from __future__ import annotations + +import asyncio +import contextlib +import functools +import json +import math +import os +import subprocess +import time +import uuid +from typing import Any, AsyncIterator, Awaitable, Callable, Dict, List, Tuple + +import pytest + +try: + import asyncpg +except ImportError: # pragma: no cover - optional dependency + asyncpg = None # type: ignore[assignment] + +import jvspatial +from jvspatial.core import context as context_module +from jvspatial.core.context import GraphContext, set_default_context +from jvspatial.core.entities import Edge, Node +from jvspatial.core.utils import generate_id +from jvspatial.db import create_database +from jvspatial.observability import db_op_counter + +_DSN = os.getenv( + "JVSPATIAL_POSTGRES_TEST_DSN", + "postgresql://jvspatial:jvspatial@localhost:5432/jvspatial", +) +_RESULTS_PATH = os.getenv("JVSPATIAL_BENCH_RESULTS") + +pytestmark = [pytest.mark.bench] + +# Latency samples per measured operation. +_WRITE_ITERATIONS = 50 +_CONCURRENT_CONNECTS = 32 +_SEED_CHUNK = 25_000 + + +class BenchHub(Node): + """Hub (and sink) node — the high fan-out endpoint.""" + + label: str = "" + counter: int = 0 + + +class BenchLeaf(Node): + """Leaf node hanging off the hub.""" + + idx: int = 0 + title: str = "" + + +class BenchContains(Edge): + """Typed containment edge (hub -> leaf, leaf -> sink).""" + + +# ---- helpers ---------------------------------------------------------------- + + +def _pct(samples: List[float], p: float) -> float: + ordered = sorted(samples) + k = max(0, math.ceil(p / 100.0 * len(ordered)) - 1) + return ordered[k] + + +def _summary(samples_ms: List[float]) -> Dict[str, float]: + return { + "p50_ms": round(_pct(samples_ms, 50), 3), + "p95_ms": round(_pct(samples_ms, 95), 3), + "max_ms": round(max(samples_ms), 3), + "n": len(samples_ms), + } + + +async def _timed(fn: Callable[[], Awaitable[Any]]) -> Tuple[float, int, Any]: + """Run ``fn`` once; return (elapsed_ms, db round trips, result).""" + token = db_op_counter.set(0) + try: + t0 = time.perf_counter() + result = await fn() + elapsed = (time.perf_counter() - t0) * 1000.0 + return elapsed, db_op_counter.get(), result + finally: + db_op_counter.reset(token) + + +async def _dsn_reachable() -> bool: + if asyncpg is None: + return False + try: + conn = await asyncio.wait_for(asyncpg.connect(dsn=_DSN), timeout=2.0) + except Exception: + return False + await conn.close() + return True + + +def _git_sha() -> str: + try: + return ( + subprocess.check_output( + ["git", "rev-parse", "--short", "HEAD"], stderr=subprocess.DEVNULL + ) + .decode() + .strip() + ) + except Exception: + return "unknown" + + +@contextlib.asynccontextmanager +async def _bench_graph() -> AsyncIterator[Tuple[GraphContext, Any, str]]: + """Yield ``(context, observable_db, schema)`` on a fresh schema.""" + schema = f"bench_hub_{uuid.uuid4().hex[:10]}" + admin = await asyncpg.connect(dsn=_DSN) + await admin.execute(f"CREATE SCHEMA {schema}") + db = create_database( + "postgres", + dsn=_DSN, + schema_name=schema, + min_size=1, + max_size=_CONCURRENT_CONNECTS + 8, + observe=True, + slow_query_ms=1e9, + ) + ctx = GraphContext(database=db) + set_default_context(ctx) + # ``ensure_indexes`` memoises per collection:entity in a module global; + # every tier runs on a fresh schema, so the memo must not leak across. + context_module._ensured_indexes.clear() + try: + for cls in (Node, Edge, BenchHub, BenchLeaf, BenchContains): + await ctx.ensure_indexes(cls) + yield ctx, db, schema + finally: + set_default_context(None) + context_module._ensured_indexes.clear() + with contextlib.suppress(Exception): + await db.close() # type: ignore[attr-defined] + await admin.execute(f"DROP SCHEMA {schema} CASCADE") + await admin.close() + + +def _persists_edge_ids(ctx: GraphContext) -> bool: + probe = getattr(ctx, "persists_edge_ids", None) + return bool(probe()) if callable(probe) else True + + +async def _seed(ctx: GraphContext, db: Any, degree: int, spare: int) -> Dict[str, Any]: + """Seed hub -> ``degree`` leaves and ``degree`` leaves -> sink via COPY. + + Also seeds ``spare`` unconnected leaves for the write measurements. + Records mirror the persisted format of the active adjacency mode: + with persisted edge ids the hub row carries ``degree`` ids in ``edges``. + """ + persist = _persists_edge_ids(ctx) + hub_id = generate_id("n", "BenchHub") + sink_id = generate_id("n", "BenchHub") + leaf_ids = [generate_id("n", "BenchLeaf") for _ in range(degree)] + spare_ids = [generate_id("n", "BenchLeaf") for _ in range(spare)] + out_eids = [generate_id("e", "BenchContains") for _ in range(degree)] + in_eids = [generate_id("e", "BenchContains") for _ in range(degree)] + + nodes: List[Dict[str, Any]] = [] + edges: List[Dict[str, Any]] = [] + for i, (lid, oe, ie) in enumerate(zip(leaf_ids, out_eids, in_eids)): + rec: Dict[str, Any] = { + "id": lid, + "entity": "BenchLeaf", + "context": {"idx": i, "title": f"leaf {i}"}, + } + if persist: + rec["edges"] = [oe, ie] + nodes.append(rec) + edges.append( + { + "id": oe, + "entity": "BenchContains", + "context": {}, + "source": hub_id, + "target": lid, + "bidirectional": False, + } + ) + edges.append( + { + "id": ie, + "entity": "BenchContains", + "context": {}, + "source": lid, + "target": sink_id, + "bidirectional": False, + } + ) + for j, sid in enumerate(spare_ids): + rec = { + "id": sid, + "entity": "BenchLeaf", + "context": {"idx": degree + j, "title": f"spare {j}"}, + } + if persist: + rec["edges"] = [] + nodes.append(rec) + for hid, label, eids in ((hub_id, "hub", out_eids), (sink_id, "sink", in_eids)): + rec = { + "id": hid, + "entity": "BenchHub", + "context": {"label": label, "counter": 0}, + } + if persist: + rec["edges"] = list(eids) + nodes.append(rec) + + for coll, recs in (("node", nodes), ("edge", edges)): + for off in range(0, len(recs), _SEED_CHUNK): + result = await db.bulk_save_detailed(coll, recs[off : off + _SEED_CHUNK]) + assert result.all_saved, f"seed {coll} lost rows: {result.failed_ids[:5]}" + + return { + "hub_id": hub_id, + "sink_id": sink_id, + "spare_ids": spare_ids, + "persist": persist, + } + + +async def _sizes(admin: Any, schema: str, hub_id: str) -> Dict[str, int]: + row = await admin.fetchrow( + f"SELECT pg_column_size(data) AS hub_bytes FROM {schema}.node WHERE id = $1", + hub_id, + ) + rel = await admin.fetchrow( + f""" + SELECT pg_relation_size('{schema}.node_data_gin') AS node_gin_bytes, + pg_total_relation_size('{schema}.node') AS node_total_bytes, + pg_total_relation_size('{schema}.edge') AS edge_total_bytes + """ + ) + return { + "hub_row_bytes": int(row["hub_bytes"]), + "node_gin_bytes": int(rel["node_gin_bytes"]), + "node_total_bytes": int(rel["node_total_bytes"]), + "edge_total_bytes": int(rel["edge_total_bytes"]), + } + + +# ---- the benchmark ---------------------------------------------------------- + + +@pytest.mark.parametrize( + "degree", + [ + pytest.param(1_000, id="1k"), + pytest.param(10_000, id="10k"), + pytest.param(100_000, id="100k", marks=pytest.mark.bench_slow), + ], +) +async def test_hub_node_scale(degree: int, request: pytest.FixtureRequest) -> None: + markexpr = request.config.getoption("markexpr") or "" + if degree >= 100_000 and "bench_slow" not in markexpr: + pytest.skip("100k tier is opt-in: -m 'bench or bench_slow'") + if not await _dsn_reachable(): + pytest.skip(f"Postgres unreachable at {_DSN}; set JVSPATIAL_POSTGRES_TEST_DSN") + + read_iters = 20 if degree < 100_000 else 5 + spare = _WRITE_ITERATIONS + _CONCURRENT_CONNECTS + results: Dict[str, Any] = {} + + async with _bench_graph() as (ctx, db, schema): + admin = await asyncpg.connect(dsn=_DSN) + try: + t0 = time.perf_counter() + seeded = await _seed(ctx, db, degree, spare) + seed_s = time.perf_counter() - t0 + await admin.execute(f"ANALYZE {schema}.node") + await admin.execute(f"ANALYZE {schema}.edge") + results["sizes_after_seed"] = await _sizes(admin, schema, seeded["hub_id"]) + + # -- hub hydration (the persisted edge list rides along) -- + samples: List[float] = [] + for _ in range(read_iters): + await ctx.clear_cache() + ms, _, hub = await _timed(lambda: ctx.get(BenchHub, seeded["hub_id"])) + samples.append(ms) + results["hub_get"] = _summary(samples) + assert hub is not None + sink = await ctx.get(BenchHub, seeded["sink_id"]) + assert sink is not None + + # -- neighbour reads -- + read_cases: Dict[str, Tuple[Callable[[], Awaitable[Any]], int]] = { + "nodes_list_out_limit20": ( + lambda: hub.nodes( + edge=[BenchContains], node=["BenchLeaf"], limit=20 + ), + 20, + ), + "nodes_class_out_limit20": ( + lambda: hub.nodes(edge=BenchContains, limit=20), + 20, + ), + "nodes_list_in_limit20": ( + lambda: sink.nodes( + edge=[BenchContains], + node=["BenchLeaf"], + direction="in", + limit=20, + ), + 20, + ), + "nodes_list_in_unlimited": ( + lambda: sink.nodes( + edge=[BenchContains], node=["BenchLeaf"], direction="in" + ), + degree, + ), + "count_via_len_nodes": ( + lambda: hub.nodes(edge=[BenchContains]), + degree, + ), + } + for name, (fn, expected) in read_cases.items(): + samples = [] + trips = 0 + for _ in range(read_iters): + await ctx.clear_cache() + ms, trips, out = await _timed(fn) + samples.append(ms) + assert len(out) == expected, f"{name}: {len(out)} != {expected}" + results[name] = {**_summary(samples), "round_trips": trips} + + # -- save() after a scalar change -- + samples = [] + for i in range(_WRITE_ITERATIONS): + hub.counter = i + 1 + + async def _save() -> Any: + return await hub.save() + + ms, _, _ = await _timed(_save) + samples.append(ms) + results["hub_save"] = _summary(samples) + + # -- sequential connect() of fresh leaves -- + spare_ids = seeded["spare_ids"] + seq_leaves = [ + await ctx.get(BenchLeaf, sid) for sid in spare_ids[:_WRITE_ITERATIONS] + ] + samples = [] + trips = 0 + for leaf in seq_leaves: + ms, trips, _ = await _timed( + functools.partial(hub.connect, leaf, edge=BenchContains) + ) + samples.append(ms) + results["hub_connect"] = {**_summary(samples), "round_trips": trips} + + # -- 32-way concurrent connect() to the same hub -- + conc_leaves = [ + await ctx.get(BenchLeaf, sid) for sid in spare_ids[_WRITE_ITERATIONS:] + ] + + async def _one(leaf: BenchLeaf) -> float: + t = time.perf_counter() + await hub.connect(leaf, edge=BenchContains) + return (time.perf_counter() - t) * 1000.0 + + t0 = time.perf_counter() + per_call = await asyncio.gather(*(_one(leaf) for leaf in conc_leaves)) + wall_ms = (time.perf_counter() - t0) * 1000.0 + results["concurrent_connect_32"] = { + "wall_ms": round(wall_ms, 3), + "max_call_ms": round(max(per_call), 3), + "p50_call_ms": round(_pct(list(per_call), 50), 3), + } + + final_degree = await db.count("edge", {"source": seeded["hub_id"]}) + assert final_degree == degree + spare + results["sizes_after_writes"] = await _sizes( + admin, schema, seeded["hub_id"] + ) + + pg_version = await admin.fetchval("SHOW server_version") + finally: + await admin.close() + + record = { + "degree": degree, + "jvspatial": jvspatial.__version__, + "git_sha": _git_sha(), + "edge_ids_mode": "persist" if seeded["persist"] else "derive", + "postgres": pg_version, + "seed_seconds": round(seed_s, 2), + "read_iterations": read_iters, + **results, + } + print(f"\n[hub-bench] degree={degree}\n" + json.dumps(record, indent=2)) + if _RESULTS_PATH: + with open(_RESULTS_PATH, "a", encoding="utf-8") as fh: + fh.write(json.dumps(record) + "\n") From dc2a6cc1fa5e37359630830e81ca7c4d8570c570 Mon Sep 17 00:00:00 2001 From: Eldon Marks Date: Fri, 11 Sep 2026 14:49:04 -0400 Subject: [PATCH 3/9] Derive node adjacency from the edge table on Postgres, MongoDB, SQLite Phase 1 of the hub-node scale remediation. Node rows no longer persist their incident edge ids in an `edges` array on backends whose edge collection is indexed on source/target (`edge_ids_mode="derive"`, new `Database` capability flag; JsonDB and DynamoDB keep "persist"). connect()/disconnect()/save() never rewrite or row-lock a node, so their cost no longer grows with degree and concurrent connects to one hub no longer serialise. edges(), connection_count(), cascade delete, expand_node/subgraph_bfs and Root rehydration read the edge collection; legacy arrays are ignored on read and dropped on the next save. - resolve_edge_ids_mode(): instance attribute > JVSPATIAL_NODE_EDGE_IDS env > adapter class default (wrappers unwrapped via `inner`) - `jvspatial migrate strip-node-edges` (dry run by default, --apply), backed by strip_node_edges() on PostgresDB (one PK keyset pass) and MongoDB (batched $unset) - expand_node pages from the edge collection by id; adds keyset `after` / `pagination.next_after` alongside the int cursor - Fix: QueryEngine.apply_update ignored $pull, leaving stale edge ids on Postgres rows in persist mode after disconnect/delete - Bench harness: warm pool and re-open connections before the burst; baseline re-measured, Phase 1 numbers in the bench doc Bench (100k hub, p50): connect 160 -> 1.9 ms, save 1161 -> 0.7 ms, 32x concurrent connect 10.3 s -> 19 ms wall, hub row 2.5 MB -> 136 B. Co-Authored-By: Claude Opus 5 --- CHANGELOG.md | 35 ++ SPEC.md | 3 +- docs/bench/2026-09-hub-node-baseline.md | 90 +++- docs/bench/2026-09-hub-node-phase0.jsonl | 6 +- docs/bench/2026-09-hub-node-phase1.jsonl | 3 + docs/md/entity-reference.md | 8 +- docs/md/environment-keys-reference.md | 1 + docs/md/graph-context.md | 20 +- docs/md/postgres-guide.md | 28 ++ .../api/endpoints/graph_visualization.py | 7 +- jvspatial/cli.py | 109 ++++- jvspatial/core/context.py | 59 ++- jvspatial/core/entities/node.py | 141 ++++-- jvspatial/core/entities/root.py | 7 +- jvspatial/core/graph_expansion.py | 113 +++-- jvspatial/db/README.md | 3 +- jvspatial/db/database.py | 50 +++ jvspatial/db/mongodb.py | 49 +++ jvspatial/db/postgres.py | 68 ++- jvspatial/db/query.py | 9 + jvspatial/db/sqlite.py | 5 + jvspatial/env_adapter.py | 1 + tests/benchmarks/test_hub_node_bench.py | 18 +- tests/core/test_edge_ids_derive.py | 415 ++++++++++++++++++ tests/core/test_entity_crud_and_cascade.py | 4 + tests/core/test_node_save_edges_merge.py | 3 + tests/db/test_jsondb.py | 6 + tests/db/test_query_apply_update.py | 28 ++ tests/test_node_operations.py | 3 + 29 files changed, 1161 insertions(+), 131 deletions(-) create mode 100644 docs/bench/2026-09-hub-node-phase1.jsonl create mode 100644 tests/core/test_edge_ids_derive.py create mode 100644 tests/db/test_query_apply_update.py diff --git a/CHANGELOG.md b/CHANGELOG.md index ddde33f..974c32b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -16,9 +16,44 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 pattern and 32-way concurrent `connect()`. Opt-in via `-m bench` (`bench_slow` for 100k). Baseline and per-change results: `docs/bench/2026-09-hub-node-baseline.md`. +- **Node adjacency mode API** (`jvspatial/db/database.py`, + `jvspatial/core/context.py`): `Database.edge_ids_mode` capability flag, + `resolve_edge_ids_mode(db)`, `GraphContext.persists_edge_ids()`, + `Node.edges(limit=...)`, `strip_node_edges()` on `PostgresDB` / `MongoDB`, + and the `jvspatial migrate strip-node-edges` CLI (dry run by default, + `--apply` to write). +- **`expand_node` keyset paging** (`jvspatial/core/graph_expansion.py`): + `after=` / `pagination.next_after` (and the `after` query parameter on the + graph expand endpoint) page incident edges by edge id in O(page). The + integer `cursor` offset keeps working. + +### Changed + +- **Postgres, MongoDB and SQLite no longer persist node adjacency on node + rows** (`edge_ids_mode="derive"`). The indexed edge collection is the only + source of truth: `connect()`, `disconnect()` and `save()` never rewrite or + row-lock a node, so their cost no longer grows with the node's degree and + concurrent connects to one hub no longer serialise. `Node.edges()`, + `connection_count()`, cascade delete, `expand_node` and `subgraph_bfs` read + the edge collection; `Node.edge_ids` stays an empty in-memory list. Legacy + `edges` arrays on existing rows are ignored on read and dropped on the next + save — run `jvspatial migrate strip-node-edges --dsn … --apply` (safe + online) to reclaim the space at once. JsonDB and DynamoDB keep `"persist"`. + Opt out with `JVSPATIAL_NODE_EDGE_IDS=persist` or + `db.edge_ids_mode = "persist"`. +- **`expand_node` / `subgraph_bfs` source incident edges from the edge + collection** in every mode (one query per node instead of one `get` per + edge id); `total_edge_count` is a count of incident edges, and neighbours + are fetched with a single `find_many`. ### Fixed +- **`$pull` was ignored by `QueryEngine.apply_update`** (`jvspatial/db/query.py`). + Backends that apply updates in Python — Postgres `find_one_and_update`, + and the JsonDB / SQLite / DynamoDB defaults — silently skipped it, so on + Postgres in persist mode `disconnect()`, edge deletion and cascade delete + left stale edge ids on node rows (and `connection_count()` over-reported). + - **`JVSPATIAL_WEBHOOK_API_KEY_REQUIRE_HTTPS` was allowlist-rejected** (`jvspatial/env_adapter.py`). The key was never in `ALLOWED_ENV_KEYS`, so startup warned that it was ignored and diff --git a/SPEC.md b/SPEC.md index 5a74fa5..a1f0efb 100644 --- a/SPEC.md +++ b/SPEC.md @@ -111,7 +111,7 @@ AttributeMixin + pydantic.BaseModel ### 2.2 Node — graph node `jvspatial/core/entities/node.py:34+`. Adds: -- `edge_ids: List[str]` — transient in memory, persisted at the top level as `edges` (`Node._get_top_level_fields` line 55-60) +- `edge_ids: List[str]` — transient in-memory list. Persisted at the top level as `edges` (`Node._get_top_level_fields`) **only** when the backend's adjacency mode is `persist` (JsonDB, DynamoDB). In `derive` mode (Postgres, MongoDB, SQLite default) node rows carry no `edges`, `connect()` / `disconnect()` / `save()` never write or lock a node row, and adjacency is read from the indexed edge collection (`Node.edges`, `Node.connection_count`, cascade delete, `expand_node`); `edge_ids` stays empty and a legacy stored `edges` array is ignored on read and dropped on the next save (`GraphContext.persists_edge_ids`, `context.py`; §4.2) - `_visitor_ref: weakref` — currently visiting walker, transient - `_visit_hooks: ClassVar` — populated by `__init_subclass__` from `@on_visit`-decorated methods (line 62+) @@ -251,6 +251,7 @@ cases are documented rather than normalized: Adapters declare capabilities as class attributes: - `supports_transactions: bool` — `True` for MongoDB (replica set); `False` for SQLite (best-effort), JSON, DynamoDB. +- `edge_ids_mode: str` — where node adjacency lives. `"derive"` for Postgres, MongoDB, SQLite: the edge collection (indexed on `source` / `target` by `Edge.get_indexes`) is the only source of truth. `"persist"` for JSON, DynamoDB: node rows also carry an `edges` array. Resolve the effective value with `resolve_edge_ids_mode(db)` (`db/database.py`): an `edge_ids_mode` set on the adapter instance → `JVSPATIAL_NODE_EDGE_IDS` (`persist` | `derive`) → class default; observability / cache wrappers are unwrapped via `inner`. Rows written in persist mode are converted with `jvspatial migrate strip-node-edges` (`strip_node_edges()` on `PostgresDB` / `MongoDB`). Callers branching on capabilities should test the flag, not the adapter class. diff --git a/docs/bench/2026-09-hub-node-baseline.md b/docs/bench/2026-09-hub-node-baseline.md index 2863773..e83eb76 100644 --- a/docs/bench/2026-09-hub-node-baseline.md +++ b/docs/bench/2026-09-hub-node-baseline.md @@ -29,17 +29,22 @@ Per tier (degree ∈ {1k, 10k, 100k}), on a fresh schema with the default `BenchContains` edges from the same leaves, plus 82 unconnected spare leaves. Records mirror the active adjacency mode — with persisted edge ids the hub and sink rows each carry `degree` ids in `edges`. +- **Pool** — `min_size = max_size = 40`, and connections are re-opened right + before the concurrency burst (asyncpg closes connections idle > 300 s, which + the 100k read phase exceeds), so samples measure queries and locking, not + connection establishment. - **Reads** — 20 samples (5 at 100k), context cache cleared before each so every sample is cold; round trips counted with `db_op_counter` through `create_database(..., observe=True)`. - **`save()`** — 50 samples of `hub.counter = i; await hub.save()`. - **`connect()`** — 50 sequential `hub.connect(spare_leaf, edge=BenchContains)`. - **Concurrency** — 32 `hub.connect()` calls to distinct spare leaves under - `asyncio.gather` (pool `max_size=40`, so the pool is not the bottleneck). + `asyncio.gather`. - **Sizes** — `pg_column_size` of the hub row, `node_data_gin`, and total relation sizes, after seeding and again after the 82 hub writes. -Raw records: `2026-09-hub-node-phase0.jsonl` (one JSON object per tier). +Raw records: `2026-09-hub-node-phase.jsonl` (one JSON object per tier). +Each phase's code is benchmarked from a clean worktree with the same harness. **Machine.** Apple M1 Pro (10 cores, 32 GB), macOS 15.4. Postgres 16.14 (`pgvector/pgvector:pg16` in Docker Desktop, `shared_buffers=512MB`). Python @@ -51,34 +56,36 @@ multiplies, so the round-trip counts matter as much as the latencies. | operation | 1k p50 / p95 ms (trips) | 10k p50 / p95 ms (trips) | 100k p50 / p95 ms (trips) | |---|---|---|---| -| `ctx.get(Hub)` (hydrate hub) | 1.0 / 3.6 | 4.2 / 9.2 | 34.1 / 43.6 | -| `hub.connect(leaf, edge=E)` | 7.7 / 11.7 (5) | 21.5 / 48.5 (5) | 162.3 / 558.5 (5) | -| `hub.save()` after scalar change | 11.3 / 13.7 | 107.0 / 123.6 | 1151.3 / 1535.1 | -| `hub.nodes(edge=[E], node=['Leaf'], limit=20)` | 39.5 / 64.3 (3) | 448.2 / 488.6 (21) | 4821.4 / 4835.3 (201) | -| `hub.nodes(edge=E, limit=20)` | 1.3 / 3.7 (1) | 1.1 / 4.0 (1) | 2.0 / 3.9 (1) | -| `sink.nodes(edge=[E], node=['Leaf'], direction='in', limit=20)` | 46.7 / 70.8 (3) | 469.9 / 513.6 (21) | 5050.2 / 5104.4 (201) | -| `sink.nodes(edge=[E], node=['Leaf'], direction='in')` | 39.2 / 72.3 (3) | 485.2 / 513.4 (21) | 4972.2 / 5101.1 (201) | -| `len(await hub.nodes(edge=[E]))` | 42.8 / 70.8 (3) | 467.8 / 508.6 (21) | 4682.5 / 4927.7 (201) | -| 32× concurrent `connect()` (ms) | 573 wall / 568 max | 1131 wall / 1121 max | 11665 wall / 11657 max | +| `ctx.get(Hub)` (hydrate hub) | 1.0 / 1.5 | 3.9 / 5.1 | 32.4 / 47.7 | +| `hub.connect(leaf, edge=E)` | 6.6 / 7.9 (5) | 23.2 / 52.3 (5) | 160.2 / 581.5 (5) | +| `hub.save()` after scalar change | 11.5 / 13.6 | 110.9 / 137.2 | 1160.6 / 1661.6 | +| `hub.nodes(edge=[E], node=['Leaf'], limit=20)` | 40.7 / 64.5 (3) | 439.6 / 485.6 (21) | 4936.1 / 5184.8 (201) | +| `hub.nodes(edge=E, limit=20)` | 1.1 / 2.0 (1) | 1.1 / 1.7 (1) | 1.3 / 3.0 (1) | +| `sink.nodes(edge=[E], node=['Leaf'], direction='in', limit=20)` | 45.6 / 72.1 (3) | 474.4 / 511.8 (21) | 5317.0 / 5338.7 (201) | +| `sink.nodes(edge=[E], node=['Leaf'], direction='in')` | 39.3 / 70.5 (3) | 470.0 / 523.0 (21) | 5218.3 / 5379.3 (201) | +| `len(await hub.nodes(edge=[E]))` | 43.3 / 69.0 (3) | 472.2 / 511.0 (21) | 4807.5 / 4881.5 (201) | +| 32× concurrent `connect()` (ms) | 134 wall / 133 max | 813 wall / 813 max | 10276 wall / 10244 max | | size (after seed → after 82 hub writes) | 1k | 10k | 100k | |---|---|---|---| -| hub row `pg_column_size(data)` | 25.9 KB → 27.9 KB | 257.2 KB → 251.0 KB | 2.5 MB → 2.4 MB | -| `node_data_gin` size | 192.0 KB → 3.3 MB | 1.6 MB → 8.6 MB | 23.4 MB → 47.3 MB | -| `node` table total | 784.0 KB → 7.6 MB | 6.4 MB → 47.1 MB | 71.0 MB → 337.3 MB | -| `edge` table total | 1.7 MB → 1.8 MB | 16.3 MB → 18.6 MB | 163.8 MB → 163.9 MB | +| hub row `pg_column_size(data)` | 26.0 KB → 27.8 KB | 257.5 KB → 251.0 KB | 2.5 MB → 2.4 MB | +| `node_data_gin` size | 192.0 KB → 3.3 MB | 1.6 MB → 8.5 MB | 25.5 MB → 47.3 MB | +| `node` table total | 784.0 KB → 7.6 MB | 6.4 MB → 47.0 MB | 72.8 MB → 346.4 MB | +| `edge` table total | 1.8 MB → 1.9 MB | 16.3 MB → 18.6 MB | 163.9 MB → 164.0 MB | (`sink.nodes(..., direction='in')` stands in for the brief's `hub.nodes(..., direction="in")`: the hub has no incoming edges, so the sink -carries the same fan-out inbound.) +carries the same fan-out inbound. The first recorded run used a cold +`min_size=1` pool, which folded asyncpg connection setup into the samples; +it was superseded by this warm-pool run of the same code.) ### Interpretation The expected shape holds. Every cost that should scale with the request instead scales with the hub's degree. -- **Writes.** `connect()` stays at 5 round trips but gets about 21× slower - from 1k to 100k. `save()` gets about 100× slower (11 ms → 1.15 s). +- **Writes.** `connect()` stays at 5 round trips but gets about 24× slower + from 1k to 100k (p50). `save()` gets about 100× slower (11.5 ms → 1.16 s). - Each `connect()` makes two `atomic_add_edge_id` calls. Each call takes `SELECT … FOR UPDATE` on the node row, decodes the whole `edges` array, applies `$addToSet` in Python, and rewrites the row. @@ -86,24 +93,59 @@ instead scales with the hub's degree. 2.5 MB row. - **Row-lock serialisation.** 32 concurrent `connect()`s to one hub queue behind each other on the row lock. The slowest call waits for the other 31 - (568 ms at 1k, 11.7 s at 100k), so wall time equals the worst single call. + (133 ms at 1k, 10.2 s at 100k), so wall time equals the worst single call. - **List-form reads.** `nodes(edge=[E], node=[…], limit=20)` skips the join fast path. It loads every outgoing edge, hydrates every neighbour in 500-id `find_many` chunks, then filters and slices in Python. - - Round trips are `1 + ⌈degree/500⌉ + 1`: 201 at 100k, which is 4.8 s to + - Round trips are `1 + ⌈degree/500⌉ + 1`: 201 at 100k, which is ~5 s to return 20 rows. - The class form (`edge=E`) takes the join fast path and stays at 1 round trip and about 1–2 ms at every tier. - The `len(nodes())` count anti-pattern costs the same as a full listing. - **Index / heap bloat.** Only 82 rewrites of the hub row grow `node_data_gin` - 17× at 1k (192 KB → 3.3 MB) and double it at 100k. The `node` table grows - from 71 MB to 337 MB, all dead TOASTed copies of the hub row. The + 17× at 1k (192 KB → 3.3 MB) and nearly double it at 100k. The `node` table + grows from 73 MB to 346 MB, all dead TOASTed copies of the hub row. The whole-document `jsonb_path_ops` GIN re-indexes every array element on every rewrite. -## After Phase 1 +## After Phase 1 — adjacency derived from the edge table (`edge_ids_mode=derive`) -_Pending._ +Code: `23af4c0` + the Phase 1 change set (worktree). + +| operation | 1k p50 / p95 ms (trips) | 10k p50 / p95 ms (trips) | 100k p50 / p95 ms (trips) | +|---|---|---|---| +| `ctx.get(Hub)` (hydrate hub) | 0.5 / 0.8 | 0.8 / 1.0 | 0.6 / 3.7 | +| `hub.connect(leaf, edge=E)` | 1.8 / 2.5 (3) | 1.7 / 2.6 (3) | 1.9 / 3.0 (3) | +| `hub.save()` after scalar change | 0.8 / 1.4 | 0.8 / 3.5 | 0.7 / 1.1 | +| `hub.nodes(edge=[E], node=['Leaf'], limit=20)` | 42.1 / 64.5 (3) | 459.2 / 505.0 (21) | 4850.6 / 5425.8 (201) | +| `hub.nodes(edge=E, limit=20)` | 1.3 / 2.7 (1) | 1.2 / 1.7 (1) | 1.4 / 3.7 (1) | +| `sink.nodes(edge=[E], node=['Leaf'], direction='in', limit=20)` | 42.5 / 75.5 (3) | 474.4 / 525.2 (21) | 5038.7 / 5254.3 (201) | +| `sink.nodes(edge=[E], node=['Leaf'], direction='in')` | 45.4 / 77.5 (3) | 479.2 / 515.2 (21) | 5047.6 / 5195.6 (201) | +| `len(await hub.nodes(edge=[E]))` | 43.6 / 73.2 (3) | 466.9 / 501.1 (21) | 4642.9 / 4915.5 (201) | +| 32× concurrent `connect()` (ms) | 18 wall / 18 max | 29 wall / 29 max | 19 wall / 18 max | + +| size (after seed → after 82 hub writes) | 1k | 10k | 100k | +|---|---|---|---| +| hub row `pg_column_size(data)` | 134 B → 136 B | 134 B → 136 B | 134 B → 136 B | +| `node_data_gin` size | 104.0 KB → 112.0 KB | 1.0 MB → 1.6 MB | 13.4 MB → 13.4 MB | +| `node` table total | 504.0 KB → 512.0 KB | 4.2 MB → 4.8 MB | 44.6 MB → 44.6 MB | +| `edge` table total | 1.8 MB → 1.9 MB | 16.5 MB → 16.6 MB | 162.0 MB → 162.1 MB | + +**Gate: passed.** + +- `connect()` and `save()` are flat across 1k / 10k / 100k. + - `connect()` p95 is 2.5 / 2.6 / 3.0 ms, and the 100k tier costs 1.1× the + 1k tier (A1 asks for ≤ 1.5×). Before, it was 7.9 / 52 / 582 ms. + - `save()` p50 is 0.7–0.8 ms at every degree. Before, it was 11.5 ms → + 1.16 s. + - `connect()` drops from 5 round trips to 3: the two node-row + read-modify-writes are gone. +- The hub row is 136 B at every degree. After the 82 writes the node GIN and + table sizes no longer grow at 1k and 100k; before, they multiplied. +- 32-way concurrent `connect()` takes 18–29 ms wall at every degree (≈ 10–16× + a single call). That cost is event-loop CPU, not locking. Before, it was + 133 ms → 10.2 s, and grew with degree. +- List-form reads are unchanged, as expected. They are Phase 2's target. ## After Phase 2 diff --git a/docs/bench/2026-09-hub-node-phase0.jsonl b/docs/bench/2026-09-hub-node-phase0.jsonl index 998ea69..a730af5 100644 --- a/docs/bench/2026-09-hub-node-phase0.jsonl +++ b/docs/bench/2026-09-hub-node-phase0.jsonl @@ -1,3 +1,3 @@ -{"degree": 1000, "jvspatial": "0.0.17", "git_sha": "8fc5138", "edge_ids_mode": "persist", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 2.45, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 26559, "node_gin_bytes": 196608, "node_total_bytes": 802816, "edge_total_bytes": 1802240}, "hub_get": {"p50_ms": 1.023, "p95_ms": 3.572, "max_ms": 5.647, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 39.496, "p95_ms": 64.317, "max_ms": 65.316, "n": 20, "round_trips": 3}, "nodes_class_out_limit20": {"p50_ms": 1.323, "p95_ms": 3.678, "max_ms": 6.555, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 46.712, "p95_ms": 70.827, "max_ms": 76.017, "n": 20, "round_trips": 3}, "nodes_list_in_unlimited": {"p50_ms": 39.198, "p95_ms": 72.297, "max_ms": 73.201, "n": 20, "round_trips": 3}, "count_via_len_nodes": {"p50_ms": 42.823, "p95_ms": 70.776, "max_ms": 80.152, "n": 20, "round_trips": 3}, "hub_save": {"p50_ms": 11.318, "p95_ms": 13.688, "max_ms": 29.778, "n": 50}, "hub_connect": {"p50_ms": 7.651, "p95_ms": 11.7, "max_ms": 15.058, "n": 50, "round_trips": 5}, "concurrent_connect_32": {"wall_ms": 573.063, "max_call_ms": 568.168, "p50_call_ms": 519.322}, "sizes_after_writes": {"hub_row_bytes": 28579, "node_gin_bytes": 3440640, "node_total_bytes": 7929856, "edge_total_bytes": 1900544}} -{"degree": 10000, "jvspatial": "0.0.17", "git_sha": "8fc5138", "edge_ids_mode": "persist", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 23.99, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 263377, "node_gin_bytes": 1638400, "node_total_bytes": 6684672, "edge_total_bytes": 17072128}, "hub_get": {"p50_ms": 4.169, "p95_ms": 9.209, "max_ms": 11.369, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 448.219, "p95_ms": 488.609, "max_ms": 495.432, "n": 20, "round_trips": 21}, "nodes_class_out_limit20": {"p50_ms": 1.144, "p95_ms": 3.987, "max_ms": 5.476, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 469.865, "p95_ms": 513.627, "max_ms": 550.486, "n": 20, "round_trips": 21}, "nodes_list_in_unlimited": {"p50_ms": 485.197, "p95_ms": 513.422, "max_ms": 515.798, "n": 20, "round_trips": 21}, "count_via_len_nodes": {"p50_ms": 467.759, "p95_ms": 508.647, "max_ms": 521.43, "n": 20, "round_trips": 21}, "hub_save": {"p50_ms": 106.992, "p95_ms": 123.644, "max_ms": 137.609, "n": 50}, "hub_connect": {"p50_ms": 21.548, "p95_ms": 48.476, "max_ms": 52.092, "n": 50, "round_trips": 5}, "concurrent_connect_32": {"wall_ms": 1131.011, "max_call_ms": 1121.222, "p50_call_ms": 780.366}, "sizes_after_writes": {"hub_row_bytes": 257067, "node_gin_bytes": 8994816, "node_total_bytes": 49348608, "edge_total_bytes": 19529728}} -{"degree": 100000, "jvspatial": "0.0.17", "git_sha": "8fc5138", "edge_ids_mode": "persist", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 251.84, "read_iterations": 5, "sizes_after_seed": {"hub_row_bytes": 2634438, "node_gin_bytes": 24551424, "node_total_bytes": 74498048, "edge_total_bytes": 171761664}, "hub_get": {"p50_ms": 34.101, "p95_ms": 43.617, "max_ms": 43.617, "n": 5}, "nodes_list_out_limit20": {"p50_ms": 4821.444, "p95_ms": 4835.256, "max_ms": 4835.256, "n": 5, "round_trips": 201}, "nodes_class_out_limit20": {"p50_ms": 2.05, "p95_ms": 3.891, "max_ms": 3.891, "n": 5, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 5050.231, "p95_ms": 5104.35, "max_ms": 5104.35, "n": 5, "round_trips": 201}, "nodes_list_in_unlimited": {"p50_ms": 4972.176, "p95_ms": 5101.09, "max_ms": 5101.09, "n": 5, "round_trips": 201}, "count_via_len_nodes": {"p50_ms": 4682.547, "p95_ms": 4927.673, "max_ms": 4927.673, "n": 5, "round_trips": 201}, "hub_save": {"p50_ms": 1151.34, "p95_ms": 1535.099, "max_ms": 1589.655, "n": 50}, "hub_connect": {"p50_ms": 162.314, "p95_ms": 558.513, "max_ms": 582.77, "n": 50, "round_trips": 5}, "concurrent_connect_32": {"wall_ms": 11665.166, "max_call_ms": 11657.035, "p50_call_ms": 6012.509}, "sizes_after_writes": {"hub_row_bytes": 2465290, "node_gin_bytes": 49577984, "node_total_bytes": 353665024, "edge_total_bytes": 171835392}} +{"degree": 1000, "jvspatial": "0.0.17", "git_sha": "8fc5138", "edge_ids_mode": "persist", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 2.25, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 26633, "node_gin_bytes": 196608, "node_total_bytes": 802816, "edge_total_bytes": 1892352}, "hub_get": {"p50_ms": 1.044, "p95_ms": 1.479, "max_ms": 2.615, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 40.666, "p95_ms": 64.47, "max_ms": 65.79, "n": 20, "round_trips": 3}, "nodes_class_out_limit20": {"p50_ms": 1.115, "p95_ms": 2.045, "max_ms": 2.392, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 45.583, "p95_ms": 72.069, "max_ms": 75.564, "n": 20, "round_trips": 3}, "nodes_list_in_unlimited": {"p50_ms": 39.259, "p95_ms": 70.509, "max_ms": 77.64, "n": 20, "round_trips": 3}, "count_via_len_nodes": {"p50_ms": 43.309, "p95_ms": 68.951, "max_ms": 71.525, "n": 20, "round_trips": 3}, "hub_save": {"p50_ms": 11.494, "p95_ms": 13.597, "max_ms": 24.191, "n": 50}, "hub_connect": {"p50_ms": 6.607, "p95_ms": 7.902, "max_ms": 10.426, "n": 50, "round_trips": 5}, "concurrent_connect_32": {"wall_ms": 133.819, "max_call_ms": 132.509, "p50_call_ms": 75.482}, "sizes_after_writes": {"hub_row_bytes": 28493, "node_gin_bytes": 3440640, "node_total_bytes": 7929856, "edge_total_bytes": 2015232}} +{"degree": 10000, "jvspatial": "0.0.17", "git_sha": "8fc5138", "edge_ids_mode": "persist", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 23.26, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 263664, "node_gin_bytes": 1638400, "node_total_bytes": 6668288, "edge_total_bytes": 17047552}, "hub_get": {"p50_ms": 3.945, "p95_ms": 5.141, "max_ms": 5.279, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 439.637, "p95_ms": 485.561, "max_ms": 488.057, "n": 20, "round_trips": 21}, "nodes_class_out_limit20": {"p50_ms": 1.137, "p95_ms": 1.708, "max_ms": 2.802, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 474.404, "p95_ms": 511.823, "max_ms": 513.261, "n": 20, "round_trips": 21}, "nodes_list_in_unlimited": {"p50_ms": 470.006, "p95_ms": 523.042, "max_ms": 571.798, "n": 20, "round_trips": 21}, "count_via_len_nodes": {"p50_ms": 472.215, "p95_ms": 510.991, "max_ms": 527.017, "n": 20, "round_trips": 21}, "hub_save": {"p50_ms": 110.905, "p95_ms": 137.193, "max_ms": 231.894, "n": 50}, "hub_connect": {"p50_ms": 23.168, "p95_ms": 52.342, "max_ms": 336.508, "n": 50, "round_trips": 5}, "concurrent_connect_32": {"wall_ms": 813.037, "max_call_ms": 812.608, "p50_call_ms": 430.787}, "sizes_after_writes": {"hub_row_bytes": 256983, "node_gin_bytes": 8962048, "node_total_bytes": 49291264, "edge_total_bytes": 19529728}} +{"degree": 100000, "jvspatial": "0.0.17", "git_sha": "8fc5138", "edge_ids_mode": "persist", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 262.07, "read_iterations": 5, "sizes_after_seed": {"hub_row_bytes": 2634003, "node_gin_bytes": 26787840, "node_total_bytes": 76308480, "edge_total_bytes": 171909120}, "hub_get": {"p50_ms": 32.39, "p95_ms": 47.729, "max_ms": 47.729, "n": 5}, "nodes_list_out_limit20": {"p50_ms": 4936.08, "p95_ms": 5184.805, "max_ms": 5184.805, "n": 5, "round_trips": 201}, "nodes_class_out_limit20": {"p50_ms": 1.344, "p95_ms": 3.033, "max_ms": 3.033, "n": 5, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 5317.014, "p95_ms": 5338.701, "max_ms": 5338.701, "n": 5, "round_trips": 201}, "nodes_list_in_unlimited": {"p50_ms": 5218.259, "p95_ms": 5379.298, "max_ms": 5379.298, "n": 5, "round_trips": 201}, "count_via_len_nodes": {"p50_ms": 4807.526, "p95_ms": 4881.467, "max_ms": 4881.467, "n": 5, "round_trips": 201}, "hub_save": {"p50_ms": 1160.649, "p95_ms": 1661.605, "max_ms": 1746.441, "n": 50}, "hub_connect": {"p50_ms": 160.187, "p95_ms": 581.505, "max_ms": 597.428, "n": 50, "round_trips": 5}, "concurrent_connect_32": {"wall_ms": 10276.447, "max_call_ms": 10244.43, "p50_call_ms": 4726.249}, "sizes_after_writes": {"hub_row_bytes": 2464513, "node_gin_bytes": 49561600, "node_total_bytes": 363249664, "edge_total_bytes": 171966464}} diff --git a/docs/bench/2026-09-hub-node-phase1.jsonl b/docs/bench/2026-09-hub-node-phase1.jsonl new file mode 100644 index 0000000..cd9f928 --- /dev/null +++ b/docs/bench/2026-09-hub-node-phase1.jsonl @@ -0,0 +1,3 @@ +{"degree": 1000, "jvspatial": "0.0.17", "git_sha": "23af4c0", "edge_ids_mode": "derive", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 2.31, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 134, "node_gin_bytes": 106496, "node_total_bytes": 516096, "edge_total_bytes": 1875968}, "hub_get": {"p50_ms": 0.541, "p95_ms": 0.779, "max_ms": 3.108, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 42.08, "p95_ms": 64.546, "max_ms": 66.334, "n": 20, "round_trips": 3}, "nodes_class_out_limit20": {"p50_ms": 1.343, "p95_ms": 2.69, "max_ms": 4.956, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 42.506, "p95_ms": 75.457, "max_ms": 84.249, "n": 20, "round_trips": 3}, "nodes_list_in_unlimited": {"p50_ms": 45.372, "p95_ms": 77.519, "max_ms": 79.347, "n": 20, "round_trips": 3}, "count_via_len_nodes": {"p50_ms": 43.583, "p95_ms": 73.171, "max_ms": 74.113, "n": 20, "round_trips": 3}, "hub_save": {"p50_ms": 0.796, "p95_ms": 1.354, "max_ms": 5.306, "n": 50}, "hub_connect": {"p50_ms": 1.766, "p95_ms": 2.52, "max_ms": 3.946, "n": 50, "round_trips": 3}, "concurrent_connect_32": {"wall_ms": 18.32, "max_call_ms": 17.782, "p50_call_ms": 16.352}, "sizes_after_writes": {"hub_row_bytes": 136, "node_gin_bytes": 114688, "node_total_bytes": 524288, "edge_total_bytes": 1974272}} +{"degree": 10000, "jvspatial": "0.0.17", "git_sha": "23af4c0", "edge_ids_mode": "derive", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 24.05, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 134, "node_gin_bytes": 1081344, "node_total_bytes": 4448256, "edge_total_bytes": 17334272}, "hub_get": {"p50_ms": 0.78, "p95_ms": 0.98, "max_ms": 1.648, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 459.237, "p95_ms": 504.987, "max_ms": 507.852, "n": 20, "round_trips": 21}, "nodes_class_out_limit20": {"p50_ms": 1.247, "p95_ms": 1.717, "max_ms": 3.05, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 474.439, "p95_ms": 525.195, "max_ms": 532.879, "n": 20, "round_trips": 21}, "nodes_list_in_unlimited": {"p50_ms": 479.165, "p95_ms": 515.222, "max_ms": 531.191, "n": 20, "round_trips": 21}, "count_via_len_nodes": {"p50_ms": 466.91, "p95_ms": 501.051, "max_ms": 508.005, "n": 20, "round_trips": 21}, "hub_save": {"p50_ms": 0.78, "p95_ms": 3.48, "max_ms": 4.774, "n": 50}, "hub_connect": {"p50_ms": 1.667, "p95_ms": 2.622, "max_ms": 4.505, "n": 50, "round_trips": 3}, "concurrent_connect_32": {"wall_ms": 29.041, "max_call_ms": 28.707, "p50_call_ms": 26.515}, "sizes_after_writes": {"hub_row_bytes": 136, "node_gin_bytes": 1638400, "node_total_bytes": 5013504, "edge_total_bytes": 17416192}} +{"degree": 100000, "jvspatial": "0.0.17", "git_sha": "23af4c0", "edge_ids_mode": "derive", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 255.56, "read_iterations": 5, "sizes_after_seed": {"hub_row_bytes": 134, "node_gin_bytes": 14016512, "node_total_bytes": 46727168, "edge_total_bytes": 169852928}, "hub_get": {"p50_ms": 0.567, "p95_ms": 3.692, "max_ms": 3.692, "n": 5}, "nodes_list_out_limit20": {"p50_ms": 4850.617, "p95_ms": 5425.813, "max_ms": 5425.813, "n": 5, "round_trips": 201}, "nodes_class_out_limit20": {"p50_ms": 1.366, "p95_ms": 3.67, "max_ms": 3.67, "n": 5, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 5038.662, "p95_ms": 5254.32, "max_ms": 5254.32, "n": 5, "round_trips": 201}, "nodes_list_in_unlimited": {"p50_ms": 5047.58, "p95_ms": 5195.632, "max_ms": 5195.632, "n": 5, "round_trips": 201}, "count_via_len_nodes": {"p50_ms": 4642.871, "p95_ms": 4915.509, "max_ms": 4915.509, "n": 5, "round_trips": 201}, "hub_save": {"p50_ms": 0.729, "p95_ms": 1.094, "max_ms": 5.637, "n": 50}, "hub_connect": {"p50_ms": 1.917, "p95_ms": 2.982, "max_ms": 4.07, "n": 50, "round_trips": 3}, "concurrent_connect_32": {"wall_ms": 19.11, "max_call_ms": 18.287, "p50_call_ms": 17.086}, "sizes_after_writes": {"hub_row_bytes": 136, "node_gin_bytes": 14016512, "node_total_bytes": 46735360, "edge_total_bytes": 169967616}} diff --git a/docs/md/entity-reference.md b/docs/md/entity-reference.md index 6527a05..16cebd7 100644 --- a/docs/md/entity-reference.md +++ b/docs/md/entity-reference.md @@ -51,7 +51,8 @@ class Node(Object): async def connect(other: "Node", edge: Type["Edge"] = Edge, direction: str = "out", **kwargs) -> "Edge" - async def edges(direction: str = "") -> List["Edge"] + async def edges(direction: str = "", limit: Optional[int] = None) -> List["Edge"] + async def connection_count() -> int # degree; one COUNT in derive mode async def nodes(direction: str = "both", node: Optional[...] = None, edge: Optional[...] = None, **kwargs) -> List["Node"] async def node(direction: str = "out", node: Optional[...] = None, @@ -63,6 +64,11 @@ class Node(Object): async def count(cls, query: Optional[dict] = None, **kwargs) -> int # Inherited from Object ``` +**Adjacency:** on Postgres, MongoDB and SQLite (derive mode) the edge +collection is the source of truth — `edge_ids` stays empty in memory and node +rows carry no `edges` array. Use `edges()`, `connection_count()` and `nodes()` +rather than reading `edge_ids`. See [graph-context.md](graph-context.md). + **Key Methods:** - **`nodes()`**: Returns a list of connected nodes with filtering options diff --git a/docs/md/environment-keys-reference.md b/docs/md/environment-keys-reference.md index c579ca8..f6ff44d 100644 --- a/docs/md/environment-keys-reference.md +++ b/docs/md/environment-keys-reference.md @@ -50,6 +50,7 @@ For full examples and default values, see: - `JVSPATIAL_DYNAMODB_REGION` - DynamoDB region. - `JVSPATIAL_DYNAMODB_ENDPOINT_URL` - DynamoDB endpoint (e.g. LocalStack). - `JVSPATIAL_DYNAMODB_WAIT_FOR_INDEX` - Wait-for-index toggle. +- `JVSPATIAL_NODE_EDGE_IDS` - Node adjacency mode: `derive` (edge collection only; Postgres / MongoDB / SQLite default) or `persist` (node rows also store `edges`; JSON / DynamoDB default). An `edge_ids_mode` set on the adapter instance wins. ### Auth and rate limit - `JVSPATIAL_AUTH_ENABLED` - Enables auth. diff --git a/docs/md/graph-context.md b/docs/md/graph-context.md index bcb7148..d87cd05 100644 --- a/docs/md/graph-context.md +++ b/docs/md/graph-context.md @@ -673,12 +673,20 @@ edge = await ctx.create_edge(Friendship, source=user1, target=user2) ### Batch Operations `GraphContext.save_batch()` routes **node** entities through `save()` so -edge-merge semantics apply (concurrent `atomic_add_edge_id` updates are not -clobbered). Non-node entities still use the adapter bulk path when available. - -`get_batch()` chunks id lists at 500 per round trip. `atomic_add_edge_id` / -`atomic_remove_edge_id` use native `find_one_and_update` on MongoDB and -Postgres; other backends fall back to read-modify-write. +edge-merge semantics apply in persist mode (concurrent `atomic_add_edge_id` +updates are not clobbered). Non-node entities still use the adapter bulk path +when available. + +`get_batch()` chunks id lists at 500 per round trip. + +**Node adjacency mode.** `ctx.persists_edge_ids()` reports whether node rows +carry an `edges` array. Postgres, MongoDB and SQLite default to *derive*: +adjacency is read from the indexed edge collection, node saves never merge or +lock an edge list, and `atomic_add_edge_id` / `atomic_remove_edge_id` are +no-ops. JsonDB and DynamoDB default to *persist*, where those helpers use +native `find_one_and_update` (MongoDB, Postgres) or read-modify-write. Override +per adapter (`db.edge_ids_mode = "persist"`) or process-wide +(`JVSPATIAL_NODE_EDGE_IDS=persist|derive`). **Fast deserialize (opt-in):** set `JVSPATIAL_FAST_DESERIALIZE=true` to hydrate trusted DB rows via `model_construct` instead of full Pydantic diff --git a/docs/md/postgres-guide.md b/docs/md/postgres-guide.md index 06b7fb5..95ed1bc 100644 --- a/docs/md/postgres-guide.md +++ b/docs/md/postgres-guide.md @@ -159,6 +159,33 @@ await db.enable_vector_column("doc", "embedding", dim=1536) See [vector-store.md](vector-store.md) for details. +### Node adjacency lives in the `edge` table + +`PostgresDB.edge_ids_mode = "derive"`: node rows carry no `edges` array. +Adjacency is read from the `edge` table through the `source` / `target` +functional indexes that `Edge.get_indexes()` declares, so `connect()`, +`disconnect()` and `save()` cost the same on a node with ten edges as on one +with a hundred thousand, and concurrent writers never queue on a hub's row +lock. Keep index auto-creation on (`JVSPATIAL_AUTO_CREATE_INDEXES`, default on +outside serverless) or run `ensure_indexes(Edge)` at deploy time — without +those indexes every adjacency read scans the `edge` table. + +Rows written by 0.0.17 and earlier still carry the array. Reads ignore it and +the next save of each node drops it; to reclaim the space in one pass: + +```bash +jvspatial migrate strip-node-edges --dsn "$JVSPATIAL_POSTGRES_DSN" # dry run: count rows +jvspatial migrate strip-node-edges --dsn "$JVSPATIAL_POSTGRES_DSN" --apply # strip, batched +``` + +It is idempotent and safe while the application serves traffic (derive mode +never writes the array back). Tables with `FORCE ROW LEVEL SECURITY` must be +migrated by a role that bypasses RLS. Afterwards run `VACUUM (ANALYZE) node` +and `REINDEX INDEX CONCURRENTLY node_data_gin` to return the space. + +To keep the pre-0.0.18 behaviour set `JVSPATIAL_NODE_EDGE_IDS=persist` (or +`db.edge_ids_mode = "persist"` on the adapter). + ## Pool tuning `PostgresDB` runs on a single `asyncpg.Pool` per instance, created lazily on @@ -256,6 +283,7 @@ precedence): | `JVSPATIAL_POSTGRES_MIN_POOL_SIZE` | Pool min size override | | `JVSPATIAL_POSTGRES_MAX_POOL_SIZE` | Pool max size override | | `JVSPATIAL_POSTGRES_POOLER_MODE` | `"session"` (default) or `"transaction"` | +| `JVSPATIAL_NODE_EDGE_IDS` | `"derive"` (Postgres default) or `"persist"` | ## Operational tips diff --git a/jvspatial/api/endpoints/graph_visualization.py b/jvspatial/api/endpoints/graph_visualization.py index 21a934b..773486e 100644 --- a/jvspatial/api/endpoints/graph_visualization.py +++ b/jvspatial/api/endpoints/graph_visualization.py @@ -45,7 +45,11 @@ async def graph_expand( cursor: int = Query( # noqa: B008 default=0, ge=0, - description="Offset into the node's edge-id list", + description="Offset into the node's id-sorted incident edges", + ), + after: str = Query( # noqa: B008 + default="", + description="Keyset cursor: pagination.next_after from the previous page (overrides cursor)", ), detail_level: str = Query( # noqa: B008 default="full", @@ -60,6 +64,7 @@ async def graph_expand( direction=direction, limit=limit, cursor=cursor, + after=after or None, detail_level=detail_level, ) except Exception as e: diff --git a/jvspatial/cli.py b/jvspatial/cli.py index 44badf6..58290be 100644 --- a/jvspatial/cli.py +++ b/jvspatial/cli.py @@ -14,6 +14,10 @@ jvspatial migrate --collection node --entity User --dry-run jvspatial migrate --collection node # all entities in collection jvspatial migrate --collection node --apply # actually persist changes + + # derive-mode adjacency: drop the legacy ``edges`` array from node rows + jvspatial migrate strip-node-edges --dsn postgresql://... --apply + jvspatial migrate strip-node-edges --dsn mongodb://... --db-name app --apply """ from __future__ import annotations @@ -220,6 +224,71 @@ async def _run_migrate(args: argparse.Namespace) -> int: return 0 if failed == 0 else 1 +# ---- migrate strip-node-edges ---------------------------------------------- + + +async def _run_strip_node_edges(args: argparse.Namespace) -> int: + """Strip the legacy ``edges`` array from node rows (derive-mode migration). + + Returns process exit code (0 on success, >0 on failure). + """ + from jvspatial.db.factory import create_database + from jvspatial.db.manager import get_database_manager + + dsn = args.dsn or "" + is_postgres = dsn.startswith(("postgres://", "postgresql://")) + if is_postgres: + db = create_database("postgres", dsn=dsn, schema_name=args.schema) + elif dsn.startswith(("mongodb://", "mongodb+srv://")): + db = create_database("mongodb", uri=dsn, db_name=args.db_name) + elif dsn: + logger.error("Unsupported --dsn scheme: expected postgresql:// or mongodb://") + return 2 + else: + try: + db = get_database_manager().get_prime_database() + except Exception: + logger.error("No database configured; pass --dsn.") + return 2 + + strip = getattr(db, "strip_node_edges", None) + if not callable(strip): + logger.error( + "%s keeps adjacency on node rows (edge_ids_mode=persist); " + "nothing to strip.", + type(db).__name__, + ) + return 2 + + collection = args.collection or "node" + try: + count = await strip(collection, batch_size=args.batch, dry_run=args.dry_run) + finally: + close = getattr(db, "close", None) + if callable(close): + await close() + + if args.dry_run: + logger.info( + "[dry-run] %d %s record(s) carry a legacy 'edges' array; " + "re-run with --apply to strip them.", + count, + collection, + ) + return 0 + logger.info("stripped 'edges' from %d %s record(s)", count, collection) + if is_postgres and count: + logger.info( + "Reclaim space when convenient: VACUUM (ANALYZE) %s.%s; " + "REINDEX INDEX CONCURRENTLY %s.%s_data_gin;", + args.schema, + collection, + args.schema, + collection, + ) + return 0 + + # ---- entry point ----------------------------------------------------------- @@ -237,10 +306,42 @@ def build_parser() -> argparse.ArgumentParser: "migrate", help="Apply schema migrations to existing records", ) + mig.add_argument( + "action", + nargs="?", + choices=["strip-node-edges"], + help=( + "Optional data migration. strip-node-edges: remove the legacy " + "'edges' array from node rows (Postgres / MongoDB derive mode)." + ), + ) mig.add_argument( "--collection", - required=True, - help='Collection to scan ("node" / "edge" / "object" / "walker").', + help=( + 'Collection to scan ("node" / "edge" / "object" / "walker"). ' + 'Required for schema migrations; strip-node-edges defaults to "node".' + ), + ) + mig.add_argument( + "--dsn", + help="strip-node-edges: postgresql:// DSN or mongodb:// URI " + "(default: the configured prime database).", + ) + mig.add_argument( + "--schema", + default="public", + help="strip-node-edges: Postgres schema holding the collection tables.", + ) + mig.add_argument( + "--db-name", + default="jvdb", + help="strip-node-edges: MongoDB database name.", + ) + mig.add_argument( + "--batch", + type=int, + default=5000, + help="strip-node-edges: rows per UPDATE batch.", ) mig.add_argument( "--entity", @@ -283,6 +384,10 @@ def main(argv: Optional[List[str]] = None) -> int: _configure_logging(args.verbose) if args.cmd == "migrate": + if args.action == "strip-node-edges": + return asyncio.run(_run_strip_node_edges(args)) + if not args.collection: + parser.error("migrate: --collection is required") return asyncio.run(_run_migrate(args)) parser.error(f"Unknown command: {args.cmd!r}") diff --git a/jvspatial/core/context.py b/jvspatial/core/context.py index b201043..506ca86 100644 --- a/jvspatial/core/context.py +++ b/jvspatial/core/context.py @@ -23,7 +23,11 @@ cast, ) -from jvspatial.db.database import Database, resolve_sort_value +from jvspatial.db.database import ( + Database, + resolve_edge_ids_mode, + resolve_sort_value, +) from jvspatial.db.factory import create_database, get_current_database from jvspatial.db.manager import get_database_manager @@ -426,6 +430,18 @@ async def _node_edge_write_guard(self, node_id: str): async with lock: yield + def persists_edge_ids(self) -> bool: + """Whether node documents persist their incident edge ids (``edges``). + + ``False`` when the active backend derives adjacency from the edge + collection (``edge_ids_mode="derive"`` — the Postgres, MongoDB and + SQLite default): ``connect()`` / ``save()`` then never touch the node + rows and ``Node.edge_ids`` stays an unpopulated in-memory list. ``True`` + for ``"persist"`` (JsonDB, DynamoDB). See + :func:`jvspatial.db.database.resolve_edge_ids_mode`. + """ + return resolve_edge_ids_mode(self.database) == "persist" + @property def database(self) -> Database: """Get the database instance, initializing if needed.""" @@ -892,6 +908,15 @@ async def save( hasattr(entity, "type_code") and getattr(entity, "type_code", "") == "n" ) + if is_node and not self.persists_edge_ids(): + # Derive mode: adjacency lives in the edge collection. Never write + # an ``edges`` array (a legacy in-memory copy must not re-persist + # it), and there is no array merge to serialise per node. + record.pop("edges", None) + await db.save(collection, record) + await self._add_to_cache(entity.id, entity) + return entity + async def _merge_edges_and_write() -> None: # Merge node edge lists with the DB so full-document saves do not clobber # edge IDs added concurrently via atomic_add_edge_id (or another writer). @@ -952,8 +977,14 @@ async def delete(self, entity, cascade: bool = False) -> None: if isinstance(entity, Node): # Check if this is a recursive call from Node.delete() by checking - # if cascade=False and the node has no edges (cleaned up by Node.delete()) - if not cascade and len(entity.edge_ids) == 0: + # if cascade=False and the node has no edges (cleaned up by Node.delete()). + # In derive mode ``edge_ids`` is never populated, so ask the edge + # collection instead. + if not cascade and ( + len(entity.edge_ids) == 0 + if self.persists_edge_ids() + else await entity.connection_count() == 0 + ): # Node.delete() has cleaned up edges, just delete the entity collection = self._get_collection_name("n") await self.database.delete(collection, entity.id) @@ -1072,6 +1103,7 @@ async def expand_node( direction: str = "both", limit: int = 50, cursor: int = 0, + after: Optional[str] = None, detail_level: str = "full", ) -> Dict[str, Any]: """Return a page of incident edges and neighbor summaries for progressive UIs. @@ -1088,6 +1120,7 @@ async def expand_node( direction=direction, limit=limit, cursor=cursor, + after=after, detail_level=detail_level, # type: ignore[arg-type] ) @@ -1228,8 +1261,11 @@ async def atomic_add_edge_id(self, node_id: str, edge_id: str) -> bool: Falls back to read-modify-write when the database does not support atomic updates or when the document is not found. - Returns True on success, False on failure. + Returns True on success, False on failure. A no-op returning True in + derive mode (:meth:`persists_edge_ids` is ``False``). """ + if not self.persists_edge_ids(): + return True db = self.database if self._is_mongodb(db) or self._is_postgres(db): try: @@ -1271,8 +1307,11 @@ async def atomic_remove_edge_id(self, node_id: str, edge_id: str) -> bool: Falls back to read-modify-write when the database does not support atomic updates or when the document is not found. - Returns True on success, False on failure. + Returns True on success, False on failure. A no-op returning True in + derive mode (:meth:`persists_edge_ids` is ``False``). """ + if not self.persists_edge_ids(): + return True db = self.database if self._is_mongodb(db) or self._is_postgres(db): try: @@ -1839,9 +1878,15 @@ async def _deserialize_entity( # entity_type_code already computed above + # Legacy rows may still carry an ``edges`` array; in derive mode it + # is stale (no longer maintained) and is ignored. + node_edge_ids: List[str] = [] + if entity_type_code == "n" and self.persists_edge_ids(): + node_edge_ids = data.get("edges", []) + if self._fast_deserialize_enabled(): if entity_type_code == "n": - edge_ids = data.get("edges", []) + edge_ids = node_edge_ids context_data.pop("edge_ids", None) context_data.pop("id", None) context_data.pop("type_code", None) @@ -1873,7 +1918,7 @@ async def _deserialize_entity( # Handle Node-specific logic # Extract edge_ids from data (stored as "edges" at top level) # Edges are included in database exports but excluded from default exports - edge_ids = data.get("edges", []) + edge_ids = node_edge_ids # Remove edge_ids, id, and type_code from context_data as they're handled separately context_data.pop("edge_ids", None) diff --git a/jvspatial/core/entities/node.py b/jvspatial/core/entities/node.py index 8576ae8..75c37b6 100644 --- a/jvspatial/core/entities/node.py +++ b/jvspatial/core/entities/node.py @@ -196,8 +196,14 @@ async def connect( matching_edge = existing_edge break + # Derive mode: the edge row is the adjacency — node rows and the + # in-memory ``edge_ids`` lists are left untouched (no hub rewrite). + persist = context.persists_edge_ids() + # If an existing edge is found, return it instead of creating a duplicate if matching_edge: + if not persist: + return matching_edge # Ensure edge IDs are in both nodes' edge_ids lists (in case they're missing) if matching_edge.id not in self.edge_ids: await context.atomic_add_edge_id(self.id, matching_edge.id) @@ -226,6 +232,9 @@ async def connect( return retry_edges[0] raise + if not persist: + return connection + # Atomically update both nodes' edge_ids await context.atomic_add_edge_id(self.id, connection.id) if connection.id not in self.edge_ids: @@ -237,21 +246,49 @@ async def connect( return connection - async def edges(self: "Node", direction: str = "") -> List["Edge"]: + async def edges( + self: "Node", direction: str = "", limit: Optional[int] = None + ) -> List["Edge"]: """Get edges connected to this node. + In derive mode (Postgres / MongoDB / SQLite default) the edge + collection is queried by ``source`` / ``target`` — index-backed, one + round trip. In persist mode the node's ``edge_ids`` are fetched. + Args: direction: Filter edges by direction ('in', 'out', 'both') + limit: Maximum number of edges to return (default: all) Returns: List of edge instances """ - if not self.edge_ids: - return [] + context = await self.get_context() - from ..context import get_default_context + if not context.persists_edge_ids(): + query: Dict[str, Any] + if direction == "out": + query = {"source": self.id} + elif direction == "in": + query = {"target": self.id} + else: + query = {"$or": [{"source": self.id}, {"target": self.id}]} + rows = await context.database.find("edge", query, limit=limit) + if limit is None and len(rows) > 10_000: + logger.debug( + "Node.edges(%s) loaded %d edges; pass limit= or use " + "connection_count() for hub nodes", + self.id, + len(rows), + ) + derived: List["Edge"] = [] + for row in rows: + edge_obj = await context._deserialize_entity(Edge, row) + if edge_obj: + derived.append(edge_obj) + return derived - context = get_default_context() + if not self.edge_ids: + return [] # Use batch query for efficiency (N+1 -> 1 query) edge_results = await context.database.find( @@ -271,11 +308,24 @@ async def edges(self: "Node", direction: str = "") -> List["Edge"]: # Filter by direction if specified if direction == "out": - return [e for e in edges if e.source == self.id] + edges = [e for e in edges if e.source == self.id] elif direction == "in": - return [e for e in edges if e.target == self.id] - else: - return edges + edges = [e for e in edges if e.target == self.id] + return edges if limit is None else edges[:limit] + + async def _incident_edges(self: "Node", context: "GraphContext") -> List["Edge"]: + """Every edge touching this node, either endpoint (used by cascade delete).""" + if context.persists_edge_ids(): + found: List["Edge"] = [] + for edge_id in self.edge_ids: + try: + edge = await Edge.get(edge_id) + if edge: + found.append(edge) + except Exception: + continue + return found + return await self.edges() async def nodes( self, @@ -1132,16 +1182,18 @@ async def disconnect( try: context = await self.get_context() edges = await context.find_edges_between(self.id, other.id, edge_type) + persist = context.persists_edge_ids() for found_edge in edges: - # Atomically remove edge_id from both nodes, then delete the edge - await context.atomic_remove_edge_id(self.id, found_edge.id) - if found_edge.id in self.edge_ids: - self.edge_ids.remove(found_edge.id) + if persist: + # Atomically remove edge_id from both nodes, then delete the edge + await context.atomic_remove_edge_id(self.id, found_edge.id) + if found_edge.id in self.edge_ids: + self.edge_ids.remove(found_edge.id) - await context.atomic_remove_edge_id(other.id, found_edge.id) - if found_edge.id in other.edge_ids: - other.edge_ids.remove(found_edge.id) + await context.atomic_remove_edge_id(other.id, found_edge.id) + if found_edge.id in other.edge_ids: + other.edge_ids.remove(found_edge.id) # Delete the edge document (context.delete already handles # edge_ids cleanup, but we already did it atomically above, @@ -1182,10 +1234,18 @@ async def is_connected_to( async def connection_count(self) -> int: """Get the number of connections (edges) for this node. + In derive mode this is one ``COUNT`` over the edge collection's + ``source`` / ``target`` indexes — the canonical degree query. + Returns: Number of connected edges """ - return len(self.edge_ids) + context = await self.get_context() + if context.persists_edge_ids(): + return len(self.edge_ids) + return await context.database.count( + "edge", {"$or": [{"source": self.id}, {"target": self.id}]} + ) async def delete(self: "Node", cascade: bool = True) -> None: """Delete this node and cascade deletion of all related edges and dependent nodes. @@ -1237,8 +1297,9 @@ async def delete(self: "Node", cascade: bool = True) -> None: except Exception: continue - # Also check edges from edge_ids where this node is the target - for edge_id in self.edge_ids: + # Persist mode: also check edges from edge_ids where this node is the + # target (the edge query above already covers derive mode). + for edge_id in self.edge_ids if context.persists_edge_ids() else []: try: edge = await Edge.get(edge_id) # Only include edges where this node is the target (incoming) @@ -1307,14 +1368,10 @@ async def delete(self: "Node", cascade: bool = True) -> None: if not node: continue # Only follow outgoing edges (where this node is the source) - for edge_id in node.edge_ids: # type: ignore[attr-defined] - try: - edge = await Edge.get(edge_id) - if edge and edge.source == node_id: - # Only add nodes reachable via outgoing edges - nodes_to_check.add(edge.target) - except Exception: - continue + for edge in await node._incident_edges(context): # type: ignore[attr-defined] + if edge.source == node_id: + # Only add nodes reachable via outgoing edges + nodes_to_check.add(edge.target) except Exception: continue @@ -1329,14 +1386,9 @@ async def delete(self: "Node", cascade: bool = True) -> None: continue # Get all edges of the candidate node - candidate_edges = [] - for edge_id in candidate_node.edge_ids: # type: ignore[attr-defined] - try: - edge = await Edge.get(edge_id) - if edge: - candidate_edges.append(edge) - except Exception: - continue + candidate_edges = await candidate_node._incident_edges( # type: ignore[attr-defined] + context + ) # If the node has no edges, it's orphaned and should be deleted if not candidate_edges: @@ -1375,14 +1427,9 @@ async def is_node_only_connected_to_deletion_set( if not node: return True - node_edges = [] - for edge_id in node.edge_ids: # type: ignore[attr-defined] - try: - edge = await Edge.get(edge_id) - if edge: - node_edges.append(edge) - except Exception: - continue + node_edges = await node._incident_edges( # type: ignore[attr-defined] + context + ) # If no edges, it's orphaned and should be deleted if not node_edges: @@ -1488,12 +1535,14 @@ async def is_node_only_connected_to_deletion_set( # Continue even if dependent node deletion fails continue - # Clear edge_ids before final deletion to avoid recursion in context.delete() # All edges have already been deleted from the database self.edge_ids = [] - # Finally, delete this node itself (no cascade needed, we've already handled it) - await context.delete(self, cascade=False) + # Finally, delete this node itself. Direct delete rather than + # ``context.delete`` — its "no edges left?" check would cost a COUNT in + # derive mode, and the edges are already gone. + await context.database.delete(context._get_collection_name("n"), self.id) + await context._remove_from_cache(self.id) @classmethod async def create_and_connect( diff --git a/jvspatial/core/entities/root.py b/jvspatial/core/entities/root.py index 33de5b8..25b583a 100644 --- a/jvspatial/core/entities/root.py +++ b/jvspatial/core/entities/root.py @@ -47,8 +47,11 @@ async def get(cls: Type["Root"], id: Optional[str] = None) -> "Root": # type: i if not isinstance(context_data, dict): context_data = {} - # Handle edge_ids from database format (stored as "edges" at top level) - edge_ids = node_data.get("edges", []) + # Handle edge_ids from database format (stored as "edges" at top + # level); ignored in derive mode, where the array is not maintained. + edge_ids = ( + node_data.get("edges", []) if context.persists_edge_ids() else [] + ) if not isinstance(edge_ids, list): edge_ids = [] diff --git a/jvspatial/core/graph_expansion.py b/jvspatial/core/graph_expansion.py index 0e7e3f1..fd9cf0c 100644 --- a/jvspatial/core/graph_expansion.py +++ b/jvspatial/core/graph_expansion.py @@ -2,6 +2,7 @@ from __future__ import annotations +import asyncio from collections import deque from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple @@ -27,6 +28,30 @@ def _coerce_edge_id_list(value: Any) -> List[str]: return [] +def _incident_query(node_id: str) -> Dict[str, Any]: + """Edge-collection query for every edge touching ``node_id``.""" + return {"$or": [{"source": node_id}, {"target": node_id}]} + + +async def _node_degrees( + context: GraphContext, records: Dict[str, Dict[str, Any]] +) -> Dict[str, int]: + """Degree per node record. + + Persist mode reads the stored ``edges`` array; derive mode counts the + edge collection (the array, if a legacy row still has one, is stale). + """ + if context.persists_edge_ids(): + return { + nid: len(_coerce_edge_id_list(rec.get("edges"))) + for nid, rec in records.items() + } + ids = list(records) + db = context.database + counts = await asyncio.gather(*(db.count("edge", _incident_query(n)) for n in ids)) + return dict(zip(ids, counts)) + + def _edge_matches_direction( edge_doc: Dict[str, Any], node_id: str, direction: str ) -> bool: @@ -83,19 +108,23 @@ async def expand_node( direction: str = "both", limit: int = 50, cursor: int = 0, + after: Optional[str] = None, detail_level: DetailLevel = "full", ) -> Dict[str, Any]: """Load the center node and a page of incident edges plus neighbor summaries. - Uses the node's persisted ``edges`` list and batch ``get`` calls — O(limit), - not O(|E|). + Pages come from the edge collection sorted by edge id (index-backed on + SQL backends), so a page costs O(limit) rows regardless of the node's + degree when paging by keyset. Args: context: Active graph context node_id: Node to expand around direction: ``both`` (default), ``out``, or ``in`` (for non-bidirectional edges) limit: Max edges in this page (capped at 500) - cursor: Offset into the sorted edge-id list + cursor: Offset into the id-sorted incident edges (costs O(cursor + limit)) + after: Keyset cursor — an edge id from a previous page's + ``pagination.next_after``; takes precedence over ``cursor`` detail_level: ``summary`` (no context) or ``full`` (trimmed context on all nodes/edges) Returns: @@ -113,6 +142,7 @@ async def expand_node( "pagination": { "cursor": cursor, "next_cursor": None, + "next_after": None, "has_more": False, "total_edge_count": 0, "returned_edges": 0, @@ -120,28 +150,41 @@ async def expand_node( "found": False, } - all_edge_ids = sorted(_coerce_edge_id_list(center_raw.get("edges"))) - total = len(all_edge_ids) - page_ids = all_edge_ids[cursor : cursor + limit] - - edge_docs: List[Dict[str, Any]] = [] - for eid in page_ids: - doc = await db.get("edge", eid) - if doc and _edge_matches_direction(doc, node_id, direction): - edge_docs.append(doc) - - neighbor_ids: List[str] = [] - for doc in edge_docs: - other = _other_endpoint(doc, node_id) - if other and other != node_id: - neighbor_ids.append(other) - - neighbor_ids_unique = sorted(set(neighbor_ids)) - neighbor_records: Dict[str, Dict[str, Any]] = {} - for nid in neighbor_ids_unique: - nraw = await db.get("node", nid) - if nraw: - neighbor_records[nid] = nraw + incident = _incident_query(node_id) + total = await db.count("edge", incident) + page: List[Dict[str, Any]] = [] + if limit: + if after: + rows = await db.find( + "edge", + {"$and": [incident, {"id": {"$gt": after}}]}, + sort=[("id", 1)], + limit=limit + 1, + ) + has_more = len(rows) > limit + page = rows[:limit] + else: + rows = await db.find( + "edge", incident, sort=[("id", 1)], limit=cursor + limit + ) + page = rows[cursor:] + has_more = cursor + len(page) < total + else: + has_more = cursor < total + + edge_docs = [d for d in page if _edge_matches_direction(d, node_id, direction)] + + neighbor_ids_unique = sorted( + { + other + for doc in edge_docs + if (other := _other_endpoint(doc, node_id)) and other != node_id + } + ) + neighbor_records = ( + await db.find_many("node", neighbor_ids_unique) if neighbor_ids_unique else {} + ) + degrees = await _node_degrees(context, neighbor_records) nodes_out: List[Dict[str, Any]] = [ node_record_to_payload( @@ -152,13 +195,11 @@ async def expand_node( ] for nid in neighbor_ids_unique: if nid in neighbor_records: - nedges = neighbor_records[nid].get("edges") or [] - deg = len(nedges) if isinstance(nedges, list) else 0 nodes_out.append( node_record_to_payload( neighbor_records[nid], detail_level=detail_level, - degree=deg, + degree=degrees.get(nid, 0), ) ) else: @@ -182,8 +223,8 @@ async def expand_node( ) for doc in edge_docs ] - next_cursor = cursor + len(page_ids) if cursor + len(page_ids) < total else None - has_more = next_cursor is not None + next_cursor = cursor + len(page) if has_more and not after else None + next_after = str(page[-1].get("id")) if has_more and page else None return { "center_id": node_id, @@ -192,6 +233,7 @@ async def expand_node( "pagination": { "cursor": cursor, "next_cursor": next_cursor, + "next_after": next_after, "has_more": has_more, "total_edge_count": total, "returned_edges": len(edges_out), @@ -213,6 +255,7 @@ async def subgraph_bfs( Stops when ``max_depth`` or ``max_nodes`` would be exceeded. Each node follows at most ``max_edges_per_node`` incident edges (sorted by edge id). + Incident edges come from one edge-collection query per expanded node. Args: context: Active graph context @@ -253,9 +296,9 @@ async def subgraph_bfs( if d >= max_depth: continue - all_eids = _coerce_edge_id_list((raw or {}).get("edges")) eid_docs: List[Tuple[str, Optional[Dict[str, Any]]]] = [ - (eid, await db.get("edge", eid)) for eid in all_eids + (str(doc.get("id")), doc) + for doc in await db.find("edge", _incident_query(vid)) ] eid_docs.sort(key=lambda t: _bfs_spine_edge_sort_key(t[0], t[1], vid)) selected = eid_docs[:max_edges_per_node] @@ -272,13 +315,15 @@ async def subgraph_bfs( if other not in seen: q.append((other, d + 1)) + degrees = await _node_degrees(context, nodes_by_id) node_payloads: List[Dict[str, Any]] = [] for nid in sorted(seen): rec = nodes_by_id.get(nid) if rec: - ecount = len(_coerce_edge_id_list(rec.get("edges"))) node_payloads.append( - node_record_to_payload(rec, detail_level=detail_level, degree=ecount) + node_record_to_payload( + rec, detail_level=detail_level, degree=degrees.get(nid, 0) + ) ) else: ent = entity_type_from_node_id(nid) diff --git a/jvspatial/db/README.md b/jvspatial/db/README.md index 2f80a94..d7fcffd 100644 --- a/jvspatial/db/README.md +++ b/jvspatial/db/README.md @@ -59,9 +59,10 @@ db/ ## Invariants - **`Database.supports_transactions` is a capability flag.** Branch on it; do not sniff adapter class. (`database.py:84`) +- **`Database.edge_ids_mode` says where node adjacency lives.** `"derive"` (Postgres, MongoDB, SQLite): the edge collection only; node rows carry no `edges`. `"persist"` (JSON, DynamoDB): node rows also store `edges`. Read the effective value with `resolve_edge_ids_mode(db)` — instance attribute → `JVSPATIAL_NODE_EDGE_IDS` → class default. `strip_node_edges()` (Postgres, MongoDB) removes legacy arrays; CLI `jvspatial migrate strip-node-edges`. - **`find_many` and `bulk_save` are public and benefit from native overrides.** Defaults exist but are slow. (`database.py:176+`) - **`find_one_and_update` / `find_one_and_delete` are NOT atomic by default.** MongoDB and Postgres override with native atomic versions (`FOR UPDATE` on Postgres). -- **Postgres-only helpers** (via `getattr`, not on the ABC): `traverse`, `find_connected_nodes`, `save_with_edge_merge`. +- **Postgres-only helpers** (via `getattr`, not on the ABC): `traverse`, `find_connected_nodes`, `save_with_edge_merge` (persist mode only). - **Atomic JSON writes use `temp + fsync + rename + fsync(dir)`.** No partial records survive a crash. (`_atomic.py`) - **Per-file locks serialize concurrent writes to the same record only.** Different files run in parallel. (`_path_locks.py`) - **`QueryEngine` LRU is bounded.** Default 1024; configurable. Unbounded query construction will not leak memory. (`query.py`) diff --git a/jvspatial/db/database.py b/jvspatial/db/database.py index 525ee49..fe56989 100644 --- a/jvspatial/db/database.py +++ b/jvspatial/db/database.py @@ -197,10 +197,22 @@ class Database(ABC): with ACID semantics (e.g. MongoDB replica set). ``False`` for adapters where transactions are unavailable or only available in a weak buffered form. Default ``False``. + + ``edge_ids_mode`` + Where node adjacency lives. ``"persist"``: every node document + carries its incident edge ids in a top-level ``edges`` array that + ``connect()`` / ``disconnect()`` rewrite — needed when the edge + collection is not indexed on ``source`` / ``target``. ``"derive"``: + the edge collection is the only source of truth and adjacency is + queried from it, so ``connect()`` / ``save()`` cost O(1) in node + degree. Default ``"persist"``. Resolve the effective value with + :func:`resolve_edge_ids_mode` (honors instance overrides and the + ``JVSPATIAL_NODE_EDGE_IDS`` env var). """ # Capability flags. Override in subclasses. supports_transactions: bool = False + edge_ids_mode: str = "persist" @abstractmethod async def save(self, collection: str, data: Dict[str, Any]) -> Dict[str, Any]: @@ -642,14 +654,52 @@ async def drop_deprecated_indexes(self, deprecated: Dict[str, List[str]]) -> Non return None +EDGE_IDS_MODES = ("persist", "derive") + + +def resolve_edge_ids_mode(db: Any) -> str: + """Return the effective node-adjacency mode (``"persist"`` / ``"derive"``). + + Precedence: an ``edge_ids_mode`` set on the adapter *instance* (explicit + config) → the ``JVSPATIAL_NODE_EDGE_IDS`` env var → the adapter class + default (:attr:`Database.edge_ids_mode`). Wrapping adapters + (observability, caching) are unwrapped through their ``inner`` attribute + so the innermost backend's capability decides. + """ + candidate = db + for _ in range(6): + explicit = getattr(candidate, "__dict__", {}).get("edge_ids_mode") + if explicit in EDGE_IDS_MODES: + return str(explicit) + inner = getattr(candidate, "inner", None) + if inner is None: + break + candidate = inner + + from jvspatial.env import env + + override = (env("JVSPATIAL_NODE_EDGE_IDS") or "").strip().lower() + if override in EDGE_IDS_MODES: + return override + if override: + logger.warning( + "Ignoring JVSPATIAL_NODE_EDGE_IDS=%r (expected 'persist' or 'derive')", + override, + ) + mode = getattr(type(candidate), "edge_ids_mode", "persist") + return mode if mode in EDGE_IDS_MODES else "persist" + + __all__ = [ "Database", "DatabaseError", "VersionConflictError", "BulkSaveResult", + "EDGE_IDS_MODES", "encode_cursor", "decode_cursor", "finalize_find_results", + "resolve_edge_ids_mode", "resolve_sort_value", ] diff --git a/jvspatial/db/mongodb.py b/jvspatial/db/mongodb.py index 8350f6f..2471b6e 100644 --- a/jvspatial/db/mongodb.py +++ b/jvspatial/db/mongodb.py @@ -70,6 +70,11 @@ class MongoDB(Database): # (audit §5.9 / SPEC §4.2). supports_transactions: bool = True + # Node adjacency is derived from the edge collection (indexed on + # source/target by ``Edge.get_indexes``); node documents carry no + # ``edges`` array. See ``jvspatial.db.database.resolve_edge_ids_mode``. + edge_ids_mode: str = "derive" + def __init__( self, uri: str = "mongodb://localhost:27017", @@ -295,6 +300,50 @@ async def _delete_op() -> None: await self._run_with_reconnect("delete", _delete_op) + async def strip_node_edges( + self, + collection: str = "node", + *, + batch_size: int = 5000, + dry_run: bool = False, + ) -> int: + """Remove the legacy ``edges`` array from node documents (derive-mode migration). + + Batched ``$unset`` — idempotent and safe to run while the application + serves traffic in derive mode (which never writes the array back). + + Returns: + Documents stripped, or — with ``dry_run`` — documents that still + carry the array. + """ + if batch_size < 1: + raise ValueError(f"batch_size must be >= 1, got {batch_size}") + legacy = {"edges": {"$exists": True}} + + async def _strip_op() -> int: + await self._ensure_connected() + if self._db is None: + raise DatabaseError("MongoDB database connection not established") + collection_obj = self._db[collection] + if dry_run: + return int(await collection_obj.count_documents(legacy)) + stripped = 0 + while True: + ids = [ + doc["_id"] + async for doc in collection_obj.find(legacy, {"_id": 1}).limit( + batch_size + ) + ] + if not ids: + return stripped + result = await collection_obj.update_many( + {"_id": {"$in": ids}}, {"$unset": {"edges": ""}} + ) + stripped += int(result.modified_count) + + return int(await self._run_with_reconnect("strip_node_edges", _strip_op)) + async def find( self, collection: str, diff --git a/jvspatial/db/postgres.py b/jvspatial/db/postgres.py index 8f01d70..d8f8daf 100644 --- a/jvspatial/db/postgres.py +++ b/jvspatial/db/postgres.py @@ -271,6 +271,11 @@ class PostgresDB(Database): # at the operation level (atomic single-row update). supports_transactions: bool = True + # Node adjacency is derived from the indexed edge table; node rows do not + # carry an ``edges`` array, so connect()/save() never rewrite (or row-lock) + # a hub. See :func:`jvspatial.db.database.resolve_edge_ids_mode`. + edge_ids_mode: str = "derive" + def __init__( self, dsn: Optional[str] = None, @@ -721,7 +726,12 @@ async def save(self, collection: str, data: Dict[str, Any]) -> Dict[str, Any]: async def save_with_edge_merge( self, collection: str, data: Dict[str, Any] ) -> Dict[str, Any]: - """Upsert a record, unioning ``edges`` with any existing row in one statement.""" + """Upsert a record, unioning ``edges`` with any existing row in one statement. + + Only used in ``edge_ids_mode="persist"``. In the default ``"derive"`` + mode node rows carry no ``edges`` array and ``GraphContext.save`` + writes them with a plain :meth:`save`. + """ await self._bootstrap_collection(collection) rec_id, entity, tenant, _ = self._split_payload(data) col = _safe_collection(collection) @@ -766,6 +776,62 @@ async def save_with_edge_merge( result = self._record_from_row(row) if row is not None else data return result if result is not None else data + async def strip_node_edges( + self, + collection: str = "node", + *, + batch_size: int = 5000, + dry_run: bool = False, + ) -> int: + """Remove the legacy ``edges`` array from node rows (derive-mode migration). + + One keyset pass over the primary key, stripping ``edges`` from each + batch with ``UPDATE … SET data = data - 'edges'`` — idempotent and safe + to run while the application serves traffic in derive mode (which + never writes the array back). Tables with ``FORCE ROW LEVEL SECURITY`` + must be migrated by a role that bypasses RLS. + + Afterwards run ``VACUUM (ANALYZE) `` and + ``REINDEX INDEX CONCURRENTLY _data_gin`` to reclaim the + space; neither is run automatically. + + Returns: + Rows stripped, or — with ``dry_run`` — rows that still carry + the array. + """ + if batch_size < 1: + raise ValueError(f"batch_size must be >= 1, got {batch_size}") + await self._bootstrap_collection(collection) + col = _safe_collection(collection) + schema = _safe_collection(self.schema_name) + pool = await self._ensure_pool() + async with pool.acquire() as conn: + if dry_run: + return int( + await conn.fetchval( + f"SELECT COUNT(*) FROM {schema}.{col} WHERE data ? 'edges'" + ) + ) + stripped = 0 + last_id = "" + while True: + rows = await conn.fetch( + f"SELECT id FROM {schema}.{col} WHERE id > $1 ORDER BY id LIMIT $2", + last_id, + batch_size, + ) + if not rows: + return stripped + ids = [r["id"] for r in rows] + last_id = ids[-1] + status = await conn.execute( + f"UPDATE {schema}.{col} SET data = data - 'edges', " + f"updated_at = NOW() " + f"WHERE id = ANY($1::text[]) AND data ? 'edges'", + ids, + ) + stripped += int(status.rsplit(" ", 1)[1]) + async def get(self, collection: str, id: str) -> Optional[Dict[str, Any]]: """Fetch a single record by id.""" await self._bootstrap_collection(collection) diff --git a/jvspatial/db/query.py b/jvspatial/db/query.py index 93f7d73..bc72582 100644 --- a/jvspatial/db/query.py +++ b/jvspatial/db/query.py @@ -579,6 +579,15 @@ def apply_update( if item not in arr: arr.append(item) QueryEngine.set_field_value(document, field, arr) + elif op == "$pull": + # Scalar-value form only (``{"$pull": {"tags": "x"}}``) — + # removes every element equal to ``item``. + for field, item in payload.items(): + arr = QueryEngine.get_field_value(document, field) + if isinstance(arr, list): + QueryEngine.set_field_value( + document, field, [v for v in arr if v != item] + ) else: continue return document diff --git a/jvspatial/db/sqlite.py b/jvspatial/db/sqlite.py index 170d062..cf1f2d0 100644 --- a/jvspatial/db/sqlite.py +++ b/jvspatial/db/sqlite.py @@ -67,6 +67,11 @@ class SQLiteDB(Database): service; SQLite does not.) """ + # Node adjacency is derived from the edge collection, which + # ``Edge.get_indexes`` covers with ``json_extract`` indexes on + # source/target. See ``jvspatial.db.database.resolve_edge_ids_mode``. + edge_ids_mode: str = "derive" + def __init__( self, db_path: Optional[Union[str, Path]] = None, diff --git a/jvspatial/env_adapter.py b/jvspatial/env_adapter.py index 0a9d384..32940ce 100644 --- a/jvspatial/env_adapter.py +++ b/jvspatial/env_adapter.py @@ -273,6 +273,7 @@ def server_config_overrides_from_env() -> Dict[str, Any]: "JVSPATIAL_DYNAMODB_ENDPOINT_URL", "JVSPATIAL_DYNAMODB_WAIT_FOR_INDEX", "JVSPATIAL_AUTO_CREATE_INDEXES", + "JVSPATIAL_NODE_EDGE_IDS", # Auth "JVSPATIAL_AUTH_ENABLED", "JVSPATIAL_AUTH_STRICT_HASHING", diff --git a/tests/benchmarks/test_hub_node_bench.py b/tests/benchmarks/test_hub_node_bench.py index 3d1d03e..1386d71 100644 --- a/tests/benchmarks/test_hub_node_bench.py +++ b/tests/benchmarks/test_hub_node_bench.py @@ -145,7 +145,9 @@ async def _bench_graph() -> AsyncIterator[Tuple[GraphContext, Any, str]]: "postgres", dsn=_DSN, schema_name=schema, - min_size=1, + # Warm pool: open every connection up front so samples (and the + # 32-way burst) measure queries, not connection establishment. + min_size=_CONCURRENT_CONNECTS + 8, max_size=_CONCURRENT_CONNECTS + 8, observe=True, slow_query_ms=1e9, @@ -346,6 +348,11 @@ async def test_hub_node_scale(degree: int, request: pytest.FixtureRequest) -> No degree, ), } + if hasattr(hub, "count_nodes"): # 0.0.18+: the replacement for len() + read_cases["count_nodes"] = ( + lambda: hub.count_nodes(edge=[BenchContains]), + degree, + ) for name, (fn, expected) in read_cases.items(): samples = [] trips = 0 @@ -353,7 +360,8 @@ async def test_hub_node_scale(degree: int, request: pytest.FixtureRequest) -> No await ctx.clear_cache() ms, trips, out = await _timed(fn) samples.append(ms) - assert len(out) == expected, f"{name}: {len(out)} != {expected}" + got = out if isinstance(out, int) else len(out) + assert got == expected, f"{name}: {got} != {expected}" results[name] = {**_summary(samples), "round_trips": trips} # -- save() after a scalar change -- @@ -392,6 +400,12 @@ async def _one(leaf: BenchLeaf) -> float: await hub.connect(leaf, edge=BenchContains) return (time.perf_counter() - t) * 1000.0 + # asyncpg closes connections idle > 300 s, which the long 100k + # read phase exceeds; re-open them so the burst measures locking, + # not reconnects. + await asyncio.gather( + *(db.get("node", seeded["hub_id"]) for _ in range(len(conc_leaves))) + ) t0 = time.perf_counter() per_call = await asyncio.gather(*(_one(leaf) for leaf in conc_leaves)) wall_ms = (time.perf_counter() - t0) * 1000.0 diff --git a/tests/core/test_edge_ids_derive.py b/tests/core/test_edge_ids_derive.py new file mode 100644 index 0000000..210fb61 --- /dev/null +++ b/tests/core/test_edge_ids_derive.py @@ -0,0 +1,415 @@ +"""Node adjacency modes (``edge_ids_mode``) — parity across backends. + +Every scenario runs on each reachable backend in both modes: ``persist`` +(node rows carry an ``edges`` array) and ``derive`` (the edge collection is +the only source of truth). Results must match; derive mode must never write +an ``edges`` array or lock a node row on ``connect()``. + +Postgres runs when ``JVSPATIAL_POSTGRES_TEST_DSN`` is reachable, MongoDB when +``JVSPATIAL_MONGODB_TEST_URI`` is set and reachable; otherwise those params skip. +""" + +from __future__ import annotations + +import asyncio +import os +import uuid +from typing import Any, Dict, List, Tuple + +import pytest + +from jvspatial.cli import _run_strip_node_edges, build_parser +from jvspatial.core import context as context_module +from jvspatial.core.context import GraphContext, set_default_context +from jvspatial.core.entities import Edge, Node, Root +from jvspatial.core.graph_expansion import expand_node +from jvspatial.db._observable import ObservableDatabase +from jvspatial.db.database import resolve_edge_ids_mode +from jvspatial.db.jsondb import JsonDB +from jvspatial.db.sqlite import SQLiteDB + +pytestmark = pytest.mark.asyncio(loop_scope="module") + +_PG_DSN = os.getenv( + "JVSPATIAL_POSTGRES_TEST_DSN", + "postgresql://jvspatial:jvspatial@localhost:5432/jvspatial", +) +_MONGO_URI = os.getenv("JVSPATIAL_MONGODB_TEST_URI") + + +class DerivePerson(Node): + name: str = "" + + +class DeriveKnows(Edge): + pass + + +class DeriveLikes(Edge): + pass + + +async def _pg_admin() -> Any: + try: + import asyncpg + except ImportError: + pytest.skip("asyncpg not installed") + try: + return await asyncio.wait_for(asyncpg.connect(dsn=_PG_DSN), timeout=2.0) + except Exception: + pytest.skip(f"Postgres unreachable at {_PG_DSN}") + + +async def _mongo_client() -> Any: + if not _MONGO_URI: + pytest.skip("JVSPATIAL_MONGODB_TEST_URI not set") + from motor.motor_asyncio import AsyncIOMotorClient + + client = AsyncIOMotorClient(_MONGO_URI, serverSelectionTimeoutMS=2000) + try: + await client.admin.command("ping") + except Exception: + client.close() + pytest.skip(f"MongoDB unreachable at {_MONGO_URI}") + return client + + +_PARAMS = [ + (backend, mode) + for backend in ("json", "sqlite", "postgres", "mongodb") + for mode in ("persist", "derive") +] + + +@pytest.fixture(params=_PARAMS, ids=[f"{b}-{m}" for b, m in _PARAMS]) +async def graph(request, tmp_path): + """Yield ``(ctx, db, mode, backend)`` with the mode set explicitly.""" + backend, mode = request.param + cleanup: List[Any] = [] + if backend == "json": + db: Any = JsonDB(base_path=str(tmp_path / "json")) + elif backend == "sqlite": + db = SQLiteDB(db_path=str(tmp_path / "graph.db")) + elif backend == "postgres": + from jvspatial.db.postgres import PostgresDB + + admin = await _pg_admin() + schema = f"t_derive_{uuid.uuid4().hex[:10]}" + await admin.execute(f"CREATE SCHEMA {schema}") + db = PostgresDB(dsn=_PG_DSN, schema_name=schema, min_size=1, max_size=8) + + async def _drop_pg() -> None: + await admin.execute(f"DROP SCHEMA {schema} CASCADE") + await admin.close() + + cleanup.append(_drop_pg) + else: + from jvspatial.db.mongodb import MongoDB + + client = await _mongo_client() + name = f"jvs_derive_{uuid.uuid4().hex[:10]}" + db = MongoDB(uri=_MONGO_URI, db_name=name) + + async def _drop_mongo() -> None: + await client.drop_database(name) + client.close() + + cleanup.append(_drop_mongo) + + db.edge_ids_mode = mode + ctx = GraphContext(database=db) + set_default_context(ctx) + context_module._ensured_indexes.clear() + try: + yield ctx, db, mode, backend + finally: + set_default_context(None) + context_module._ensured_indexes.clear() + close = getattr(db, "close", None) + if callable(close): + result = close() + if asyncio.iscoroutine(result): + await result + for fn in cleanup: + await fn() + + +async def _fresh(ctx: GraphContext, node: Node) -> Node: + """Re-read ``node`` from the database, bypassing the entity cache.""" + await ctx.clear_cache() + loaded = await ctx.get(type(node), node.id) + assert loaded is not None + return loaded + + +async def _triangle() -> ( + Tuple[DerivePerson, DerivePerson, DerivePerson, Dict[str, str]] +): + a = await DerivePerson.create(name="a") + b = await DerivePerson.create(name="b") + c = await DerivePerson.create(name="c") + ab = await a.connect(b, edge=DeriveKnows) + ac = await a.connect(c, edge=DeriveLikes) + cb = await c.connect(b, edge=DeriveKnows) + return a, b, c, {"ab": ab.id, "ac": ac.id, "cb": cb.id} + + +async def _dangling_edges(db: Any) -> List[str]: + node_ids = {n["id"] for n in await db.find("node", {})} + return [ + e["id"] + for e in await db.find("edge", {}) + if e.get("source") not in node_ids or e.get("target") not in node_ids + ] + + +async def test_mode_resolution(graph): + ctx, _db, mode, _backend = graph + assert ctx.persists_edge_ids() is (mode == "persist") + + +async def test_connect_and_read_parity(graph): + ctx, _db, _mode, _backend = graph + a, b, c, eids = await _triangle() + + assert {e.target for e in await a.edges("out")} == {b.id, c.id} + assert {e.source for e in await b.edges("in")} == {a.id, c.id} + assert {e.id for e in await c.edges()} == {eids["ac"], eids["cb"]} + assert len(await a.edges(limit=1)) == 1 + + for node, degree in ((a, 2), (b, 2), (c, 2)): + assert await node.connection_count() == degree + assert await (await _fresh(ctx, node)).connection_count() == degree + + a = await _fresh(ctx, a) + assert {n.id for n in await a.nodes()} == {b.id, c.id} + assert {n.id for n in await a.nodes(edge=DeriveKnows)} == {b.id} + assert {n.id for n in await b.nodes(direction="in")} == {a.id, c.id} + + again = await a.connect(b, edge=DeriveKnows) + assert again.id == eids["ab"] + assert await a.connection_count() == 2 + + +async def test_disconnect_parity(graph): + ctx, db, _mode, _backend = graph + a, b, c, eids = await _triangle() + + assert await a.disconnect(b) is True + assert await db.get("edge", eids["ab"]) is None + assert await (await _fresh(ctx, a)).connection_count() == 1 + assert await (await _fresh(ctx, b)).connection_count() == 1 + assert {n.id for n in await a.nodes()} == {c.id} + + +async def test_save_writes_edges_array_only_in_persist_mode(graph): + _ctx, db, mode, _backend = graph + a, _b, _c, eids = await _triangle() + a.name = "renamed" + await a.save() + + raw = await db.get("node", a.id) + assert raw["context"]["name"] == "renamed" + if mode == "derive": + assert "edges" not in raw + assert a.edge_ids == [] + else: + assert set(raw["edges"]) == {eids["ab"], eids["ac"]} + + +async def test_derive_connect_takes_no_node_row_lock(graph, monkeypatch): + _ctx, db, mode, _backend = graph + if mode != "derive": + pytest.skip("derive-mode property") + hub = await DerivePerson.create(name="hub") + leaves = [await DerivePerson.create(name=f"l{i}") for i in range(3)] + + node_writes: List[str] = [] + real_save = db.save + + async def _spy_find_one_and_update(*args: Any, **kwargs: Any) -> Any: + raise AssertionError("connect() must not read-modify-write a node row") + + async def _spy_save(collection: str, data: Dict[str, Any]) -> Any: + if collection == "node": + node_writes.append(data["id"]) + return await real_save(collection, data) + + monkeypatch.setattr(db, "find_one_and_update", _spy_find_one_and_update) + monkeypatch.setattr(db, "save", _spy_save) + for leaf in leaves: + await hub.connect(leaf, edge=DeriveKnows) + await asyncio.gather(*(hub.connect(leaf, edge=DeriveLikes) for leaf in leaves)) + await hub.disconnect(leaves[0]) + + assert node_writes == [] + assert await hub.connection_count() == 4 + + +async def test_concurrent_connect_to_one_hub(graph): + ctx, db, mode, _backend = graph + hub = await DerivePerson.create(name="hub") + leaves = [await DerivePerson.create(name=f"l{i}") for i in range(12)] + await asyncio.gather(*(hub.connect(leaf, edge=DeriveKnows) for leaf in leaves)) + + assert await (await _fresh(ctx, hub)).connection_count() == 12 + assert {n.id for n in await hub.nodes()} == {leaf.id for leaf in leaves} + if mode == "persist": + assert len((await db.get("node", hub.id))["edges"]) == 12 + + +async def test_legacy_edges_array_is_ignored_in_derive_mode(graph): + ctx, db, mode, _backend = graph + if mode != "derive": + pytest.skip("derive-mode property") + a = await DerivePerson.create(name="a") + b = await DerivePerson.create(name="b") + await a.connect(b, edge=DeriveKnows) + + raw = await db.get("node", a.id) + raw["edges"] = ["e.DeriveKnows.stale000000000000000000"] + await db.save("node", raw) + + legacy = await _fresh(ctx, a) + assert legacy.edge_ids == [] + assert await legacy.connection_count() == 1 + assert {n.id for n in await legacy.nodes()} == {b.id} + + legacy.name = "resaved" + await legacy.save() + assert "edges" not in await db.get("node", a.id) + + +async def test_cascade_delete_parity(graph): + ctx, db, _mode, _backend = graph + parent = await DerivePerson.create(name="parent") + a = await DerivePerson.create(name="a") + only_via_a = await DerivePerson.create(name="b") + shared = await DerivePerson.create(name="c") + outsider = await DerivePerson.create(name="x") + await parent.connect(a, edge=DeriveKnows) + await a.connect(only_via_a, edge=DeriveKnows) + await a.connect(shared, edge=DeriveKnows) + await outsider.connect(shared, edge=DeriveKnows) + + await a.delete(cascade=True) + await ctx.clear_cache() + + assert await db.get("node", a.id) is None + assert await db.get("node", only_via_a.id) is None + for survivor in (parent, shared, outsider): + assert await db.get("node", survivor.id) is not None + assert await (await _fresh(ctx, parent)).connection_count() == 0 + assert await (await _fresh(ctx, shared)).connection_count() == 1 + assert await _dangling_edges(db) == [] + + +async def test_context_delete_without_cascade_removes_edges(graph): + ctx, db, _mode, _backend = graph + a = await DerivePerson.create(name="a") + b = await DerivePerson.create(name="b") + await a.connect(b, edge=DeriveKnows) + + await ctx.delete(a) + + assert await db.get("node", a.id) is None + assert await (await _fresh(ctx, b)).connection_count() == 0 + assert await _dangling_edges(db) == [] + + +async def test_root_rehydrates(graph): + _ctx, _db, mode, _backend = graph + root = await Root.get() + app = await DerivePerson.create(name="app") + await root.connect(app, edge=DeriveKnows) + + again = await Root.get() + assert {n.id for n in await again.nodes()} == {app.id} + if mode == "derive": + assert again.edge_ids == [] + + +async def test_expand_node_pages_from_edge_collection(graph): + ctx, _db, _mode, _backend = graph + hub = await DerivePerson.create(name="hub") + kids = [await DerivePerson.create(name=f"k{i}") for i in range(5)] + for kid in kids: + await hub.connect(kid, edge=DeriveKnows) + + p1 = await expand_node(ctx, hub.id, limit=2) + assert p1["pagination"]["total_edge_count"] == 5 + assert p1["pagination"]["has_more"] is True + assert p1["pagination"]["next_cursor"] == 2 + center = next(n for n in p1["nodes"] if n["id"] == hub.id) + assert center["degree"] == 5 + assert all(n["degree"] == 1 for n in p1["nodes"] if n["id"] != hub.id) + + seen = [e["id"] for e in p1["edges"]] + after = p1["pagination"]["next_after"] + while after: + page = await expand_node(ctx, hub.id, limit=2, after=after) + seen += [e["id"] for e in page["edges"]] + after = page["pagination"]["next_after"] + assert len(seen) == len(set(seen)) == 5 + + offset = await expand_node(ctx, hub.id, limit=10, cursor=4) + assert offset["pagination"]["returned_edges"] == 1 + assert offset["pagination"]["has_more"] is False + + +async def test_strip_node_edges_migration(graph): + ctx, db, mode, backend = graph + if backend not in ("postgres", "mongodb") or mode != "persist": + pytest.skip("migration runs on Postgres/MongoDB rows written in persist mode") + a, b, c, _eids = await _triangle() + before = {n.id for n in await a.nodes()} + assert "edges" in await db.get("node", a.id) + + assert await db.strip_node_edges(dry_run=True) == 3 + db.edge_ids_mode = "derive" + assert await db.strip_node_edges(batch_size=2) == 3 + assert await db.strip_node_edges() == 0 + for node in (a, b, c): + assert "edges" not in await db.get("node", node.id) + + fresh_a = await _fresh(ctx, a) + assert {n.id for n in await fresh_a.nodes()} == before + assert await fresh_a.connection_count() == 2 + + # CLI: dry run by default, --apply strips. + raw = await db.get("node", b.id) + raw["edges"] = ["e.DeriveKnows.legacy"] + await db.save("node", raw) + target = ( + ["--dsn", _PG_DSN, "--schema", db.schema_name] + if backend == "postgres" + else ["--dsn", _MONGO_URI, "--db-name", db.db_name] + ) + parser = build_parser() + dry = parser.parse_args(["migrate", "strip-node-edges", *target]) + assert await _run_strip_node_edges(dry) == 0 + assert "edges" in await db.get("node", b.id) + apply = parser.parse_args(["migrate", "strip-node-edges", *target, "--apply"]) + assert await _run_strip_node_edges(apply) == 0 + assert "edges" not in await db.get("node", b.id) + + +async def test_mode_resolution_precedence(tmp_path, monkeypatch): + monkeypatch.delenv("JVSPATIAL_NODE_EDGE_IDS", raising=False) + sqlite = SQLiteDB(db_path=str(tmp_path / "p.db")) + json_db = JsonDB(base_path=str(tmp_path / "j")) + assert resolve_edge_ids_mode(sqlite) == "derive" + assert resolve_edge_ids_mode(json_db) == "persist" + assert resolve_edge_ids_mode(ObservableDatabase(sqlite)) == "derive" + + monkeypatch.setenv("JVSPATIAL_NODE_EDGE_IDS", "persist") + assert resolve_edge_ids_mode(sqlite) == "persist" + monkeypatch.setenv("JVSPATIAL_NODE_EDGE_IDS", "derive") + assert resolve_edge_ids_mode(ObservableDatabase(json_db)) == "derive" + + json_db.edge_ids_mode = "persist" # explicit instance config beats env + assert resolve_edge_ids_mode(ObservableDatabase(json_db)) == "persist" + + monkeypatch.setenv("JVSPATIAL_NODE_EDGE_IDS", "bogus") + assert resolve_edge_ids_mode(sqlite) == "derive" + await sqlite.close() diff --git a/tests/core/test_entity_crud_and_cascade.py b/tests/core/test_entity_crud_and_cascade.py index 09c75ec..c6b51ab 100644 --- a/tests/core/test_entity_crud_and_cascade.py +++ b/tests/core/test_entity_crud_and_cascade.py @@ -422,6 +422,10 @@ def temp_context(self): unique_path = f"{tmpdir}/test_{uuid.uuid4().hex}" config = {"db_type": "json", "db_config": {"base_path": unique_path}} database = create_database(config["db_type"], **config["db_config"]) + # Assertions read ``edge_ids``; pin persist so a + # JVSPATIAL_NODE_EDGE_IDS=derive run does not switch it off + # (derive-mode cascade is covered in test_edge_ids_derive.py). + database.edge_ids_mode = "persist" context = GraphContext(database=database) # Set as default context so entity methods use it set_default_context(context) diff --git a/tests/core/test_node_save_edges_merge.py b/tests/core/test_node_save_edges_merge.py index 8368fd2..ece3b05 100644 --- a/tests/core/test_node_save_edges_merge.py +++ b/tests/core/test_node_save_edges_merge.py @@ -24,6 +24,9 @@ class Widget(Node): async def graph_context(): with tempfile.TemporaryDirectory() as tmpdir: db = JsonDB(base_path=tmpdir) + # These tests cover the persisted ``edges`` merge contract; pin the + # mode so a JVSPATIAL_NODE_EDGE_IDS=derive run does not switch it off. + db.edge_ids_mode = "persist" ctx = GraphContext(database=db) set_default_context(ctx) yield ctx, db diff --git a/tests/db/test_jsondb.py b/tests/db/test_jsondb.py index 8b25f65..6bd78cc 100644 --- a/tests/db/test_jsondb.py +++ b/tests/db/test_jsondb.py @@ -697,6 +697,9 @@ async def context(self, jsondb): """Create GraphContext with JsonDB for testing.""" from jvspatial.core.context import GraphContext, set_default_context + # Assertions read the persisted ``edges`` array; pin the mode so a + # JVSPATIAL_NODE_EDGE_IDS=derive run does not switch it off. + jsondb.edge_ids_mode = "persist" ctx = GraphContext(database=jsondb) set_default_context(ctx) return ctx @@ -942,6 +945,9 @@ async def context(self, jsondb): """Create GraphContext with JsonDB for testing.""" from jvspatial.core.context import GraphContext, set_default_context + # Assertions read the persisted ``edges`` array; pin the mode so a + # JVSPATIAL_NODE_EDGE_IDS=derive run does not switch it off. + jsondb.edge_ids_mode = "persist" ctx = GraphContext(database=jsondb) set_default_context(ctx) return ctx diff --git a/tests/db/test_query_apply_update.py b/tests/db/test_query_apply_update.py new file mode 100644 index 0000000..b79f1f1 --- /dev/null +++ b/tests/db/test_query_apply_update.py @@ -0,0 +1,28 @@ +"""QueryEngine.apply_update — Mongo-style update operators applied in Python. + +Backends without native update operators (Postgres ``find_one_and_update``, +JsonDB / SQLite / DynamoDB defaults) rely on this, so every operator the +library emits must be honored rather than silently skipped. +""" + +from jvspatial.db.query import QueryEngine + + +def test_add_to_set_and_pull_round_trip(): + doc = {"id": "n.X.1", "edges": ["e.1"]} + QueryEngine.apply_update(doc, {"$addToSet": {"edges": "e.2"}}) + QueryEngine.apply_update(doc, {"$addToSet": {"edges": "e.2"}}) + assert doc["edges"] == ["e.1", "e.2"] + + QueryEngine.apply_update(doc, {"$pull": {"edges": "e.1"}}) + assert doc["edges"] == ["e.2"] + + +def test_pull_removes_every_match_and_ignores_missing_field(): + doc = {"tags": ["a", "b", "a"], "context": {"n": [1, 2, 1]}} + QueryEngine.apply_update(doc, {"$pull": {"tags": "a", "context.n": 1}}) + assert doc["tags"] == ["b"] + assert doc["context"]["n"] == [2] + + QueryEngine.apply_update(doc, {"$pull": {"absent": "x"}}) + assert "absent" not in doc diff --git a/tests/test_node_operations.py b/tests/test_node_operations.py index b6da782..b6ce379 100644 --- a/tests/test_node_operations.py +++ b/tests/test_node_operations.py @@ -20,6 +20,9 @@ async def context(): with tempfile.TemporaryDirectory() as tmpdir: db = JsonDB(os.path.join(tmpdir, "test.json")) + # Assertions read ``edge_ids``; pin persist so a + # JVSPATIAL_NODE_EDGE_IDS=derive run does not switch it off. + db.edge_ids_mode = "persist" ctx = GraphContext(database=db) # Set this as the default context so all entities use this database set_default_context(ctx) From e7dfe6fcea3220a20bf9b3ef0ee0aec26cb6e143 Mon Sep 17 00:00:00 2001 From: Eldon Marks Date: Fri, 11 Sep 2026 14:58:50 -0400 Subject: [PATCH 4/9] Push neighbour filters, limits, counts and pages into the database Phase 2 of the hub-node scale remediation. Node.nodes() normalises every filter shape - classes (subclass-inclusive on the node side), names, lists, {Name: criteria} dicts and property kwargs - to entity lists plus record-path queries and sends them, limit included, to the backend's find_connected_nodes in one round trip, direction="both" included. - PostgresDB: shared SQL builder; join form (one edge type, one direction: duplicate-free under the (source, target, entity) unique index, LIMIT streams) or id IN (...) semi-join (dedups); filters that do not translate raise NotImplementedError so nothing is dropped; count_connected_nodes, find_connected_nodes_bulk (ROW_NUMBER per source) - MongoDB ($lookup aggregation) and SQLite (json_extract join) implement the same contract; table-alias support in both SQL translators - Node.count_nodes(), Node.nodes_page() (keyset; cursor helpers moved to core/pager.py and shared with GraphContext.find_page), nodes_bulk(limit_per_source=); count_neighbors delegates to count_nodes - Python fallback keeps every filter and filters in the database find - Fix: list-form edge filters were ignored by nodes() - Fix: ObservableDatabase exposed find_connected_nodes/traverse even when the wrapped backend lacks them Bench (100k hub): nodes(edge=[E], node=["Leaf"], limit=20) 4.9 s / 201 round trips -> 1.2 ms / 1; count 4.8 s (len(nodes())) -> 83 ms / 1. Co-Authored-By: Claude Opus 5 --- CHANGELOG.md | 25 + SPEC.md | 10 + docs/bench/2026-09-hub-node-baseline.md | 41 +- docs/bench/2026-09-hub-node-phase2.jsonl | 3 + docs/md/entity-reference.md | 4 +- docs/md/optimization.md | 26 + jvspatial/core/context.py | 142 ++--- jvspatial/core/entities/node.py | 630 +++++++++++-------- jvspatial/core/pager.py | 97 ++- jvspatial/db/README.md | 3 +- jvspatial/db/_observable.py | 95 ++- jvspatial/db/_postgres_translate.py | 53 +- jvspatial/db/_sqlite_translate.py | 43 +- jvspatial/db/mongodb.py | 163 ++++- jvspatial/db/postgres.py | 287 ++++++++- jvspatial/db/sqlite.py | 207 +++++- tests/core/test_connected_nodes_fast_path.py | 22 +- tests/core/test_node_query_pushdown.py | 474 ++++++++++++++ 18 files changed, 1876 insertions(+), 449 deletions(-) create mode 100644 docs/bench/2026-09-hub-node-phase2.jsonl create mode 100644 tests/core/test_node_query_pushdown.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 974c32b..5aff56f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -22,6 +22,13 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 `Node.edges(limit=...)`, `strip_node_edges()` on `PostgresDB` / `MongoDB`, and the `jvspatial migrate strip-node-edges` CLI (dry run by default, `--apply` to write). +- **Neighbour counts and pages in the database** (`jvspatial/core/entities/node.py`): + `Node.count_nodes(...)` (one `COUNT` round trip — use it instead of + `len(await n.nodes(...))`), `Node.nodes_page(sort=..., cursor=..., limit=...)` + (keyset-paginated neighbours, cursor encoding shared with + `GraphContext.find_page` via `jvspatial.core.pager`), and + `nodes_bulk(limit_per_source=...)` (one windowed query on Postgres). + `count_neighbors` now delegates to `count_nodes`. - **`expand_node` keyset paging** (`jvspatial/core/graph_expansion.py`): `after=` / `pagination.next_after` (and the `after` query parameter on the graph expand endpoint) page incident edges by edge id in O(page). The @@ -41,6 +48,14 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 online) to reclaim the space at once. JsonDB and DynamoDB keep `"persist"`. Opt out with `JVSPATIAL_NODE_EDGE_IDS=persist` or `db.edge_ids_mode = "persist"`. +- **`Node.nodes()` pushes every filter shape to the database.** The list + form (`edge=[E], node=["Leaf"]`), strings, subclass-inclusive node classes, + `{Name: criteria}` dicts and property kwargs all become one + `find_connected_nodes` round trip on Postgres, MongoDB and SQLite, with + `limit` pushed down; `direction="both"` included. Previously only the class + form took the join and everything else loaded every incident edge and + hydrated every neighbour before slicing in Python. Backends without the + pushdown keep the Python path, now filtering in the database `find`. - **`expand_node` / `subgraph_bfs` source incident edges from the edge collection** in every mode (one query per node instead of one `get` per edge id); `total_edge_count` is a count of incident edges, and neighbours @@ -48,6 +63,16 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Fixed +- **List-form edge filters were dropped by `Node.nodes()`** + (`jvspatial/core/entities/node.py`). `nodes(edge=[E])`, `edge=["E"]` and + `edge=[{"E": {...}}]` returned neighbours reached through *any* edge type, + and `direction="both"` capped at 10 000 edges. Edge types and criteria are + now always applied. +- **`ObservableDatabase` advertised graph pushdowns its backend lacks** + (`jvspatial/db/_observable.py`). `find_connected_nodes` / `traverse` were + plain methods, so `getattr(db, "find_connected_nodes", None)` looked + callable over JsonDB and the call then raised. They now exist on the + wrapper only when the wrapped adapter implements them. - **`$pull` was ignored by `QueryEngine.apply_update`** (`jvspatial/db/query.py`). Backends that apply updates in Python — Postgres `find_one_and_update`, and the JsonDB / SQLite / DynamoDB defaults — silently skipped it, so on diff --git a/SPEC.md b/SPEC.md index a1f0efb..2ae5282 100644 --- a/SPEC.md +++ b/SPEC.md @@ -365,6 +365,16 @@ Enabling prefetch may enqueue neighbors before hook-driven `visit()` calls; visi `await node.neighborhood(depth=k, direction=..., edge=..., node=..., limit=...)` returns hydrated `Node` instances within `k` hops. Postgres uses `Database.traverse` + `get_batch`; other backends use per-hop `nodes()` BFS. +### 6.8 Neighbour queries (`Node.nodes` / `count_nodes` / `nodes_page`) + +`nodes(direction, node=, edge=, limit=, **props)` normalises every filter shape to `(edge_entities, edge_query, node_entities, node_query)` (`_neighbor_filter_spec`, `node.py`): + +- a name or list of names → exact entity match; an **edge** class → its exact entity; a **node** class → the class and every loaded subclass (`isinstance` semantics; `Node` itself means any type); +- `{Name: criteria}` items → an `$or` of `{"entity": Name, }` branches; property kwargs AND onto the node side; +- bare criteria keys map to `context.` (the attribute `_matches_property_filter` reads); `id`, `entity` (and `source`, `target`, `bidirectional` on edges) are top-level record paths. + +When the backend exposes `find_connected_nodes` (Postgres, MongoDB, SQLite) the spec plus `limit` go down in **one** round trip, `direction="both"` included. Each neighbour appears once. Postgres picks a join for one edge type in one direction (duplicate-free under the `(source, target, entity)` unique index, so `LIMIT` streams) and an `id IN (...)` semi-join otherwise. A criterion that does not translate raises `NotImplementedError` in the adapter and `nodes()` takes the Python path — one edge `find` plus a node `find` that still applies every filter; no filter is ever dropped. `count_nodes(...)` is the matching `COUNT` (`count_connected_nodes`; `count_neighbors` is an alias). `nodes_page(sort=, cursor=, limit=)` pages neighbours by keyset with the cursor encoding of `GraphContext.find_page` (`core/pager.py`); the keyset is expressed as a node-side query, so it follows the same pushdown. `nodes_bulk(..., limit_per_source=N)` caps neighbours per source (`ROW_NUMBER() OVER (PARTITION BY source)` on Postgres). + --- ## 7. GraphContext and Dependency Injection diff --git a/docs/bench/2026-09-hub-node-baseline.md b/docs/bench/2026-09-hub-node-baseline.md index e83eb76..3a7d6a2 100644 --- a/docs/bench/2026-09-hub-node-baseline.md +++ b/docs/bench/2026-09-hub-node-baseline.md @@ -147,9 +147,46 @@ Code: `23af4c0` + the Phase 1 change set (worktree). 133 ms → 10.2 s, and grew with degree. - List-form reads are unchanged, as expected. They are Phase 2's target. -## After Phase 2 +## After Phase 2 — neighbour filters, limits and counts pushed into SQL -_Pending._ +Code: `dc2a6cc` + the Phase 2 change set. + +| operation | 1k p50 / p95 ms (trips) | 10k p50 / p95 ms (trips) | 100k p50 / p95 ms (trips) | +|---|---|---|---| +| `ctx.get(Hub)` (hydrate hub) | 0.5 / 4.0 | 0.8 / 1.4 | 0.6 / 6.4 | +| `hub.connect(leaf, edge=E)` | 1.8 / 2.3 (3) | 1.5 / 1.9 (3) | 1.6 / 2.8 (3) | +| `hub.save()` after scalar change | 0.7 / 1.0 | 0.7 / 3.2 | 0.9 / 2.2 | +| `hub.nodes(edge=[E], node=['Leaf'], limit=20)` | 2.1 / 9.7 (1) | 1.2 / 1.8 (1) | 1.2 / 3.1 (1) | +| `hub.nodes(edge=E, limit=20)` | 1.1 / 2.4 (1) | 1.1 / 1.5 (1) | 1.2 / 6.1 (1) | +| `sink.nodes(edge=[E], node=['Leaf'], direction='in', limit=20)` | 1.0 / 1.3 (1) | 1.0 / 1.3 (1) | 1.4 / 2.4 (1) | +| `sink.nodes(edge=[E], node=['Leaf'], direction='in')` | 18.7 / 44.7 (1) | 230.9 / 256.1 (1) | 2128.5 / 2187.0 (1) | +| `len(await hub.nodes(edge=[E]))` | 21.7 / 45.4 (1) | 226.0 / 246.9 (1) | 2112.8 / 2182.1 (1) | +| `hub.count_nodes(edge=[E])` | 3.0 / 3.3 (1) | 33.3 / 37.0 (1) | 82.8 / 100.0 (1) | +| 32× concurrent `connect()` (ms) | 16 wall / 15 max | 16 wall / 15 max | 16 wall / 16 max | + +| size (after seed → after 82 hub writes) | 1k | 10k | 100k | +|---|---|---|---| +| hub row `pg_column_size(data)` | 134 B → 136 B | 134 B → 136 B | 134 B → 136 B | +| `node_data_gin` size | 104.0 KB → 112.0 KB | 808.0 KB → 816.0 KB | 13.0 MB → 13.0 MB | +| `node` table total | 504.0 KB → 512.0 KB | 3.9 MB → 3.9 MB | 44.2 MB → 44.2 MB | +| `edge` table total | 1.7 MB → 1.9 MB | 16.6 MB → 16.7 MB | 162.2 MB → 162.3 MB | + +**Gate: passed.** + +- The list form `nodes(edge=[E], node=["Leaf"], limit=20)` is 1 round trip in + both directions at every tier. Its p50 is 1.0–2.1 ms and flat across tiers. + - Before, it cost 201 round trips and ~5 s at 100k. + - The outlier is the 1k out-direction p95 of 9.7 ms. That tier's other + samples sit at 1–2 ms, so it looks like noise. +- `count_nodes(edge=[E])` replaces `len(await nodes(...))` and is 1 round trip. + - It takes 83 ms at 100k, against 2.1 s for the `len()` pattern here and + 4.8 s on 0.0.17. + - It still grows with degree. `COUNT` over the `edge ⋈ node` join has to + visit every matching row. +- Unbounded listings are now 1 round trip too (2.1 s at 100k, down from ~5 s + and 201 trips). What remains is hydrating 100k `Node` objects in Python. + Page instead: `nodes_page(...)`. +- The Phase 1 numbers (connect, save, concurrency, sizes) are unchanged. ## After Phase 3 diff --git a/docs/bench/2026-09-hub-node-phase2.jsonl b/docs/bench/2026-09-hub-node-phase2.jsonl new file mode 100644 index 0000000..1a49985 --- /dev/null +++ b/docs/bench/2026-09-hub-node-phase2.jsonl @@ -0,0 +1,3 @@ +{"degree": 1000, "jvspatial": "0.0.17", "git_sha": "dc2a6cc", "edge_ids_mode": "derive", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 2.14, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 134, "node_gin_bytes": 106496, "node_total_bytes": 516096, "edge_total_bytes": 1826816}, "hub_get": {"p50_ms": 0.533, "p95_ms": 4.029, "max_ms": 5.628, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 2.055, "p95_ms": 9.724, "max_ms": 11.835, "n": 20, "round_trips": 1}, "nodes_class_out_limit20": {"p50_ms": 1.07, "p95_ms": 2.441, "max_ms": 3.794, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 1.029, "p95_ms": 1.287, "max_ms": 1.595, "n": 20, "round_trips": 1}, "nodes_list_in_unlimited": {"p50_ms": 18.731, "p95_ms": 44.683, "max_ms": 44.73, "n": 20, "round_trips": 1}, "count_via_len_nodes": {"p50_ms": 21.736, "p95_ms": 45.447, "max_ms": 55.71, "n": 20, "round_trips": 1}, "count_nodes": {"p50_ms": 3.033, "p95_ms": 3.292, "max_ms": 3.311, "n": 20, "round_trips": 1}, "hub_save": {"p50_ms": 0.705, "p95_ms": 0.976, "max_ms": 2.5, "n": 50}, "hub_connect": {"p50_ms": 1.756, "p95_ms": 2.313, "max_ms": 2.841, "n": 50, "round_trips": 3}, "concurrent_connect_32": {"wall_ms": 15.529, "max_call_ms": 15.417, "p50_call_ms": 13.809}, "sizes_after_writes": {"hub_row_bytes": 136, "node_gin_bytes": 114688, "node_total_bytes": 524288, "edge_total_bytes": 1966080}} +{"degree": 10000, "jvspatial": "0.0.17", "git_sha": "dc2a6cc", "edge_ids_mode": "derive", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 22.19, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 134, "node_gin_bytes": 827392, "node_total_bytes": 4104192, "edge_total_bytes": 17399808}, "hub_get": {"p50_ms": 0.781, "p95_ms": 1.373, "max_ms": 1.429, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 1.209, "p95_ms": 1.829, "max_ms": 3.471, "n": 20, "round_trips": 1}, "nodes_class_out_limit20": {"p50_ms": 1.052, "p95_ms": 1.491, "max_ms": 1.935, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 1.035, "p95_ms": 1.267, "max_ms": 1.935, "n": 20, "round_trips": 1}, "nodes_list_in_unlimited": {"p50_ms": 230.94, "p95_ms": 256.148, "max_ms": 258.194, "n": 20, "round_trips": 1}, "count_via_len_nodes": {"p50_ms": 226.038, "p95_ms": 246.9, "max_ms": 258.188, "n": 20, "round_trips": 1}, "count_nodes": {"p50_ms": 33.318, "p95_ms": 36.954, "max_ms": 41.222, "n": 20, "round_trips": 1}, "hub_save": {"p50_ms": 0.693, "p95_ms": 3.16, "max_ms": 5.457, "n": 50}, "hub_connect": {"p50_ms": 1.495, "p95_ms": 1.856, "max_ms": 2.561, "n": 50, "round_trips": 3}, "concurrent_connect_32": {"wall_ms": 15.931, "max_call_ms": 15.294, "p50_call_ms": 13.99}, "sizes_after_writes": {"hub_row_bytes": 136, "node_gin_bytes": 835584, "node_total_bytes": 4120576, "edge_total_bytes": 17489920}} +{"degree": 100000, "jvspatial": "0.0.17", "git_sha": "dc2a6cc", "edge_ids_mode": "derive", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 266.0, "read_iterations": 5, "sizes_after_seed": {"hub_row_bytes": 134, "node_gin_bytes": 13647872, "node_total_bytes": 46391296, "edge_total_bytes": 170098688}, "hub_get": {"p50_ms": 0.601, "p95_ms": 6.433, "max_ms": 6.433, "n": 5}, "nodes_list_out_limit20": {"p50_ms": 1.205, "p95_ms": 3.059, "max_ms": 3.059, "n": 5, "round_trips": 1}, "nodes_class_out_limit20": {"p50_ms": 1.212, "p95_ms": 6.148, "max_ms": 6.148, "n": 5, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 1.387, "p95_ms": 2.425, "max_ms": 2.425, "n": 5, "round_trips": 1}, "nodes_list_in_unlimited": {"p50_ms": 2128.484, "p95_ms": 2187.017, "max_ms": 2187.017, "n": 5, "round_trips": 1}, "count_via_len_nodes": {"p50_ms": 2112.814, "p95_ms": 2182.099, "max_ms": 2182.099, "n": 5, "round_trips": 1}, "count_nodes": {"p50_ms": 82.751, "p95_ms": 100.028, "max_ms": 100.028, "n": 5, "round_trips": 1}, "hub_save": {"p50_ms": 0.909, "p95_ms": 2.227, "max_ms": 3.361, "n": 50}, "hub_connect": {"p50_ms": 1.602, "p95_ms": 2.78, "max_ms": 7.358, "n": 50, "round_trips": 3}, "concurrent_connect_32": {"wall_ms": 15.696, "max_call_ms": 15.556, "p50_call_ms": 14.265}, "sizes_after_writes": {"hub_row_bytes": 136, "node_gin_bytes": 13647872, "node_total_bytes": 46399488, "edge_total_bytes": 170188800}} diff --git a/docs/md/entity-reference.md b/docs/md/entity-reference.md index 16cebd7..efe721f 100644 --- a/docs/md/entity-reference.md +++ b/docs/md/entity-reference.md @@ -71,8 +71,10 @@ rather than reading `edge_ids`. See [graph-context.md](graph-context.md). **Key Methods:** -- **`nodes()`**: Returns a list of connected nodes with filtering options +- **`nodes()`**: Returns a list of connected nodes with filtering options — one database round trip on Postgres, MongoDB and SQLite for every filter shape, `limit` included - **`node()`**: Returns a single connected node (first match) or None - convenience method when you expect only one result +- **`count_nodes()`**: Counts connected nodes with the same filters as `nodes()` (one `COUNT`; prefer it to `len(await n.nodes())`) +- **`nodes_page(sort=, cursor=, limit=)`**: Keyset-paginated neighbours, returns `(nodes, next_cursor)` - **`neighborhood(depth)`**: Multi-hop neighbor fetch (Postgres `traverse` fast path or BFS fallback) - **`nodes_bulk(node_ids)`**: Batch neighbor fetch for many source IDs in two queries - **`delete(cascade=True)`**: Deletes the node and cascades deletion of all connected edges and dependent nodes diff --git a/docs/md/optimization.md b/docs/md/optimization.md index 5c7857f..2118cdf 100644 --- a/docs/md/optimization.md +++ b/docs/md/optimization.md @@ -57,6 +57,32 @@ user = connected_nodes[0] if connected_nodes else None # Inefficient user = await current_node.node(node=User, direction="out") # Efficient ``` +#### Count and page neighbours in the database + +On Postgres, MongoDB and SQLite every `nodes()` filter shape — classes, +names, lists, `{Name: criteria}` dicts, property kwargs — plus `limit` runs as +one query, so `nodes(edge=[Contains], node=["Entry"], limit=20)` costs the same +on a node with 20 neighbours as on one with 100 000. Don't count or page by +listing: + +```python +# Bad: hydrates every neighbour just to count or slice them +total = len(await track.nodes(edge=[Contains])) +recent = (await track.nodes(edge=[Contains]))[:20] + +# Good: one COUNT, and keyset pages that stay O(page) +total = await track.count_nodes(edge=[Contains]) +page, cursor = await track.nodes_page( + edge=[Contains], sort=[("context.created_at", -1)], limit=20 +) +more, cursor = await track.nodes_page( + edge=[Contains], sort=[("context.created_at", -1)], limit=20, cursor=cursor +) +``` + +A sorted page still has to order every matching neighbour; an unsorted +`nodes(limit=N)` with a single edge type stops after N rows. + #### Bulk Query Optimization ```python diff --git a/jvspatial/core/context.py b/jvspatial/core/context.py index 506ca86..9f3dbea 100644 --- a/jvspatial/core/context.py +++ b/jvspatial/core/context.py @@ -1,10 +1,8 @@ """GraphContext for managing database dependencies.""" import asyncio -import base64 import contextvars import inspect -import json import logging import time from contextlib import asynccontextmanager, contextmanager, suppress @@ -23,11 +21,7 @@ cast, ) -from jvspatial.db.database import ( - Database, - resolve_edge_ids_mode, - resolve_sort_value, -) +from jvspatial.db.database import Database, resolve_edge_ids_mode from jvspatial.db.factory import create_database, get_current_database from jvspatial.db.manager import get_database_manager @@ -1450,69 +1444,20 @@ async def find_page( ``sort`` is expected to include at least one field; ``id`` is appended as a deterministic tiebreaker when missing. """ - page_limit = max(1, int(limit)) - if not sort: - sort = [("id", 1)] - sort_fields: List[Tuple[str, int]] = list(sort) - if not any(field == "id" for field, _ in sort_fields): - sort_fields.append(("id", sort_fields[0][1])) + from .pager import ( + decode_keyset_cursor, + encode_keyset_cursor, + keyset_filter, + keyset_sort_fields, + ) + page_limit = max(1, int(limit)) + sort_fields = keyset_sort_fields(sort) final_query: Dict[str, Any] = dict(query or {}) - cursor_payload: Optional[Dict[str, Any]] = None - if isinstance(after, str) and after: - try: - cursor_payload = json.loads( - base64.urlsafe_b64decode(after.encode()).decode() - ) - except Exception: - cursor_payload = None - elif isinstance(after, dict): - cursor_payload = after - - primary_field, primary_dir = sort_fields[0] - id_dir = sort_fields[-1][1] - if cursor_payload and "id" in cursor_payload and "sort" in cursor_payload: - sort_op = "$lt" if primary_dir < 0 else "$gt" - id_op = "$lt" if id_dir < 0 else "$gt" - cursor_sort = cursor_payload["sort"] - keyset_branches: List[Dict[str, Any]] - if cursor_sort is None: - # The cursor sits in the trailing run of records that have - # no value for the sort field. Records missing the sort - # field sort last in both directions (see - # ``finalize_find_results``), so everything still ahead of - # us is also missing it — walk that run by id alone. - keyset_branches = [ - { - "$and": [ - {primary_field: None}, - {"id": {id_op: cursor_payload["id"]}}, - ] - } - ] - else: - keyset_branches = [ - {primary_field: {sort_op: cursor_sort}}, - # ``{field: None}`` matches both an explicit null and a - # missing key. Without this branch the nulls-last tail - # is unreachable: ``{field: {"$lt": v}}`` never matches - # a record that has no value at all, so iteration would - # stop at the last record that does. - {primary_field: None}, - { - "$and": [ - {primary_field: cursor_sort}, - {"id": {id_op: cursor_payload["id"]}}, - ] - }, - ] - keyset_filter: Dict[str, Any] = ( - keyset_branches[0] - if len(keyset_branches) == 1 - else {"$or": keyset_branches} - ) + after_cursor = keyset_filter(sort_fields, decode_keyset_cursor(after)) + if after_cursor is not None: final_query = ( - {"$and": [final_query, keyset_filter]} if final_query else keyset_filter + {"$and": [final_query, after_cursor]} if final_query else after_cursor ) rows = await self.database.find( @@ -1520,21 +1465,11 @@ async def find_page( ) has_more = len(rows) > page_limit page_rows = rows[:page_limit] - - next_cursor: Optional[str] = None - if has_more and page_rows: - last = page_rows[-1] - # Dotted sort fields (``context.started_at``) need the same - # path walk the adapters use; a flat ``.get`` would mint a - # ``None`` sort value for every cursor and stall pagination. - payload = { - "id": last.get("id"), - "sort": resolve_sort_value(last, primary_field), - } - next_cursor = base64.urlsafe_b64encode( - json.dumps(payload, separators=(",", ":")).encode() - ).decode() - + next_cursor = ( + encode_keyset_cursor(page_rows[-1], sort_fields) + if has_more and page_rows + else None + ) return page_rows, next_cursor async def nodes_bulk( @@ -1547,8 +1482,15 @@ async def nodes_bulk( edge_filter: Optional[Dict[str, Any]] = None, node_filter: Optional[Dict[str, Any]] = None, limit: Optional[int] = None, + limit_per_source: Optional[int] = None, ) -> Dict[str, List[Any]]: - """Batch traversal for many source IDs in one edge-query pass.""" + """Batch traversal for many source IDs in one edge-query pass. + + ``limit_per_source`` caps the neighbours returned per source id. On + backends with ``find_connected_nodes_bulk`` (Postgres) an ``out`` / + ``in`` traversal without the legacy global ``limit`` is one round trip + (``ROW_NUMBER() OVER (PARTITION BY source)``). + """ from .entities.edge import Edge from .entities.node import Node @@ -1571,6 +1513,36 @@ async def nodes_bulk( if direction not in ("out", "in", "both"): direction = "out" + bulk = getattr(self.database, "find_connected_nodes_bulk", None) + if callable(bulk) and limit is None and direction in ("out", "in"): + try: + rows_by_source = await bulk( + self._get_collection_name("n"), + self._get_collection_name("e"), + unique_ids, + direction=direction, + edge_entities=edge_entities or None, + node_entities=node_entities or None, + edge_query={ + f"context.{k}": v for k, v in (edge_filter or {}).items() + } + or None, + node_query={ + f"context.{k}": v for k, v in (node_filter or {}).items() + } + or None, + limit_per_source=limit_per_source, + ) + except NotImplementedError: + rows_by_source = None + if rows_by_source is not None: + for src_id, rows in rows_by_source.items(): + for row in rows: + obj = await self._deserialize_entity(Node, row) + if obj is not None: + out.setdefault(src_id, []).append(obj) + return out + edge_query: Dict[str, Any] = {} if direction == "out": edge_query["source"] = {"$in": unique_ids} @@ -1637,6 +1609,8 @@ async def nodes_bulk( if obj is not None: node_by_id[doc["id"]] = obj + caps = [c for c in (limit, limit_per_source) if c is not None] + cap = max(1, min(caps)) if caps else None for src_id, target_ids in links.items(): matched: List[Any] = [] for tid in target_ids: @@ -1644,7 +1618,7 @@ async def nodes_bulk( if obj is None: continue matched.append(obj) - if limit is not None and len(matched) >= max(1, limit): + if cap is not None and len(matched) >= cap: break out[src_id] = matched return out diff --git a/jvspatial/core/entities/node.py b/jvspatial/core/entities/node.py index 75c37b6..1d987ed 100644 --- a/jvspatial/core/entities/node.py +++ b/jvspatial/core/entities/node.py @@ -1,7 +1,6 @@ """Node class for jvspatial graph entities.""" import logging -import re import weakref from typing import ( TYPE_CHECKING, @@ -11,10 +10,13 @@ Dict, List, Optional, + Tuple, Type, Union, ) +from jvspatial.db.database import finalize_find_results + from ..annotations import attribute from .edge import Edge from .object import Object @@ -27,6 +29,124 @@ logger = logging.getLogger(__name__) +# Record paths at the top level of node / edge documents. Any other bare +# property name in a neighbour filter refers to ``context.`` — the same +# attribute ``_matches_property_filter`` reads. +_NODE_TOP_LEVEL_KEYS = frozenset({"id", "entity"}) +_EDGE_TOP_LEVEL_KEYS = frozenset({"id", "entity", "source", "target", "bidirectional"}) + +# (edge_entities, edge_query, node_entities, node_query) — ``None`` entities +# mean "any type", ``None`` queries mean "no property filter". +NeighborSpec = Tuple[ + Optional[List[str]], + Optional[Dict[str, Any]], + Optional[List[str]], + Optional[Dict[str, Any]], +] + + +def _record_query(criteria: Dict[str, Any], top_level: frozenset) -> Dict[str, Any]: + """Map attribute-style criteria (``population``) to record paths (``context.population``).""" + return { + ( + key + if key.startswith(("context.", "$")) or key in top_level + else f"context.{key}" + ): value + for key, value in criteria.items() + } + + +def _entity_name_of(cls: type) -> str: + resolver = getattr(cls, "_entity_name", None) + return resolver() if callable(resolver) else cls.__name__ + + +def _node_class_entities(cls: type) -> Optional[List[str]]: + """Entity names ``isinstance(x, cls)`` accepts: ``cls`` and every loaded subclass. + + ``None`` for the ``Node`` base itself, which every node matches. + """ + if cls is Node: + return None + names = set() + stack: List[type] = [cls] + seen: set = set() + while stack: + klass = stack.pop() + if klass in seen: + continue + seen.add(klass) + names.add(_entity_name_of(klass)) + stack.extend(klass.__subclasses__()) + return sorted(names) + + +def _neighbor_filter_spec( + flt: Any, *, node_side: bool +) -> Tuple[Optional[List[str]], Optional[Dict[str, Any]]]: + """Normalise a ``nodes()`` node / edge filter to ``(entities, query)``. + + Accepts a name, a class, or a list of names, classes and + ``{name: criteria}`` dicts. Node classes match subclasses (``isinstance`` + semantics); edge classes, names and dict keys match exactly. Criteria + dicts become an ``$or`` of ``{"entity": name, }`` + branches, alongside ``{"entity": {"$in": [...]}}`` for plain items. + """ + if flt is None: + return None, None + items = list(flt) if isinstance(flt, (list, tuple)) else [flt] + if not items and not node_side: + return None, None # ``edge=[]`` has always meant "any edge type" + top_level = _NODE_TOP_LEVEL_KEYS if node_side else _EDGE_TOP_LEVEL_KEYS + entities: List[str] = [] + branches: List[Dict[str, Any]] = [] + for item in items: + if isinstance(item, str): + entities.append(item) + elif isinstance(item, type): + if not node_side: + entities.append(_entity_name_of(item)) + continue + names = _node_class_entities(item) + if names is None: + return None, None + entities.extend(names) + elif isinstance(item, dict): + for name, criteria in item.items(): + entity = name if isinstance(name, str) else _entity_name_of(name) + branches.append( + {"entity": entity, **_record_query(dict(criteria or {}), top_level)} + ) + else: + side = "node" if node_side else "edge" + raise TypeError(f"unsupported {side} filter item: {item!r}") + entities = sorted(set(entities)) + if not branches: + return entities, None + if entities: + branches.append({"entity": {"$in": entities}}) + return None, {"$or": branches} + + +def _deserialize_hint(node_filter: Any) -> type: + """Class to hydrate neighbours as when the filter names exactly one class. + + ``_deserialize_entity`` resolves the stored ``entity`` against this + class's subtree first, so when two Node subclasses share an entity name + (an app ``User`` and an embedded-agent ``User``) the row hydrates as the + class the caller asked for rather than the first global name match. + """ + if isinstance(node_filter, type): + return node_filter + if ( + isinstance(node_filter, (list, tuple)) + and len(node_filter) == 1 + and isinstance(node_filter[0], type) + ): + return node_filter[0] + return Node + class Node(Object): """Graph node with visitor tracking and connection capabilities. @@ -419,11 +539,13 @@ async def nodes_bulk( edge_filter: Optional[Dict[str, Any]] = None, node_filter: Optional[Dict[str, Any]] = None, limit: Optional[int] = None, + limit_per_source: Optional[int] = None, ) -> Dict[str, List["Node"]]: """Batch version of ``nodes()`` for many source IDs. Returns a mapping of ``source_id -> connected nodes`` using a single - edge-query pass in the active GraphContext. + edge-query pass in the active GraphContext. ``limit_per_source`` caps + the neighbours per source (one windowed query on Postgres). """ from ..context import get_default_context @@ -436,6 +558,7 @@ async def nodes_bulk( edge_filter=edge_filter, node_filter=node_filter, limit=limit, + limit_per_source=limit_per_source, ) async def neighborhood( @@ -541,72 +664,122 @@ async def count_neighbors( ] = None, **kwargs: Any, ) -> int: - """Count neighbors (same filters as :meth:`nodes`). + """Count neighbors (same filters as :meth:`nodes`) — alias of :meth:`count_nodes`. Named ``count_neighbors`` so this does not shadow :meth:`Object.count` on Node subclasses (e.g. ``User.count(query)`` remains the DB count API). - Fast path: when ``node`` is a single entity name/class and no ``edge`` - filter or extra ``kwargs`` are provided, this issues ``count`` queries on - the ``edge`` collection using ``source`` / ``target`` plus a regex on the - peer node id (pattern ``^n..``). Persisted edges do not store - separate ``target_entity`` / ``source_entity`` fields. - Returns: Number of matching connected nodes. """ - # Fast-path: single entity type filter, no edge filter or property kwargs. - if ( - not kwargs - and edge is None - and node is not None - and not isinstance(node, list) - ): - entity_name: Optional[str] = None - if isinstance(node, str): - entity_name = node - elif isinstance(node, type): - # Honor ``__entity_name__`` so subclasses with custom - # discriminators match their persisted ID prefix. - resolver = getattr(node, "_entity_name", None) - entity_name = resolver() if callable(resolver) else node.__name__ - if entity_name is not None: - try: - from ..context import get_default_context - - ctx = get_default_context() - db = ctx.database - node_type_re = { - "$regex": {"pattern": rf"^n\.{re.escape(entity_name)}\."} - } - if direction in ("out", "both"): - q_out: Dict[str, Any] = { - "source": self.id, - "target": node_type_re, - } - out_count = await db.count("edge", q_out) - else: - out_count = 0 - if direction in ("in", "both"): - q_in: Dict[str, Any] = { - "target": self.id, - "source": node_type_re, - } - in_count = await db.count("edge", q_in) - else: - in_count = 0 - return out_count + in_count - except Exception: - pass # Fall through to full hydration on any error. + return await self.count_nodes( + direction=direction, node=node, edge=edge, **kwargs + ) - return len( - await self.nodes( - direction=direction, - node=node, - edge=edge, - **kwargs, - ) + async def count_nodes( + self, + direction: str = "out", + node: Optional[ + Union[str, type, List[Union[str, type, Dict[str, Dict[str, Any]]]]] + ] = None, + edge: Optional[ + Union[ + str, + Type["Edge"], + List[Union[str, Type["Edge"], Dict[str, Dict[str, Any]]]], + ] + ] = None, + **kwargs: Any, + ) -> int: + """Count neighbours matching the same filters as :meth:`nodes`. + + One ``COUNT`` round trip on backends with ``count_connected_nodes`` + (Postgres, MongoDB, SQLite); elsewhere the neighbours are listed and + counted. Use this instead of ``len(await node.nodes(...))``, which + hydrates every neighbour. + + Returns: + Number of distinct matching neighbours. + """ + context = await self.get_context() + spec = self._neighbor_spec(node, edge, kwargs) + counter = getattr(context.database, "count_connected_nodes", None) + if callable(counter): + edge_entities, edge_query, node_entities, node_query = spec + try: + return int( + await counter( + context._get_collection_name("n"), + context._get_collection_name("e"), + self.id, + direction=direction, + edge_entities=edge_entities, + node_entities=node_entities, + edge_query=edge_query, + node_query=node_query, + ) + ) + except NotImplementedError as exc: + logger.debug("count_nodes(): no pushdown (%s); listing instead", exc) + return len(await self._neighbors(context, direction=direction, spec=spec)) + + async def nodes_page( + self, + *, + direction: str = "out", + node: Optional[ + Union[str, type, List[Union[str, type, Dict[str, Dict[str, Any]]]]] + ] = None, + edge: Optional[ + Union[ + str, + Type["Edge"], + List[Union[str, Type["Edge"], Dict[str, Dict[str, Any]]]], + ] + ] = None, + sort: Optional[List[Tuple[str, int]]] = None, + cursor: Optional[str] = None, + limit: int = 20, + **kwargs: Any, + ) -> Tuple[List["Node"], Optional[str]]: + """One keyset-paginated page of neighbours: ``(nodes, next_cursor)``. + + Same filters as :meth:`nodes`. ``sort`` takes record paths + (``[("context.created_at", -1)]``; default ``[("id", 1)]``, with ``id`` + appended as tiebreaker). Pass ``next_cursor`` back as ``cursor`` for + the next page; ``None`` means there are no more. Neighbours inserted + before the cursor never shift later pages. Cursors use the same + opaque encoding as :meth:`GraphContext.find_page`. + """ + from ..pager import ( + decode_keyset_cursor, + encode_keyset_cursor, + keyset_filter, + keyset_sort_fields, + ) + + context = await self.get_context() + page_limit = max(1, int(limit)) + sort_fields = keyset_sort_fields(sort) + edge_entities, edge_query, node_entities, node_query = self._neighbor_spec( + node, edge, kwargs + ) + after = keyset_filter(sort_fields, decode_keyset_cursor(cursor)) + if after is not None: + node_query = after if node_query is None else {"$and": [node_query, after]} + found = await self._neighbors( + context, + direction=direction, + spec=(edge_entities, edge_query, node_entities, node_query), + deserialize_as=_deserialize_hint(node), + sort=sort_fields, + limit=page_limit + 1, ) + page = found[:page_limit] + next_cursor = None + if len(found) > page_limit and page: + next_cursor = encode_keyset_cursor(await page[-1].export(), sort_fields) + return page, next_cursor async def node( self, @@ -678,7 +851,13 @@ async def _node_query( limit: Optional[int] = None, **kwargs: Any, ) -> List["Node"]: - """Execute optimized database query to find connected nodes. + """Find connected nodes matching the node / edge filters. + + Every filter shape is normalised to entity lists plus record-path + queries (:func:`_neighbor_filter_spec`) and pushed to the backend's + ``find_connected_nodes`` in one round trip — ``limit`` included — when + it has one. Otherwise (or when a criterion does not translate) the + Python path still applies every filter; nothing is dropped. Args: context: GraphContext instance for database operations @@ -689,207 +868,162 @@ async def _node_query( **kwargs: Simple property filters for connected nodes Returns: - List of connected nodes matching the criteria + List of connected nodes matching the criteria (each at most once) """ - from .node import Node as NodeClass - - if ( - not kwargs - and direction in ("out", "in") - and not isinstance(edge_filter, list) - ): - db = context.database - find_connected = getattr(db, "find_connected_nodes", None) - edge_entity: Optional[str] - if isinstance(edge_filter, type): - from .edge import Edge as EdgeClassForEntity - - resolver = getattr(edge_filter, "_entity_name", None) - edge_entity = resolver() if callable(resolver) else edge_filter.__name__ - edge_cls = edge_filter - elif isinstance(edge_filter, str): - edge_entity = edge_filter - from .edge import Edge as EdgeClassForEntity - - edge_cls = EdgeClassForEntity - elif edge_filter is None: - from .edge import Edge as EdgeClassForEntity - - edge_entity = None - edge_cls = EdgeClassForEntity - else: - edge_entity = "__skip_fast_path__" + return await self._neighbors( + context, + direction=direction, + spec=self._neighbor_spec(node_filter, edge_filter, kwargs), + deserialize_as=_deserialize_hint(node_filter), + limit=limit, + ) - if callable(find_connected) and edge_entity != "__skip_fast_path__": - node_coll = context._get_collection_name( - context._get_entity_type_code(NodeClass) - ) - edge_coll = context._get_collection_name( - context._get_entity_type_code(edge_cls) - ) - # Only push ``limit`` into the DB scan when there is no node - # filter. With a node filter the type match happens in Python - # below, so a DB-side limit would truncate the candidate set - # BEFORE filtering (e.g. limit=1 fetches one neighbor that may - # not match the type, yielding an empty result even though - # matching neighbors exist). Fetch all, filter, then slice. - db_limit = limit if node_filter is None else None + @staticmethod + def _neighbor_spec( + node_filter: Any, edge_filter: Any, properties: Dict[str, Any] + ) -> NeighborSpec: + """Normalise ``nodes()``-style filters plus property kwargs to a spec.""" + edge_entities, edge_query = _neighbor_filter_spec(edge_filter, node_side=False) + node_entities, node_query = _neighbor_filter_spec(node_filter, node_side=True) + if properties: + props = _record_query(properties, _NODE_TOP_LEVEL_KEYS) + node_query = props if node_query is None else {"$and": [node_query, props]} + return edge_entities, edge_query, node_entities, node_query + + async def _neighbors( + self, + context: "GraphContext", + *, + direction: str, + spec: NeighborSpec, + deserialize_as: Optional[type] = None, + sort: Optional[List[Tuple[str, int]]] = None, + limit: Optional[int] = None, + ) -> List["Node"]: + """Neighbours for a normalised spec — pushed down when the backend can.""" + if direction not in ("out", "in", "both"): + raise ValueError( + f"direction must be 'out', 'in' or 'both', got {direction!r}" + ) + hint = deserialize_as or Node + edge_entities, edge_query, node_entities, node_query = spec + records: Optional[List[Dict[str, Any]]] = None + find_connected = getattr(context.database, "find_connected_nodes", None) + if callable(find_connected): + try: records = await find_connected( - node_coll, - edge_coll, + context._get_collection_name("n"), + context._get_collection_name("e"), self.id, direction=direction, - edge_entity=edge_entity, - limit=db_limit, + edge_entities=edge_entities, + node_entities=node_entities, + edge_query=edge_query, + node_query=node_query, + sort=sort, + limit=limit, ) - # Deserialize with the caller's requested concrete type as the - # hint (not the base ``Node``) when the node filter names one. - # ``_deserialize_entity`` resolves the stored ``entity`` name - # against this class's subtree first, so when two Node - # subclasses across embedded graphs share an entity name - # (e.g. an app ``User`` and an embedded-agent ``User``), the - # row hydrates as the class the caller asked for instead of the - # first global name match — otherwise ``_matches_node_filter``'s - # ``isinstance`` check silently drops a valid neighbor. - deser_class: type = NodeClass - if isinstance(node_filter, type): - deser_class = node_filter - elif ( - isinstance(node_filter, list) - and len(node_filter) == 1 - and isinstance(node_filter[0], type) - ): - deser_class = node_filter[0] - - connected_nodes: List["Node"] = [] - for data in records: - try: - node_obj: Optional["Node"] = await context._deserialize_entity( - deser_class, data - ) - if node_obj: - connected_nodes.append(node_obj) - await context._add_to_cache(node_obj.id, node_obj) - except Exception as e: - logger.debug( - "Skipping invalid node during connected-node join: %s", e - ) - continue - if node_filter is not None: - connected_nodes = [ - n - for n in connected_nodes - if self._matches_node_filter(n, node_filter) - ] - if limit is not None: - connected_nodes = connected_nodes[:limit] - return connected_nodes - - # Find edges connected to this node - from .edge import Edge as EdgeClass - - edges = [] - edge_cls = edge_filter if isinstance(edge_filter, type) else EdgeClass + except NotImplementedError as exc: + logger.debug("nodes(): no pushdown (%s); Python path", exc) + if records is None: + return await self._neighbors_fallback( + context, + direction=direction, + spec=spec, + deserialize_as=hint, + sort=sort, + limit=limit, + ) + return await self._hydrate_neighbors(context, records, hint) - # Optimization: Use single combined query for bidirectional traversal - if direction == "both": - # Build combined query for both directions - edge_query: Dict[str, Any] = { - "$or": [{"source": self.id}, {"target": self.id}] - } - if isinstance(edge_filter, type): - # Persisted discriminator field is ``entity``; honor - # ``__entity_name__`` override (SPEC §1.2). - resolver = getattr(edge_filter, "_entity_name", None) - edge_query["entity"] = ( - resolver() if callable(resolver) else edge_filter.__name__ + async def _hydrate_neighbors( + self, + context: "GraphContext", + records: List[Dict[str, Any]], + deserialize_as: type, + ) -> List["Node"]: + found: List["Node"] = [] + for data in records: + try: + node_obj: Optional["Node"] = await context._deserialize_entity( + deserialize_as, data ) + except Exception as e: + logger.debug("Skipping invalid neighbour record: %s", e) + continue + if node_obj: + found.append(node_obj) + await context._add_to_cache(node_obj.id, node_obj) + return found - # Cap edge fan-out so hub nodes don't accidentally load unbounded sets. - edge_results = await context.database.find( - "edge", edge_query, limit=limit if limit is not None else 10000 - ) - for edge_data in edge_results: - try: - edge_obj: Optional["Edge"] = await context._deserialize_entity( - edge_cls, edge_data - ) - if edge_obj: - edges.append(edge_obj) - except Exception as e: - logger.debug( - f"Skipping invalid edge during bidirectional traversal: {e}" - ) - continue - else: - # Single direction queries - if direction == "out": - # Find outgoing edges - outgoing_edges = await context.find_edges_between( - source_id=self.id, - edge_class=edge_filter if isinstance(edge_filter, type) else None, - ) - edges.extend(outgoing_edges) + async def _neighbors_fallback( + self, + context: "GraphContext", + *, + direction: str, + spec: NeighborSpec, + deserialize_as: type, + sort: Optional[List[Tuple[str, int]]], + limit: Optional[int], + ) -> List["Node"]: + """Python traversal for backends without ``find_connected_nodes``. - elif direction == "in": - # Find incoming edges (where this node is the target) - query: Dict[str, Any] = {"target": self.id} - if isinstance(edge_filter, type): - resolver = getattr(edge_filter, "_entity_name", None) - query["entity"] = ( - resolver() if callable(resolver) else edge_filter.__name__ - ) + One edge ``find`` (endpoint + entity ``$in`` + edge criteria), then the + neighbours — filtered by entity and criteria in the node ``find`` + before hydration. + """ + edge_entities, edge_query, node_entities, node_query = spec + db = context.database + endpoints: Dict[str, Dict[str, Any]] = { + "out": {"source": self.id}, + "in": {"target": self.id}, + "both": {"$or": [{"source": self.id}, {"target": self.id}]}, + } + edge_parts: List[Dict[str, Any]] = [endpoints[direction]] + if edge_entities is not None: + edge_parts.append({"entity": {"$in": edge_entities}}) + if edge_query: + edge_parts.append(edge_query) + edge_docs = await db.find( + context._get_collection_name("e"), + edge_parts[0] if len(edge_parts) == 1 else {"$and": edge_parts}, + ) + neighbor_ids: List[str] = [] + for doc in edge_docs: + src, tgt = doc.get("source"), doc.get("target") + if direction in ("out", "both") and src == self.id and tgt: + neighbor_ids.append(tgt) + if direction in ("in", "both") and tgt == self.id and src: + neighbor_ids.append(src) + neighbor_ids = list(dict.fromkeys(neighbor_ids)) + if not neighbor_ids: + return [] - edge_results = await context.database.find("edge", query) - for edge_data in edge_results: - try: - edge_obj = await context._deserialize_entity( - edge_cls, edge_data - ) - if edge_obj: - edges.append(edge_obj) - except Exception as e: - logger.debug( - f"Skipping invalid incoming edge during traversal: {e}" - ) - continue + if node_entities is None and node_query is None and not sort: + # Unfiltered: the cache-aware batch fetch (identity-map friendly). + found = await context.get_batch(Node, neighbor_ids) + return found if limit is None else found[:limit] - # Get unique connected node IDs - connected_node_ids = set() - for edge in edges: - if direction in ["out", "both"] and hasattr(edge, "target"): - connected_node_ids.add(edge.target) - if direction in ["in", "both"] and hasattr(edge, "source"): - connected_node_ids.add(edge.source) - - # Find the actual nodes using batch retrieval for efficiency - # This uses context.get_batch() which handles caching and batch queries - if connected_node_ids: - connected_nodes = await context.get_batch(Node, list(connected_node_ids)) + records: List[Dict[str, Any]] = [] + for off in range(0, len(neighbor_ids), 500): + node_parts: List[Dict[str, Any]] = [ + {"id": {"$in": neighbor_ids[off : off + 500]}} + ] + if node_entities is not None: + node_parts.append({"entity": {"$in": node_entities}}) + if node_query: + node_parts.append(node_query) + records.extend( + await db.find(context._get_collection_name("n"), {"$and": node_parts}) + ) + if sort: + records = finalize_find_results(records, sort=sort, limit=limit) else: - connected_nodes = [] - - # Apply node type filtering - if node_filter is not None: - filtered_nodes = [] - for node_obj in connected_nodes: - if self._matches_node_filter(node_obj, node_filter): - filtered_nodes.append(node_obj) - connected_nodes = filtered_nodes - - # Apply property filtering from kwargs - if kwargs: - filtered_nodes = [] - for node_obj in connected_nodes: - if self._matches_property_filter(node_obj, kwargs): - filtered_nodes.append(node_obj) - connected_nodes = filtered_nodes - - # Apply limit - if limit is not None: - connected_nodes = connected_nodes[:limit] - - return connected_nodes + order = {nid: i for i, nid in enumerate(neighbor_ids)} + records.sort(key=lambda r: order.get(str(r.get("id")), 0)) + if limit is not None: + records = records[:limit] + return await self._hydrate_neighbors(context, records, deserialize_as) def _matches_node_filter( self, diff --git a/jvspatial/core/pager.py b/jvspatial/core/pager.py index b8fb74e..64d1a71 100644 --- a/jvspatial/core/pager.py +++ b/jvspatial/core/pager.py @@ -5,10 +5,12 @@ Designed to integrate seamlessly with UI frameworks requiring paginated data. """ +import base64 +import json from math import ceil -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type, TypeVar +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Type, TypeVar, Union -from jvspatial.db.database import finalize_find_results +from jvspatial.db.database import finalize_find_results, resolve_sort_value if TYPE_CHECKING: from .entities import Object @@ -357,3 +359,94 @@ async def paginate_by_field( ) return await pager.get_page(page) + + +# ---- keyset cursors --------------------------------------------------------- +# Shared by ``GraphContext.find_page`` and ``Node.nodes_page`` so both mint and +# read the same opaque cursor: base64(json({"id": , "sort": })). + + +def keyset_sort_fields( + sort: Optional[List[Tuple[str, int]]], +) -> List[Tuple[str, int]]: + """Normalise a keyset sort: default ``[("id", 1)]``, ``id`` appended as tiebreaker.""" + fields: List[Tuple[str, int]] = list(sort) if sort else [("id", 1)] + if not any(field == "id" for field, _ in fields): + fields.append(("id", fields[0][1])) + return fields + + +def encode_keyset_cursor( + record: Dict[str, Any], sort_fields: List[Tuple[str, int]] +) -> str: + """Opaque cursor positioned after ``record`` for a sort led by ``sort_fields[0]``. + + Dotted sort fields (``context.started_at``) resolve through the same path + walk the adapters use; a flat ``.get`` would mint a ``None`` sort value + for every cursor and stall pagination. + """ + payload = { + "id": record.get("id"), + "sort": resolve_sort_value(record, sort_fields[0][0]), + } + return base64.urlsafe_b64encode( + json.dumps(payload, separators=(",", ":")).encode() + ).decode() + + +def decode_keyset_cursor( + cursor: Optional[Union[str, Dict[str, Any]]], +) -> Optional[Dict[str, Any]]: + """Decode a cursor from :func:`encode_keyset_cursor`; ``None`` if absent or malformed.""" + if isinstance(cursor, dict): + return cursor + if not isinstance(cursor, str) or not cursor: + return None + try: + payload = json.loads(base64.urlsafe_b64decode(cursor.encode()).decode()) + except Exception: + return None + return payload if isinstance(payload, dict) else None + + +def keyset_filter( + sort_fields: List[Tuple[str, int]], + cursor_payload: Optional[Dict[str, Any]], +) -> Optional[Dict[str, Any]]: + """Mongo-style filter selecting records strictly after the cursor. + + Honors the nulls-last contract of ``finalize_find_results``: records + with no value for the primary sort field come after every valued record + in both directions. ``None`` when there is no (usable) cursor. + """ + if not cursor_payload or "id" not in cursor_payload or "sort" not in cursor_payload: + return None + primary_field, primary_dir = sort_fields[0] + id_dir = sort_fields[-1][1] + sort_op = "$lt" if primary_dir < 0 else "$gt" + id_op = "$lt" if id_dir < 0 else "$gt" + cursor_sort = cursor_payload["sort"] + if cursor_sort is None: + # The cursor sits in the trailing run of records that have no value + # for the sort field; everything still ahead is also missing it — + # walk that run by id alone. + return { + "$and": [ + {primary_field: None}, + {"id": {id_op: cursor_payload["id"]}}, + ] + } + return { + "$or": [ + {primary_field: {sort_op: cursor_sort}}, + # ``{field: None}`` matches both an explicit null and a missing + # key. Without this branch the nulls-last tail is unreachable. + {primary_field: None}, + { + "$and": [ + {primary_field: cursor_sort}, + {"id": {id_op: cursor_payload["id"]}}, + ] + }, + ] + } diff --git a/jvspatial/db/README.md b/jvspatial/db/README.md index d7fcffd..911bbf2 100644 --- a/jvspatial/db/README.md +++ b/jvspatial/db/README.md @@ -62,7 +62,8 @@ db/ - **`Database.edge_ids_mode` says where node adjacency lives.** `"derive"` (Postgres, MongoDB, SQLite): the edge collection only; node rows carry no `edges`. `"persist"` (JSON, DynamoDB): node rows also store `edges`. Read the effective value with `resolve_edge_ids_mode(db)` — instance attribute → `JVSPATIAL_NODE_EDGE_IDS` → class default. `strip_node_edges()` (Postgres, MongoDB) removes legacy arrays; CLI `jvspatial migrate strip-node-edges`. - **`find_many` and `bulk_save` are public and benefit from native overrides.** Defaults exist but are slow. (`database.py:176+`) - **`find_one_and_update` / `find_one_and_delete` are NOT atomic by default.** MongoDB and Postgres override with native atomic versions (`FOR UPDATE` on Postgres). -- **Postgres-only helpers** (via `getattr`, not on the ABC): `traverse`, `find_connected_nodes`, `save_with_edge_merge` (persist mode only). +- **Neighbour pushdown** (via `getattr`, not on the ABC): `find_connected_nodes` / `count_connected_nodes` on Postgres, MongoDB and SQLite take `edge_entities`, `node_entities`, `edge_query`, `node_query` (record paths), `sort` and `limit`, and raise `NotImplementedError` for anything they cannot translate so `Node.nodes()` falls back without dropping a filter. Postgres adds `find_connected_nodes_bulk(limit_per_source=...)`. +- **Postgres-only helpers** (via `getattr`, not on the ABC): `traverse`, `find_connected_nodes_bulk`, `save_with_edge_merge` (persist mode only). - **Atomic JSON writes use `temp + fsync + rename + fsync(dir)`.** No partial records survive a crash. (`_atomic.py`) - **Per-file locks serialize concurrent writes to the same record only.** Different files run in parallel. (`_path_locks.py`) - **`QueryEngine` LRU is bounded.** Default 1024; configurable. Unbounded query construction will not leak memory. (`query.py`) diff --git a/jvspatial/db/_observable.py b/jvspatial/db/_observable.py index 964a4dc..68eeb59 100644 --- a/jvspatial/db/_observable.py +++ b/jvspatial/db/_observable.py @@ -335,60 +335,53 @@ async def drop_deprecated_indexes(self, deprecated: Dict[str, List[str]]) -> Non """Pass through deprecated-index cleanup to the wrapped backend.""" await self.inner.drop_deprecated_indexes(deprecated) - async def find_connected_nodes( - self, - node_collection: str, - edge_collection: str, - node_id: str, - *, - direction: str = "out", - edge_entity: Optional[str] = None, - limit: Optional[int] = None, - ) -> List[Dict[str, Any]]: - """Instrumented single-hop neighbor join when the backend supports it.""" - inner = getattr(self.inner, "find_connected_nodes", None) + # Optional graph ops exist on the wrapper only when the wrapped backend + # implements them: the properties raise ``AttributeError`` otherwise, so + # ``getattr(db, "find_connected_nodes", None)`` stays ``None`` for + # adapters without a pushdown and callers take their fallback path. + + def _optional_op( + self, name: str, result_count: Callable[[Any], int] + ) -> Callable[..., Awaitable[Any]]: + inner = getattr(self.inner, name, None) if not callable(inner): - raise AttributeError("find_connected_nodes") - return await self._instrument( - "find_connected_nodes", - node_collection, - lambda: inner( - node_collection, - edge_collection, - node_id, - direction=direction, - edge_entity=edge_entity, - limit=limit, - ), - result_count_extractor=lambda r: len(r) if isinstance(r, list) else 0, + raise AttributeError(name) + + async def _call(collection: str, *args: Any, **kwargs: Any) -> Any: + return await self._instrument( + name, + collection, + lambda: inner(collection, *args, **kwargs), + result_count_extractor=result_count, + ) + + return _call + + @property + def find_connected_nodes(self) -> Callable[..., Awaitable[Any]]: + """Instrumented single-hop neighbour query (backend permitting).""" + return self._optional_op( + "find_connected_nodes", lambda r: len(r) if isinstance(r, list) else 0 ) - async def traverse( - self, - edge_collection: str, - start_id: str, - *, - direction: str = "out", - max_depth: int = 1, - edge_filter: Optional[Dict[str, Any]] = None, - limit: Optional[int] = None, - ) -> List[Dict[str, Any]]: - """Instrumented multi-hop graph walk when the backend supports it.""" - inner = getattr(self.inner, "traverse", None) - if not callable(inner): - raise AttributeError("traverse") - return await self._instrument( - "traverse", - edge_collection, - lambda: inner( - edge_collection, - start_id, - direction=direction, - max_depth=max_depth, - edge_filter=edge_filter, - limit=limit, - ), - result_count_extractor=lambda r: len(r) if isinstance(r, list) else 0, + @property + def count_connected_nodes(self) -> Callable[..., Awaitable[Any]]: + """Instrumented single-hop neighbour count (backend permitting).""" + return self._optional_op("count_connected_nodes", lambda r: int(r or 0)) + + @property + def find_connected_nodes_bulk(self) -> Callable[..., Awaitable[Any]]: + """Instrumented multi-source neighbour query (backend permitting).""" + return self._optional_op( + "find_connected_nodes_bulk", + lambda r: sum(len(v) for v in r.values()) if isinstance(r, dict) else 0, + ) + + @property + def traverse(self) -> Callable[..., Awaitable[Any]]: + """Instrumented multi-hop graph walk (backend permitting).""" + return self._optional_op( + "traverse", lambda r: len(r) if isinstance(r, list) else 0 ) async def find_iter( diff --git a/jvspatial/db/_postgres_translate.py b/jvspatial/db/_postgres_translate.py index 447d081..e15611c 100644 --- a/jvspatial/db/_postgres_translate.py +++ b/jvspatial/db/_postgres_translate.py @@ -130,25 +130,29 @@ def values(self) -> List[Any]: def _translate_field_clause( - field: str, condition: Any, pb: ParamBuilder + field: str, condition: Any, pb: ParamBuilder, table: str = "" ) -> Optional[str]: """Translate ``{field: condition}`` to a SQL fragment. + ``table`` is an optional table alias (``"n"``) prefixed to every column + reference — used when the fragment lands in a join. + Returns ``None`` to signal fallback (the whole query should drop to in-Python evaluation). """ if not _safe_field_path(field): return None + prefix = f"{table}." if table else "" if field in _TOP_LEVEL_TEXT_COLUMNS: # Bare column reference — lets $eq/comparators hit the real btree # index (e.g. ``{col}_entity_idx``) instead of a JSONB scan. - extract_text = field - extract_jsonb = f"to_jsonb({field})" + extract_text = f"{prefix}{field}" + extract_jsonb = f"to_jsonb({prefix}{field})" else: path = _path_literal(field) - extract_text = f"(data #>> '{path}')" # text - extract_jsonb = f"(data #> '{path}')" # jsonb + extract_text = f"({prefix}data #>> '{path}')" # text + extract_jsonb = f"({prefix}data #> '{path}')" # jsonb # Plain equality with a scalar value. if not isinstance(condition, dict): @@ -252,7 +256,7 @@ def _translate_field_clause( # then negate the whole thing. if not isinstance(operand, dict): return None - inner = _translate_field_clause(field, operand, pb) + inner = _translate_field_clause(field, operand, pb, table) if inner is None: return None fragments.append(f"NOT ({inner})") @@ -506,14 +510,16 @@ def _to_jsonb_literal(value: Any) -> str: # ---- logical translation ---------------------------------------------------- -def _translate_logical(op: str, conditions: Any, pb: ParamBuilder) -> Optional[str]: +def _translate_logical( + op: str, conditions: Any, pb: ParamBuilder, table: str = "" +) -> Optional[str]: if not isinstance(conditions, list) or not conditions: return None parts: List[str] = [] for sub in conditions: if not isinstance(sub, dict): return None - translated = _translate_query_into(sub, pb) + translated = _translate_query_into(sub, pb, table) if translated is None: return None parts.append(f"({translated})") @@ -526,7 +532,9 @@ def _translate_logical(op: str, conditions: Any, pb: ParamBuilder) -> Optional[s return None -def _translate_query_into(query: Dict[str, Any], pb: ParamBuilder) -> Optional[str]: +def _translate_query_into( + query: Dict[str, Any], pb: ParamBuilder, table: str = "" +) -> Optional[str]: """Translate a query dict into a SQL fragment, accumulating into ``pb``.""" if not query: return "" @@ -536,7 +544,7 @@ def _translate_query_into(query: Dict[str, Any], pb: ParamBuilder) -> Optional[s if key in _IGNORED_TOP_LEVEL: continue if key in ("$and", "$or", "$nor"): - piece = _translate_logical(key, value, pb) + piece = _translate_logical(key, value, pb, table) if piece is None: return None fragments.append(f"({piece})") @@ -546,14 +554,14 @@ def _translate_query_into(query: Dict[str, Any], pb: ParamBuilder) -> Optional[s # dict. if not isinstance(value, dict): return None - inner = _translate_query_into(value, pb) + inner = _translate_query_into(value, pb, table) if inner is None: return None fragments.append(f"NOT ({inner})") continue if key.startswith("$"): return None - piece = _translate_field_clause(key, value, pb) + piece = _translate_field_clause(key, value, pb, table) if piece is None: return None fragments.append(piece) @@ -561,8 +569,18 @@ def _translate_query_into(query: Dict[str, Any], pb: ParamBuilder) -> Optional[s return " AND ".join(fragments) if fragments else "" +def _check_table_alias(table: Optional[str]) -> str: + if not table: + return "" + if not _SAFE_SEGMENT_RE.match(table): + raise ValueError(f"unsafe table alias: {table!r}") + return table + + def translate_query( query: Dict[str, Any], + *, + table: Optional[str] = None, ) -> Optional[Tuple[str, List[Any]]]: """Translate a Mongo-style query dict to ``(sql_where, params)``. @@ -571,10 +589,11 @@ def translate_query( The returned SQL fragment is meant to be ANDed into a larger WHERE clause: ``WHERE ()``. When the query is empty, returns - ``("", [])``. + ``("", [])``. ``table`` qualifies every column with a table alias + (``n.data``, ``n.entity``) so the fragment can target one side of a join. """ pb = ParamBuilder() - out = _translate_query_into(query, pb) + out = _translate_query_into(query, pb, _check_table_alias(table)) if out is None: return None return out, pb.values @@ -582,6 +601,8 @@ def translate_query( def translate_sort( sort: Optional[List[Tuple[str, int]]], + *, + table: Optional[str] = None, ) -> Optional[str]: """Translate a Mongo-style sort spec into an ORDER BY fragment. @@ -596,6 +617,8 @@ def translate_sort( """ if not sort: return None + alias = _check_table_alias(table) + prefix = f"{alias}." if alias else "" parts: List[str] = [] for field, direction in sort: if direction not in (1, -1): @@ -603,7 +626,7 @@ def translate_sort( if not _safe_field_path(field): return None path = _path_literal(field) - col = f"(data #>> '{path}')" + col = f"({prefix}data #>> '{path}')" if direction == 1: parts.append(f"{col} ASC NULLS LAST") else: diff --git a/jvspatial/db/_sqlite_translate.py b/jvspatial/db/_sqlite_translate.py index ea48ac0..0d878ee 100644 --- a/jvspatial/db/_sqlite_translate.py +++ b/jvspatial/db/_sqlite_translate.py @@ -90,12 +90,22 @@ def _safe_field_path(field: str) -> bool: return all(_SAFE_SEGMENT_RE.match(seg) for seg in field.split(".")) -def _json_extract(field: str) -> str: +def _json_extract(field: str, table: str = "") -> str: """Return the SQL fragment for ``json_extract(data, '$.field.path')``. Caller must have already verified the path with :func:`_safe_field_path`. + ``table`` optionally qualifies the column (``e.data``) for joins. """ - return f"json_extract(data, '$.{field}')" + prefix = f"{table}." if table else "" + return f"json_extract({prefix}data, '$.{field}')" + + +def _check_table_alias(table: Optional[str]) -> str: + if not table: + return "" + if not _SAFE_SEGMENT_RE.match(table): + raise ValueError(f"unsafe table alias: {table!r}") + return table def _is_scalar(value: Any) -> bool: @@ -114,7 +124,7 @@ def _scalar_param(value: Any) -> Any: def _translate_field_clause( - field: str, condition: Any + field: str, condition: Any, table: str = "" ) -> Optional[Tuple[str, List[Any]]]: """Translate ``{field: condition}`` for one field. @@ -123,7 +133,7 @@ def _translate_field_clause( if not _safe_field_path(field): return None - column = _json_extract(field) + column = _json_extract(field, table) # Plain equality with a scalar value. if not isinstance(condition, dict): @@ -204,7 +214,9 @@ def _translate_field_clause( return " AND ".join(fragments), params -def _translate_logical(op: str, conditions: Any) -> Optional[Tuple[str, List[Any]]]: +def _translate_logical( + op: str, conditions: Any, table: str = "" +) -> Optional[Tuple[str, List[Any]]]: """Translate ``$and`` / ``$or`` recursively.""" if not isinstance(conditions, list) or not conditions: return None @@ -213,7 +225,7 @@ def _translate_logical(op: str, conditions: Any) -> Optional[Tuple[str, List[Any for sub in conditions: if not isinstance(sub, dict): return None - translated = translate_query(sub) + translated = translate_query(sub, table=table) if translated is None: return None sub_sql, sub_params = translated @@ -223,18 +235,22 @@ def _translate_logical(op: str, conditions: Any) -> Optional[Tuple[str, List[Any return joiner.join(parts), params -def translate_query(query: Dict[str, Any]) -> Optional[Tuple[str, List[Any]]]: +def translate_query( + query: Dict[str, Any], *, table: Optional[str] = None +) -> Optional[Tuple[str, List[Any]]]: """Translate a Mongo-style query dict to ``(sql_where, params)``. Returns ``None`` when any portion of the query can't be expressed in SQL we trust; the caller should fall back to in-Python filtering. The returned SQL fragment is meant to be ANDed into a larger WHERE - clause; e.g. ``WHERE collection = ? AND ()``. + clause; e.g. ``WHERE collection = ? AND ()``. ``table`` + qualifies the ``data`` column with a table alias for joins. """ if not query: return "", [] + alias = _check_table_alias(table) fragments: List[str] = [] params: List[Any] = [] @@ -242,7 +258,7 @@ def translate_query(query: Dict[str, Any]) -> Optional[Tuple[str, List[Any]]]: if key in _IGNORED_TOP_LEVEL: continue if key in ("$and", "$or"): - translated = _translate_logical(key, value) + translated = _translate_logical(key, value, alias) if translated is None: return None sub_sql, sub_params = translated @@ -252,7 +268,7 @@ def translate_query(query: Dict[str, Any]) -> Optional[Tuple[str, List[Any]]]: if key.startswith("$"): # Unknown top-level operator -> fallback. return None - translated = _translate_field_clause(key, value) + translated = _translate_field_clause(key, value, alias) if translated is None: return None sub_sql, sub_params = translated @@ -264,7 +280,9 @@ def translate_query(query: Dict[str, Any]) -> Optional[Tuple[str, List[Any]]]: return " AND ".join(fragments), params -def translate_sort(sort: Optional[List[Tuple[str, int]]]) -> Optional[str]: +def translate_sort( + sort: Optional[List[Tuple[str, int]]], *, table: Optional[str] = None +) -> Optional[str]: """Translate a sort spec to a SQL ORDER BY fragment. Returns ``None`` when the sort can't be expressed (unsafe field name, @@ -276,13 +294,14 @@ def translate_sort(sort: Optional[List[Tuple[str, int]]]) -> Optional[str]: """ if not sort: return None + alias = _check_table_alias(table) parts: List[str] = [] for field, direction in sort: if direction not in (1, -1): return None if not _safe_field_path(field): return None - column = _json_extract(field) + column = _json_extract(field, alias) if direction == 1: # ascending: NULLs last parts.append(f"({column} IS NULL), {column} ASC") diff --git a/jvspatial/db/mongodb.py b/jvspatial/db/mongodb.py index 2471b6e..355b6a8 100644 --- a/jvspatial/db/mongodb.py +++ b/jvspatial/db/mongodb.py @@ -18,7 +18,18 @@ import contextlib import logging -from typing import Any, Awaitable, Callable, Dict, List, Optional, Set, Tuple, Union +from typing import ( + Any, + Awaitable, + Callable, + Dict, + List, + Optional, + Sequence, + Set, + Tuple, + Union, +) from motor.motor_asyncio import AsyncIOMotorClient, AsyncIOMotorDatabase from pymongo.errors import ( @@ -344,6 +355,156 @@ async def _strip_op() -> int: return int(await self._run_with_reconnect("strip_node_edges", _strip_op)) + @staticmethod + def _connected_pipeline( + node_collection: str, + start_id: str, + *, + direction: str, + edge_entities: Optional[Sequence[str]], + node_entities: Optional[Sequence[str]], + edge_query: Optional[Dict[str, Any]], + node_query: Optional[Dict[str, Any]], + ) -> List[Dict[str, Any]]: + """Aggregation over the edge collection yielding neighbour node documents. + + One edge type in one direction streams edge → ``$lookup`` → node (the + ``(source, target, entity)`` unique index makes it duplicate-free, so a + trailing ``$limit`` stops early). Otherwise neighbour ids are grouped + first so each neighbour appears once. + """ + if direction not in ("out", "in", "both"): + raise ValueError( + f"direction must be 'out', 'in' or 'both', got {direction!r}" + ) + endpoints: Dict[str, Dict[str, Any]] = { + "out": {"source": start_id}, + "in": {"target": start_id}, + "both": {"$or": [{"source": start_id}, {"target": start_id}]}, + } + edge_match: List[Dict[str, Any]] = [endpoints[direction]] + if edge_entities is not None: + edge_match.append({"entity": {"$in": list(edge_entities)}}) + if edge_query: + edge_match.append(edge_query) + far: Any = { + "out": "$target", + "in": "$source", + "both": {"$cond": [{"$eq": ["$source", start_id]}, "$target", "$source"]}, + }[direction] + pipeline: List[Dict[str, Any]] = [ + {"$match": {"$and": edge_match}}, + {"$project": {"_far": far}}, + ] + if direction == "both" or edge_entities is None or len(edge_entities) != 1: + pipeline.append({"$group": {"_id": "$_far"}}) + local_field = "_id" + else: + local_field = "_far" + pipeline += [ + { + "$lookup": { + "from": node_collection, + "localField": local_field, + "foreignField": "_id", + "as": "_n", + } + }, + {"$unwind": "$_n"}, + {"$replaceRoot": {"newRoot": "$_n"}}, + ] + node_match: List[Dict[str, Any]] = [] + if node_entities is not None: + node_match.append({"entity": {"$in": list(node_entities)}}) + if node_query: + node_match.append(node_query) + if node_match: + pipeline.append({"$match": {"$and": node_match}}) + return pipeline + + async def find_connected_nodes( + self, + node_collection: str, + edge_collection: str, + start_id: str, + *, + direction: str = "out", + edge_entity: Optional[str] = None, + edge_entities: Optional[Sequence[str]] = None, + node_entities: Optional[Sequence[str]] = None, + edge_query: Optional[Dict[str, Any]] = None, + node_query: Optional[Dict[str, Any]] = None, + sort: Optional[List[Tuple[str, int]]] = None, + limit: Optional[int] = None, + ) -> List[Dict[str, Any]]: + """Single-hop neighbours of ``start_id`` in one aggregation round trip. + + Same contract as ``PostgresDB.find_connected_nodes``: entity lists and + record-path queries filter the edge / node side, ``sort`` / ``limit`` + apply to neighbours. + """ + if edge_entities is None and edge_entity is not None: + edge_entities = [edge_entity] + pipeline = self._connected_pipeline( + node_collection, + start_id, + direction=direction, + edge_entities=edge_entities, + node_entities=node_entities, + edge_query=edge_query, + node_query=node_query, + ) + if sort: + order: Dict[str, int] = {} + for field, direction_ in sort: + order[field] = direction_ + order.setdefault("_id", 1) + pipeline.append({"$sort": order}) + if limit is not None: + pipeline.append({"$limit": int(limit)}) + + async def _op() -> List[Dict[str, Any]]: + await self._ensure_connected() + if self._db is None: + raise DatabaseError("MongoDB database connection not established") + cursor = self._db[edge_collection].aggregate(pipeline) + return await cursor.to_list(length=None) + + return await self._run_with_reconnect("find_connected_nodes", _op) + + async def count_connected_nodes( + self, + node_collection: str, + edge_collection: str, + start_id: str, + *, + direction: str = "out", + edge_entities: Optional[Sequence[str]] = None, + node_entities: Optional[Sequence[str]] = None, + edge_query: Optional[Dict[str, Any]] = None, + node_query: Optional[Dict[str, Any]] = None, + ) -> int: + """Count of the neighbours :meth:`find_connected_nodes` would return.""" + pipeline = self._connected_pipeline( + node_collection, + start_id, + direction=direction, + edge_entities=edge_entities, + node_entities=node_entities, + edge_query=edge_query, + node_query=node_query, + ) + pipeline.append({"$count": "n"}) + + async def _op() -> int: + await self._ensure_connected() + if self._db is None: + raise DatabaseError("MongoDB database connection not established") + rows = await self._db[edge_collection].aggregate(pipeline).to_list(1) + return int(rows[0]["n"]) if rows else 0 + + return int(await self._run_with_reconnect("count_connected_nodes", _op)) + async def find( self, collection: str, diff --git a/jvspatial/db/postgres.py b/jvspatial/db/postgres.py index d8f8daf..0eeb991 100644 --- a/jvspatial/db/postgres.py +++ b/jvspatial/db/postgres.py @@ -78,6 +78,7 @@ Dict, List, Optional, + Sequence, Set, Tuple, Union, @@ -994,55 +995,287 @@ async def find_many( out[row["id"]] = self._record_from_row(row) return out - async def find_connected_nodes( + def _connected_nodes_sql( self, node_collection: str, edge_collection: str, start_id: str, *, direction: str = "out", - edge_entity: Optional[str] = None, + edge_entities: Optional[Sequence[str]] = None, + node_entities: Optional[Sequence[str]] = None, + edge_query: Optional[Dict[str, Any]] = None, + node_query: Optional[Dict[str, Any]] = None, + sort: Optional[List[Tuple[str, int]]] = None, limit: Optional[int] = None, - ) -> List[Dict[str, Any]]: - """Single-hop neighbor fetch via edge+node join (one round trip).""" - if direction not in ("out", "in"): + count: bool = False, + with_edges: bool = False, + ) -> Tuple[str, List[Any]]: + """Build the single-hop neighbour query shared by find / count. + + Two shapes, both one round trip: + + * **join** — ``edge ⋈ node`` driven by the ``source`` / ``target`` + index. Used for one edge type in one direction (duplicate-free under + the ``(source, target, entity)`` unique index, so ``LIMIT`` stops + after N rows whatever the node's degree) and whenever edge rows are + requested. + * **semi-join** — ``node WHERE id IN ()``. Used for + several / no edge types or ``direction="both"``, where one neighbour + can be reached through several edges and must appear once. + + Raises: + ValueError: ``direction`` not in ``{"out", "in", "both"}``. + NotImplementedError: ``edge_query`` / ``node_query`` / ``sort`` + cannot be translated — callers fall back to the Python path + rather than drop the filter. + """ + if direction not in ("out", "in", "both"): raise ValueError( - "direction must be 'out' or 'in' for find_connected_nodes, " - f"got {direction!r}" + f"direction must be 'out', 'in' or 'both', got {direction!r}" ) - await self._bootstrap_collection(edge_collection) - await self._bootstrap_collection(node_collection) edge_col = _safe_collection(edge_collection) node_col = _safe_collection(node_collection) schema = _safe_collection(self.schema_name) + params: List[Any] = [str(start_id)] - if direction == "out": - join_on = "(e.data #>> '{target}') = n.id" - where_endpoint = "(e.data #>> '{source}') = $1" + def bind(value: Any) -> str: + params.append(value) + return f"${len(params)}" + + def fragment(query: Optional[Dict[str, Any]], table: str) -> str: + if not query: + return "" + translated = translate_query(query, table=table) + if translated is None: + raise NotImplementedError( + f"find_connected_nodes: {table}-side filter does not " + f"translate to SQL: {query!r}" + ) + sql, sub_params = translated + if not sql: + return "" + shifted = _shift_placeholders(sql, shift=len(params)) + params.extend(sub_params) + return f" AND ({shifted})" + + edge_pred = "" + if edge_entities is not None: + edge_pred += f" AND e.entity = ANY({bind(list(edge_entities))}::text[])" + edge_pred += fragment(edge_query, "e") + node_pred = "" + if node_entities is not None: + node_pred += f" AND n.entity = ANY({bind(list(node_entities))}::text[])" + node_pred += fragment(node_query, "n") + + # (near endpoint = start_id, far endpoint = neighbour) per branch. + ends = { + "out": [("source", "target")], + "in": [("target", "source")], + "both": [("source", "target"), ("target", "source")], + }[direction] + join_form = with_edges or ( + direction != "both" + and edge_entities is not None + and len(edge_entities) == 1 + ) + if join_form: + edge_col_sql = ", e.data AS edge" if with_edges else "" + inner = " UNION ALL ".join( + f"SELECT n.id AS id, n.data AS data{edge_col_sql} " + f"FROM {schema}.{edge_col} e " + f"JOIN {schema}.{node_col} n ON n.id = (e.data #>> '{{{far}}}') " + f"WHERE (e.data #>> '{{{near}}}') = $1{edge_pred}{node_pred}" + for near, far in ends + ) else: - join_on = "(e.data #>> '{source}') = n.id" - where_endpoint = "(e.data #>> '{target}') = $1" + far_ids = " UNION ALL ".join( + f"SELECT (e.data #>> '{{{far}}}') FROM {schema}.{edge_col} e " + f"WHERE (e.data #>> '{{{near}}}') = $1{edge_pred}" + for near, far in ends + ) + inner = ( + f"SELECT n.id AS id, n.data AS data FROM {schema}.{node_col} n " + f"WHERE n.id IN ({far_ids}){node_pred}" + ) - clauses = [where_endpoint] - params: List[Any] = [start_id] - if edge_entity is not None: - clauses.append(f"e.entity = ${len(params) + 1}") - params.append(edge_entity) + if count: + return f"SELECT COUNT(*) FROM ({inner}) sub", params - limit_sql = "" - if limit is not None: - limit_sql = f" LIMIT ${len(params) + 1}" - params.append(int(limit)) + order_sql = "" + if sort: + sort_sql = translate_sort(sort, table="sub") + if sort_sql is None: + raise NotImplementedError( + f"find_connected_nodes: sort does not translate: {sort!r}" + ) + order_sql = f" ORDER BY {sort_sql}, sub.id ASC" + limit_sql = f" LIMIT {bind(int(limit))}" if limit is not None else "" + columns = "sub.data, sub.edge" if with_edges else "sub.data" + return f"SELECT {columns} FROM ({inner}) sub{order_sql}{limit_sql}", params - sql = ( - f"SELECT n.data FROM {schema}.{edge_col} e " - f"JOIN {schema}.{node_col} n ON {join_on} " - f"WHERE {' AND '.join(clauses)}{limit_sql}" + async def find_connected_nodes( + self, + node_collection: str, + edge_collection: str, + start_id: str, + *, + direction: str = "out", + edge_entity: Optional[str] = None, + edge_entities: Optional[Sequence[str]] = None, + node_entities: Optional[Sequence[str]] = None, + edge_query: Optional[Dict[str, Any]] = None, + node_query: Optional[Dict[str, Any]] = None, + sort: Optional[List[Tuple[str, int]]] = None, + limit: Optional[int] = None, + with_edges: bool = False, + ) -> List[Dict[str, Any]]: + """Single-hop neighbours of ``start_id`` in one round trip. + + Every filter is pushed into SQL: ``edge_entities`` / ``node_entities`` + become ``entity = ANY(...)``, ``edge_query`` / ``node_query`` (record + paths, same dialect as :meth:`find`) are translated against the edge / + node side of the join, and ``sort`` / ``limit`` apply to neighbours. + ``edge_entity`` is the pre-0.0.18 single-type spelling. + + Returns node records — or, with ``with_edges=True``, one + ``{"node": ..., "edge": ...}`` dict per connecting edge. + + Raises: + NotImplementedError: a filter or the sort cannot be translated; + fall back to the Python traversal path. + """ + if edge_entities is None and edge_entity is not None: + edge_entities = [edge_entity] + await self._bootstrap_collection(edge_collection) + await self._bootstrap_collection(node_collection) + sql, params = self._connected_nodes_sql( + node_collection, + edge_collection, + start_id, + direction=direction, + edge_entities=edge_entities, + node_entities=node_entities, + edge_query=edge_query, + node_query=node_query, + sort=sort, + limit=limit, + with_edges=with_edges, ) async with self._acquire_conn() as conn: rows = await conn.fetch(sql, *params) + if with_edges: + return [ + { + "node": self._record_from_row(r), + "edge": self._record_from_row({"data": r["edge"]}), + } + for r in rows + ] return [self._record_from_row(r) for r in rows] + async def count_connected_nodes( + self, + node_collection: str, + edge_collection: str, + start_id: str, + *, + direction: str = "out", + edge_entities: Optional[Sequence[str]] = None, + node_entities: Optional[Sequence[str]] = None, + edge_query: Optional[Dict[str, Any]] = None, + node_query: Optional[Dict[str, Any]] = None, + ) -> int: + """``COUNT`` of the neighbours :meth:`find_connected_nodes` would return.""" + await self._bootstrap_collection(edge_collection) + await self._bootstrap_collection(node_collection) + sql, params = self._connected_nodes_sql( + node_collection, + edge_collection, + start_id, + direction=direction, + edge_entities=edge_entities, + node_entities=node_entities, + edge_query=edge_query, + node_query=node_query, + count=True, + ) + async with self._acquire_conn() as conn: + return int(await conn.fetchval(sql, *params)) + + async def find_connected_nodes_bulk( + self, + node_collection: str, + edge_collection: str, + start_ids: Sequence[str], + *, + direction: str = "out", + edge_entities: Optional[Sequence[str]] = None, + node_entities: Optional[Sequence[str]] = None, + edge_query: Optional[Dict[str, Any]] = None, + node_query: Optional[Dict[str, Any]] = None, + limit_per_source: Optional[int] = None, + ) -> Dict[str, List[Dict[str, Any]]]: + """Neighbours of many sources in one round trip, capped per source. + + ``ROW_NUMBER() OVER (PARTITION BY source ORDER BY neighbour id)`` keeps + at most ``limit_per_source`` distinct neighbours per start id. + ``direction`` is ``"out"`` or ``"in"``. + + Raises: + NotImplementedError: a filter cannot be translated. + """ + if direction not in ("out", "in"): + raise ValueError(f"direction must be 'out' or 'in', got {direction!r}") + out: Dict[str, List[Dict[str, Any]]] = {str(s): [] for s in start_ids} + if not out: + return out + await self._bootstrap_collection(edge_collection) + await self._bootstrap_collection(node_collection) + edge_col = _safe_collection(edge_collection) + node_col = _safe_collection(node_collection) + schema = _safe_collection(self.schema_name) + near, far = ("source", "target") if direction == "out" else ("target", "source") + params: List[Any] = [list(out)] + preds = "" + for entities, table in ((edge_entities, "e"), (node_entities, "n")): + if entities is not None: + params.append(list(entities)) + preds += f" AND {table}.entity = ANY(${len(params)}::text[])" + for query, table in ((edge_query, "e"), (node_query, "n")): + if not query: + continue + translated = translate_query(query, table=table) + if translated is None: + raise NotImplementedError( + f"find_connected_nodes_bulk: filter does not translate: {query!r}" + ) + sql, sub_params = translated + if sql: + preds += f" AND ({_shift_placeholders(sql, shift=len(params))})" + params.extend(sub_params) + cap = "" + if limit_per_source is not None: + params.append(int(limit_per_source)) + cap = f" WHERE r.rn <= ${len(params)}" + sql = ( + f"WITH pairs AS (" + f"SELECT DISTINCT (e.data #>> '{{{near}}}') AS src, n.id AS nid " + f"FROM {schema}.{edge_col} e " + f"JOIN {schema}.{node_col} n ON n.id = (e.data #>> '{{{far}}}') " + f"WHERE (e.data #>> '{{{near}}}') = ANY($1::text[]){preds}" + f"), ranked AS (" + f"SELECT src, nid, ROW_NUMBER() OVER (PARTITION BY src ORDER BY nid) AS rn " + f"FROM pairs) " + f"SELECT r.src AS src, n.data AS data FROM ranked r " + f"JOIN {schema}.{node_col} n ON n.id = r.nid{cap} ORDER BY r.src, r.rn" + ) + async with self._acquire_conn() as conn: + rows = await conn.fetch(sql, *params) + for row in rows: + out[row["src"]].append(self._record_from_row(row)) + return out + async def bulk_save_detailed( self, collection: str, records: List[Dict[str, Any]] ) -> BulkSaveResult: diff --git a/jvspatial/db/sqlite.py b/jvspatial/db/sqlite.py index cf1f2d0..f1e4cd6 100644 --- a/jvspatial/db/sqlite.py +++ b/jvspatial/db/sqlite.py @@ -14,7 +14,17 @@ import logging import uuid from pathlib import Path -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple, Union +from typing import ( + TYPE_CHECKING, + Any, + Dict, + List, + Optional, + Sequence, + Set, + Tuple, + Union, +) from ._sqlite_translate import ( translate_partial_filter_expression, @@ -734,6 +744,201 @@ async def count( rows = await self.find(collection, q) return len(rows) + # ---- single-hop neighbour pushdown ------------------------------------- + + def _connected_nodes_sql( + self, + node_collection: str, + edge_collection: str, + start_id: str, + *, + direction: str = "out", + edge_entities: Optional[Sequence[str]] = None, + node_entities: Optional[Sequence[str]] = None, + edge_query: Optional[Dict[str, Any]] = None, + node_query: Optional[Dict[str, Any]] = None, + sort: Optional[List[Tuple[str, int]]] = None, + limit: Optional[int] = None, + count: bool = False, + ) -> Tuple[str, List[Any]]: + """Single-hop neighbour SQL over ``records`` (see ``PostgresDB``). + + Join for one edge type in one direction, ``id IN (...)`` semi-join + otherwise (one row per neighbour). Positional ``?`` parameters follow + the SQL text, so predicates repeated across ``UNION ALL`` branches + repeat their parameters. + + Raises: + ValueError: ``direction`` not in ``{"out", "in", "both"}``. + NotImplementedError: a query or the sort does not translate. + """ + if direction not in ("out", "in", "both"): + raise ValueError( + f"direction must be 'out', 'in' or 'both', got {direction!r}" + ) + + def entity_clause( + table: str, entities: Optional[Sequence[str]] + ) -> Tuple[str, List[Any]]: + if entities is None: + return "", [] + if not entities: + return " AND 0", [] + marks = ",".join("?" * len(entities)) + return ( + f" AND json_extract({table}.data, '$.entity') IN ({marks})", + list(entities), + ) + + def query_clause( + table: str, query: Optional[Dict[str, Any]] + ) -> Tuple[str, List[Any]]: + if not query: + return "", [] + translated = translate_query(query, table=table) + if translated is None: + raise NotImplementedError( + f"find_connected_nodes: {table}-side filter does not " + f"translate to SQL: {query!r}" + ) + sql, params = translated + return (f" AND ({sql})", list(params)) if sql else ("", []) + + edge_type_sql, edge_type_params = entity_clause("e", edge_entities) + edge_q_sql, edge_q_params = query_clause("e", edge_query) + edge_pred, edge_params = ( + edge_type_sql + edge_q_sql, + edge_type_params + edge_q_params, + ) + node_type_sql, node_type_params = entity_clause("n", node_entities) + node_q_sql, node_q_params = query_clause("n", node_query) + node_pred, node_params = ( + node_type_sql + node_q_sql, + node_type_params + node_q_params, + ) + + ends = { + "out": [("source", "target")], + "in": [("target", "source")], + "both": [("source", "target"), ("target", "source")], + }[direction] + params: List[Any] = [] + if ( + direction != "both" + and edge_entities is not None + and len(edge_entities) == 1 + ): + near, far = ends[0] + inner = ( + "SELECT n.id AS id, n.data AS data FROM records e " + f"JOIN records n ON n.collection = ? " + f"AND n.id = json_extract(e.data, '$.{far}') " + f"WHERE e.collection = ? AND json_extract(e.data, '$.{near}') = ?" + f"{edge_pred}{node_pred}" + ) + params += [node_collection, edge_collection, start_id] + params += edge_params + node_params + else: + far_ids = " UNION ALL ".join( + f"SELECT json_extract(e.data, '$.{far}') FROM records e " + f"WHERE e.collection = ? AND json_extract(e.data, '$.{near}') = ?" + f"{edge_pred}" + for near, far in ends + ) + inner = ( + "SELECT n.id AS id, n.data AS data FROM records n " + f"WHERE n.collection = ? AND n.id IN ({far_ids}){node_pred}" + ) + params.append(node_collection) + for _ in ends: + params += [edge_collection, start_id, *edge_params] + params += node_params + + if count: + return f"SELECT COUNT(*) FROM ({inner}) sub", params + order_sql = "" + if sort: + sort_sql = translate_sort(sort, table="sub") + if sort_sql is None: + raise NotImplementedError( + f"find_connected_nodes: sort does not translate: {sort!r}" + ) + order_sql = f" ORDER BY {sort_sql}, sub.id ASC" + limit_sql = "" + if limit is not None: + limit_sql = " LIMIT ?" + params.append(int(limit)) + return ( + f"SELECT sub.data AS data FROM ({inner}) sub{order_sql}{limit_sql}", + params, + ) + + async def find_connected_nodes( + self, + node_collection: str, + edge_collection: str, + start_id: str, + *, + direction: str = "out", + edge_entity: Optional[str] = None, + edge_entities: Optional[Sequence[str]] = None, + node_entities: Optional[Sequence[str]] = None, + edge_query: Optional[Dict[str, Any]] = None, + node_query: Optional[Dict[str, Any]] = None, + sort: Optional[List[Tuple[str, int]]] = None, + limit: Optional[int] = None, + ) -> List[Dict[str, Any]]: + """Single-hop neighbours in one query (same contract as ``PostgresDB``).""" + if edge_entities is None and edge_entity is not None: + edge_entities = [edge_entity] + sql, params = self._connected_nodes_sql( + node_collection, + edge_collection, + start_id, + direction=direction, + edge_entities=edge_entities, + node_entities=node_entities, + edge_query=edge_query, + node_query=node_query, + sort=sort, + limit=limit, + ) + connection = await self._get_connection() + cursor = await connection.execute(sql, tuple(params)) + rows = await cursor.fetchall() + await cursor.close() + return [json.loads(row["data"]) for row in rows] + + async def count_connected_nodes( + self, + node_collection: str, + edge_collection: str, + start_id: str, + *, + direction: str = "out", + edge_entities: Optional[Sequence[str]] = None, + node_entities: Optional[Sequence[str]] = None, + edge_query: Optional[Dict[str, Any]] = None, + node_query: Optional[Dict[str, Any]] = None, + ) -> int: + """Count of the neighbours :meth:`find_connected_nodes` would return.""" + sql, params = self._connected_nodes_sql( + node_collection, + edge_collection, + start_id, + direction=direction, + edge_entities=edge_entities, + node_entities=node_entities, + edge_query=edge_query, + node_query=node_query, + count=True, + ) + connection = await self._get_connection() + cursor = await connection.execute(sql, tuple(params)) + row = await cursor.fetchone() + await cursor.close() + return int(row[0]) if row else 0 + # Context manager helpers for convenience async def __aenter__(self) -> "SQLiteDB": """Async context manager entry.""" diff --git a/tests/core/test_connected_nodes_fast_path.py b/tests/core/test_connected_nodes_fast_path.py index 234989d..af06915 100644 --- a/tests/core/test_connected_nodes_fast_path.py +++ b/tests/core/test_connected_nodes_fast_path.py @@ -33,7 +33,9 @@ class _FastJsonDB(JsonDB): """JsonDB plus a ``find_connected_nodes`` so the fast path is exercised. Mirrors the Postgres backend contract: strict source/target endpoints per - ``direction`` and DB-side ``limit`` applied to the raw (unfiltered) rows. + ``direction``, entity filters applied in the "query", and the DB-side + ``limit`` applied after them. Property queries / sort are declined with + ``NotImplementedError`` (the caller's Python path takes over). """ async def find_connected_nodes( @@ -44,8 +46,17 @@ async def find_connected_nodes( *, direction: str = "out", edge_entity: Optional[str] = None, + edge_entities: Optional[List[str]] = None, + node_entities: Optional[List[str]] = None, + edge_query: Optional[Dict[str, Any]] = None, + node_query: Optional[Dict[str, Any]] = None, + sort: Optional[List[Any]] = None, limit: Optional[int] = None, ) -> List[Dict[str, Any]]: + if edge_query or node_query or sort or direction not in ("out", "in"): + raise NotImplementedError("shim supports type-only out/in queries") + if edge_entities is None and edge_entity is not None: + edge_entities = [edge_entity] edges = await self.find(edge_collection, {}) rows: List[Dict[str, Any]] = [] for e in edges: @@ -56,11 +67,14 @@ async def find_connected_nodes( other = src else: continue - if edge_entity is not None and e.get("entity") != edge_entity: + if edge_entities is not None and e.get("entity") not in edge_entities: continue n = await self.get(node_collection, other) - if n is not None: - rows.append(n) + if n is None: + continue + if node_entities is not None and n.get("entity") not in node_entities: + continue + rows.append(n) # Deterministic order (Postgres has no implicit ORDER BY, so the bug # surfaces whenever a non-matching neighbor happens to sort first). rows.sort(key=lambda r: r["id"]) diff --git a/tests/core/test_node_query_pushdown.py b/tests/core/test_node_query_pushdown.py new file mode 100644 index 0000000..40fe06e --- /dev/null +++ b/tests/core/test_node_query_pushdown.py @@ -0,0 +1,474 @@ +"""Neighbour-query pushdown — ``nodes()`` / ``count_nodes()`` / ``nodes_page()``. + +Every filter shape ``nodes()`` accepts (single / list / string / class / +``{name: criteria}``, node and edge side, plus property kwargs) runs against +every reachable backend and is compared with a pure-Python reference over the +raw records. On Postgres, type-only filters must cost exactly one round trip +and the join must not sequentially scan the edge table. + +Postgres runs when ``JVSPATIAL_POSTGRES_TEST_DSN`` is reachable, MongoDB when +``JVSPATIAL_MONGODB_TEST_URI`` is set and reachable; otherwise those skip. +""" + +from __future__ import annotations + +import asyncio +import itertools +import json +import os +import uuid +from typing import Any, Dict, List, Optional, Set + +import pytest + +from jvspatial.core import context as context_module +from jvspatial.core.context import GraphContext, set_default_context +from jvspatial.core.entities import Edge, Node +from jvspatial.core.utils import generate_id +from jvspatial.db._observable import ObservableDatabase +from jvspatial.db.jsondb import JsonDB +from jvspatial.db.sqlite import SQLiteDB +from jvspatial.observability import db_op_counter + +pytestmark = pytest.mark.asyncio(loop_scope="module") + +_PG_DSN = os.getenv( + "JVSPATIAL_POSTGRES_TEST_DSN", + "postgresql://jvspatial:jvspatial@localhost:5432/jvspatial", +) +_MONGO_URI = os.getenv("JVSPATIAL_MONGODB_TEST_URI") + + +class QHub(Node): + name: str = "" + + +class QAlpha(Node): + name: str = "" + score: int = 0 + + +class QAlphaSub(QAlpha): + pass + + +class QBeta(Node): + name: str = "" + score: int = 0 + + +class QKnows(Edge): + weight: int = 0 + + +class QLikes(Edge): + weight: int = 0 + + +_NODE_CLASSES = {"QHub": QHub, "QAlpha": QAlpha, "QAlphaSub": QAlphaSub, "QBeta": QBeta} + + +# ---- backends --------------------------------------------------------------- + + +async def _pg_admin() -> Any: + try: + import asyncpg + except ImportError: + pytest.skip("asyncpg not installed") + try: + return await asyncio.wait_for(asyncpg.connect(dsn=_PG_DSN), timeout=2.0) + except Exception: + pytest.skip(f"Postgres unreachable at {_PG_DSN}") + + +async def _mongo_client() -> Any: + if not _MONGO_URI: + pytest.skip("JVSPATIAL_MONGODB_TEST_URI not set") + from motor.motor_asyncio import AsyncIOMotorClient + + client = AsyncIOMotorClient(_MONGO_URI, serverSelectionTimeoutMS=2000) + try: + await client.admin.command("ping") + except Exception: + client.close() + pytest.skip(f"MongoDB unreachable at {_MONGO_URI}") + return client + + +@pytest.fixture(params=["json", "sqlite", "postgres", "mongodb"]) +async def graph(request, tmp_path): + """Yield ``(ctx, db, backend, extras)`` on a fresh store.""" + backend = request.param + cleanup: List[Any] = [] + extras: Dict[str, Any] = {} + if backend == "json": + db: Any = JsonDB(base_path=str(tmp_path / "json")) + elif backend == "sqlite": + db = SQLiteDB(db_path=str(tmp_path / "graph.db")) + elif backend == "postgres": + from jvspatial.db.postgres import PostgresDB + + admin = await _pg_admin() + schema = f"t_push_{uuid.uuid4().hex[:10]}" + await admin.execute(f"CREATE SCHEMA {schema}") + db = PostgresDB(dsn=_PG_DSN, schema_name=schema, min_size=1, max_size=8) + extras.update(admin=admin, schema=schema) + + async def _drop_pg() -> None: + await admin.execute(f"DROP SCHEMA {schema} CASCADE") + await admin.close() + + cleanup.append(_drop_pg) + else: + from jvspatial.db.mongodb import MongoDB + + client = await _mongo_client() + name = f"jvs_push_{uuid.uuid4().hex[:10]}" + db = MongoDB(uri=_MONGO_URI, db_name=name) + + async def _drop_mongo() -> None: + await client.drop_database(name) + client.close() + + cleanup.append(_drop_mongo) + + ctx = GraphContext(database=db) + set_default_context(ctx) + context_module._ensured_indexes.clear() + try: + yield ctx, db, backend, extras + finally: + set_default_context(None) + context_module._ensured_indexes.clear() + close = getattr(db, "close", None) + if callable(close): + result = close() + if asyncio.iscoroutine(result): + await result + for fn in cleanup: + await fn() + + +# ---- fixture graph + reference ------------------------------------------------ + + +async def _seed() -> Dict[str, Node]: + """Hub with mixed edge / neighbour types, a subclass, and both-way links.""" + n: Dict[str, Node] = {"hub": await QHub.create(name="hub")} + for key, cls, score in ( + ("a1", QAlpha, 1), + ("a2", QAlpha, 5), + ("s1", QAlphaSub, 3), + ("b1", QBeta, 2), + ("b2", QBeta, 7), + ("x", QBeta, 4), + ("z", QAlpha, 9), + ): + n[key] = await cls.create(name=key, score=score) + hub = n["hub"] + for target, edge, weight in ( + ("a1", QKnows, 1), + ("a2", QKnows, 3), + ("s1", QLikes, 2), + ("b1", QKnows, 5), + ("b2", QLikes, 1), + ("z", QKnows, 2), + ): + await hub.connect(n[target], edge=edge, weight=weight) + for source, edge, weight in (("x", QKnows, 4), ("z", QLikes, 6), ("b1", QLikes, 3)): + await n[source].connect(hub, edge=edge, weight=weight) + return n + + +def _crit_ok(values: Dict[str, Any], criteria: Dict[str, Any]) -> bool: + for key, cond in criteria.items(): + value = values.get(key.split(".", 1)[1] if key.startswith("context.") else key) + if isinstance(cond, dict): + for op, operand in cond.items(): + if value is None: + return False + ok = { + "$gt": lambda: value > operand, + "$gte": lambda: value >= operand, + "$lt": lambda: value < operand, + "$in": lambda: value in operand, + }[op]() + if not ok: + return False + elif value != cond: + return False + return True + + +def _entity_ok( + entity: str, flt: Any, criteria_values: Dict[str, Any], node: bool +) -> bool: + if flt is None: + return True + items = flt if isinstance(flt, list) else [flt] + if not items and not node: + return True + for item in items: + if isinstance(item, str) and entity == item: + return True + if isinstance(item, type): + if node and issubclass(_NODE_CLASSES[entity], item): + return True + if not node and entity == item.__name__: + return True + if isinstance(item, dict): + for name, criteria in item.items(): + if entity == name and _crit_ok(criteria_values, criteria): + return True + return False + + +async def _reference( + db: Any, hub_id: str, direction: str, node_f: Any, edge_f: Any, props: Dict +) -> Set[str]: + edges = await db.find("edge", {}) + nodes = {r["id"]: r for r in await db.find("node", {})} + out: Set[str] = set() + for e in edges: + if not _entity_ok(e["entity"], edge_f, e.get("context", {}), node=False): + continue + far = [] + if direction in ("out", "both") and e["source"] == hub_id: + far.append(e["target"]) + if direction in ("in", "both") and e["target"] == hub_id: + far.append(e["source"]) + for nid in far: + rec = nodes.get(nid) + if rec is None: + continue + ctx = rec.get("context", {}) + if _entity_ok(rec["entity"], node_f, ctx, node=True) and _crit_ok( + ctx, props + ): + out.add(nid) + return out + + +_EDGE_FILTERS: List[Any] = [ + None, + QKnows, + [QKnows], + "QKnows", + ["QKnows", "QLikes"], + [{"QKnows": {"weight": {"$gte": 2}}}], + [QLikes, {"QKnows": {"weight": {"$gte": 3}}}], +] +_NODE_FILTERS: List[Any] = [ + None, + QAlpha, + [QAlpha], + "QAlpha", + ["QAlpha", "QBeta"], + [{"QAlpha": {"score": {"$gt": 1}}}], + [QBeta, {"QAlpha": {"score": {"$gte": 5}}}], +] +_PROPS: List[Dict[str, Any]] = [{}, {"score": {"$gte": 3}}] + + +def _type_only(flt: Any) -> bool: + items = flt if isinstance(flt, list) else [flt] + return all(not isinstance(i, dict) for i in items) + + +# ---- tests ------------------------------------------------------------------ + + +async def test_nodes_matches_reference_for_every_filter_shape(graph): + ctx, db, backend, _ = graph + n = await _seed() + hub = n["hub"] + failures: List[str] = [] + for edge_f, node_f, direction, props in itertools.product( + _EDGE_FILTERS, _NODE_FILTERS, ("out", "in", "both"), _PROPS + ): + expected = await _reference(db, hub.id, direction, node_f, edge_f, props) + for limit in (None, 1, 5): + await ctx.clear_cache() + got = await hub.nodes( + direction=direction, node=node_f, edge=edge_f, limit=limit, **props + ) + ids = [x.id for x in got] + ok = len(ids) == len(set(ids)) and ( + set(ids) == expected + if limit is None + else set(ids) <= expected and len(ids) == min(limit, len(expected)) + ) + if isinstance(node_f, type) or ( + isinstance(node_f, list) and node_f and isinstance(node_f[0], type) + ): + ok = ok and all( + isinstance(x, (QAlpha, QBeta)) for x in got + ) # hydrated as concrete classes, subclass included + if not ok: + failures.append( + f"{direction} edge={edge_f!r} node={node_f!r} props={props} " + f"limit={limit}: got {sorted(ids)} expected {sorted(expected)}" + ) + count = await hub.count_nodes( + direction=direction, node=node_f, edge=edge_f, **props + ) + if count != len(expected): + failures.append( + f"count_nodes {direction} edge={edge_f!r} node={node_f!r} " + f"props={props}: {count} != {len(expected)}" + ) + assert not failures, f"{backend}: {len(failures)} mismatches, e.g.\n" + "\n".join( + failures[:10] + ) + + +async def test_subclass_filter_hydrates_subclass_instances(graph): + _ctx, _db, _backend, _ = graph + n = await _seed() + got = await n["hub"].nodes(node=QAlpha, direction="out") + assert {x.id for x in got} == {n["a1"].id, n["a2"].id, n["s1"].id, n["z"].id} + assert any(type(x) is QAlphaSub for x in got) + exact = await n["hub"].nodes(node="QAlpha", direction="out") + assert n["s1"].id not in {x.id for x in exact} + + +async def test_type_only_queries_are_one_round_trip_on_postgres(graph): + ctx, db, backend, _ = graph + if backend != "postgres": + pytest.skip("round-trip contract is Postgres-specific") + n = await _seed() + wrapped = ObservableDatabase(db, slow_query_ms=1e9) + await ctx.set_database(wrapped) + hub = n["hub"] + for edge_f, node_f, direction, limit in itertools.product( + [f for f in _EDGE_FILTERS if _type_only(f)], + [f for f in _NODE_FILTERS if _type_only(f)], + ("out", "in", "both"), + (None, 3), + ): + await ctx.clear_cache() + token = db_op_counter.set(0) + try: + await hub.nodes(direction=direction, node=node_f, edge=edge_f, limit=limit) + assert db_op_counter.get() == 1, (direction, edge_f, node_f, limit) + db_op_counter.set(0) + await hub.count_nodes(direction=direction, node=node_f, edge=edge_f) + assert db_op_counter.get() == 1, ("count", direction, edge_f, node_f) + finally: + db_op_counter.reset(token) + + +async def test_postgres_join_uses_edge_indexes(graph): + ctx, db, backend, extras = graph + if backend != "postgres": + pytest.skip("EXPLAIN check is Postgres-specific") + await _seed() # creates the Edge indexes through the normal save path + records = [] + for s in range(50): + src = generate_id("n", "QHub") + for _ in range(60): + records.append( + { + "id": generate_id("e", "QKnows"), + "entity": "QKnows", + "context": {"weight": 1}, + "source": src, + "target": generate_id("n", "QAlpha"), + "bidirectional": False, + } + ) + await db.bulk_save_detailed("edge", records) + admin, schema = extras["admin"], extras["schema"] + await admin.execute(f"ANALYZE {schema}.edge") + await admin.execute(f"ANALYZE {schema}.node") + hub_id = records[0]["source"] + for direction in ("out", "in"): + sql, params = db._connected_nodes_sql( + "node", + "edge", + hub_id, + direction=direction, + edge_entities=["QKnows"], + node_entities=["QAlpha"], + limit=20, + ) + plan = await admin.fetchval(f"EXPLAIN (FORMAT JSON) {sql}", *params) + plan_json = json.loads(plan) if isinstance(plan, str) else plan + + def _nodes(p: Any) -> List[Dict[str, Any]]: + found = [p] + for child in p.get("Plans", []): + found += _nodes(child) + return found + + scans = [ + p + for p in _nodes(plan_json[0]["Plan"]) + if p.get("Node Type") == "Seq Scan" and p.get("Relation Name") == "edge" + ] + assert not scans, f"{direction}: sequential scan on edge\n{plan_json}" + + +async def test_nodes_page_walks_every_neighbour_once(graph): + ctx, _db, _backend, _ = graph + hub = await QHub.create(name="hub") + kids = [] + for i in range(23): + kid = await QAlpha.create(name=f"n{i:02d}", score=i) + kids.append(kid) + await hub.connect(kid, edge=QKnows, weight=i) + other = await QBeta.create(name="beta") + await hub.connect(other, edge=QKnows) + + for order in (1, -1): + seen: List[str] = [] + cursor: Optional[str] = None + while True: + page, cursor = await hub.nodes_page( + node=QAlpha, + edge=[QKnows], + sort=[("context.name", order)], + cursor=cursor, + limit=5, + ) + seen += [x.name for x in page] + if cursor is None: + break + assert seen == sorted((k.name for k in kids), reverse=order == -1) + + +async def test_nodes_page_cursor_is_stable_under_inserts(graph): + ctx, _db, _backend, _ = graph + hub = await QHub.create(name="hub") + for i in range(10): + await hub.connect(await QAlpha.create(name=f"n{i:02d}"), edge=QKnows) + + first, cursor = await hub.nodes_page(sort=[("context.name", 1)], limit=4) + assert [x.name for x in first] == ["n00", "n01", "n02", "n03"] + # One neighbour lands before the cursor, one after. + await hub.connect(await QAlpha.create(name="n01x"), edge=QKnows) + await hub.connect(await QAlpha.create(name="n07x"), edge=QKnows) + + rest: List[str] = [] + while cursor: + page, cursor = await hub.nodes_page( + sort=[("context.name", 1)], cursor=cursor, limit=4 + ) + rest += [x.name for x in page] + assert rest == ["n04", "n05", "n06", "n07", "n07x", "n08", "n09"] + + +async def test_nodes_bulk_limit_per_source(graph): + _ctx, _db, _backend, _ = graph + hubs = [await QHub.create(name=f"h{i}") for i in range(3)] + for hub in hubs: + for j in range(4): + await hub.connect(await QAlpha.create(name=f"{hub.name}-{j}"), edge=QKnows) + capped = await Node.nodes_bulk( + [h.id for h in hubs], edge=[QKnows], limit_per_source=2 + ) + assert all(len(capped[h.id]) == 2 for h in hubs) + full = await Node.nodes_bulk([h.id for h in hubs], edge=[QKnows]) + assert all(len(full[h.id]) == 4 for h in hubs) + assert {x.id for x in capped[hubs[0].id]} <= {x.id for x in full[hubs[0].id]} From 905d122cdae98ffee5a4e07ba940eb31db269b03 Mon Sep 17 00:00:00 2001 From: Eldon Marks Date: Fri, 11 Sep 2026 15:14:44 -0400 Subject: [PATCH 5/9] fix(db): wrappers forward bulk_save_detailed to the backend ObservableDatabase (create_database(observe=True)) and CachingDatabase subclass Database, whose default bulk_save_detailed is a serial save() loop. Because neither wrapper defined the method, that default shadowed their __getattr__ forwarding: a wrapped Postgres COPY or Mongo bulk_write ran as one round trip per record (~0.6 ms/row on loopback). Both wrappers now forward it; the cache refreshes saved ids and drops failed ones. Co-Authored-By: Claude Opus 5 --- jvspatial/db/_cache.py | 23 ++++++++- jvspatial/db/_observable.py | 18 ++++++- tests/db/test_wrapper_bulk_save_detailed.py | 55 +++++++++++++++++++++ 3 files changed, 94 insertions(+), 2 deletions(-) create mode 100644 tests/db/test_wrapper_bulk_save_detailed.py diff --git a/jvspatial/db/_cache.py b/jvspatial/db/_cache.py index 2f80587..f447d28 100644 --- a/jvspatial/db/_cache.py +++ b/jvspatial/db/_cache.py @@ -38,7 +38,7 @@ from collections import OrderedDict from typing import Any, Dict, List, Optional, Tuple, Union -from jvspatial.db.database import Database +from jvspatial.db.database import BulkSaveResult, Database from jvspatial.runtime.serverless import is_serverless_mode logger = logging.getLogger(__name__) @@ -253,6 +253,27 @@ async def bulk_save(self, collection: str, records: List[Dict[str, Any]]) -> int self._cache_put(collection, str(rid), dict(r)) return result + async def bulk_save_detailed( + self, collection: str, records: List[Dict[str, Any]] + ) -> BulkSaveResult: + """Backend bulk path, then refresh cached entries (failed ids dropped). + + Defined explicitly: the ``Database`` default is a serial ``save`` + loop that would otherwise shadow ``__getattr__`` forwarding. + """ + result = await self.inner.bulk_save_detailed(collection, records) + if self._enabled(): + failed = set(result.failed_ids) + for r in records: + rid = r.get("id", r.get("_id")) + if rid is None: + continue + if str(rid) in failed: + self._invalidate(collection, str(rid)) + else: + self._cache_put(collection, str(rid), dict(r)) + return result + async def count( self, collection: str, diff --git a/jvspatial/db/_observable.py b/jvspatial/db/_observable.py index 68eeb59..85ad79e 100644 --- a/jvspatial/db/_observable.py +++ b/jvspatial/db/_observable.py @@ -47,7 +47,7 @@ Union, ) -from jvspatial.db.database import Database +from jvspatial.db.database import BulkSaveResult, Database from jvspatial.observability import db_op_counter from jvspatial.observability.metrics import ( MetricsRecorder, @@ -289,6 +289,22 @@ async def bulk_save(self, collection: str, records: List[Dict[str, Any]]) -> int result_count_extractor=lambda r: int(r) if r is not None else 0, ) + async def bulk_save_detailed( + self, collection: str, records: List[Dict[str, Any]] + ) -> BulkSaveResult: + """Instrumented ``bulk_save_detailed`` — the backend's native bulk path. + + Must be defined here: the ``Database`` default is a serial ``save`` + loop, which would otherwise shadow ``__getattr__`` forwarding and turn + one ``COPY`` into a round trip per record. + """ + return await self._instrument( + "bulk_save_detailed", + collection, + lambda: self.inner.bulk_save_detailed(collection, records), + result_count_extractor=lambda r: int(getattr(r, "saved", 0) or 0), + ) + async def find_one_and_delete( self, collection: str, query: Dict[str, Any] ) -> Optional[Dict[str, Any]]: diff --git a/tests/db/test_wrapper_bulk_save_detailed.py b/tests/db/test_wrapper_bulk_save_detailed.py new file mode 100644 index 0000000..40adee3 --- /dev/null +++ b/tests/db/test_wrapper_bulk_save_detailed.py @@ -0,0 +1,55 @@ +"""Wrapping adapters must forward ``bulk_save_detailed`` to the backend. + +``ObservableDatabase`` and ``CachingDatabase`` subclass ``Database``, whose +default ``bulk_save_detailed`` is a serial ``save`` loop. Without an explicit +override that default shadows ``__getattr__`` forwarding, so a wrapped +Postgres ``COPY`` silently became one round trip per record. +""" + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from jvspatial.db._cache import CachingDatabase +from jvspatial.db._observable import ObservableDatabase +from jvspatial.db.database import BulkSaveResult +from jvspatial.observability import db_op_counter + +RECORDS = [{"id": f"r{i}", "entity": "E", "context": {}} for i in range(3)] + + +def _inner() -> MagicMock: + inner = MagicMock() + inner.supports_transactions = False + inner.bulk_save_detailed = AsyncMock( + return_value=BulkSaveResult(attempted=3, saved=2, failed_ids=["r2"]) + ) + inner.save = AsyncMock(side_effect=AssertionError("per-record save used")) + return inner + + +@pytest.mark.asyncio +async def test_observable_forwards_bulk_save_detailed_as_one_op(): + inner = _inner() + token = db_op_counter.set(0) + try: + result = await ObservableDatabase(inner).bulk_save_detailed("node", RECORDS) + assert db_op_counter.get() == 1 + finally: + db_op_counter.reset(token) + assert result.saved == 2 + inner.bulk_save_detailed.assert_awaited_once_with("node", RECORDS) + + +@pytest.mark.asyncio +async def test_caching_forwards_bulk_save_detailed_and_skips_failed_ids(monkeypatch): + monkeypatch.setenv("SERVERLESS_MODE", "false") + inner = _inner() + inner.get = AsyncMock(return_value=None) + cached = CachingDatabase(inner) + result = await cached.bulk_save_detailed("node", RECORDS) + inner.bulk_save_detailed.assert_awaited_once_with("node", RECORDS) + assert result.failed_ids == ["r2"] + assert await cached.get("node", "r0") == RECORDS[0] # served from cache + assert await cached.get("node", "r2") is None # failed id not cached + inner.get.assert_awaited_once_with("node", "r2") From 844a951ad93305147d23909b1e62dec9413c4fa5 Mon Sep 17 00:00:00 2001 From: Eldon Marks Date: Fri, 11 Sep 2026 15:19:50 -0400 Subject: [PATCH 6/9] Entity-leading Postgres indexes, optional GIN, $text search Phase 3 of the hub-node scale remediation. - PostgresDB.create_index indexes entity/id/tenant_id as the real columns the translator compares, emits DESC NULLS LAST (the order find(sort=...) uses), and rebuilds indexes defined by the old rules in place (the edge (source, target, entity) unique index indexed data->entity, which no query used) - ensure_indexes scopes per-class (annotation-declared) indexes to the class's entity: (entity, ) by default, WHERE entity = ... with partial_by_entity; the unscoped pre-0.0.18 index is dropped once replaced. attribute()/compound_index() gain the scoping options - gin_index="off" / JVSPATIAL_PG_GIN_INDEX=off skips the whole-document GIN on new collections; find() warns once for $all/$elemMatch then - $text: to_tsvector('simple', ...) @@ plainto_tsquery on Postgres, GIN via @fulltext_index / attribute(fulltext=True); QueryEngine evaluates the same semantics in memory; MongoDB strips $fields - escape_regex() helper; find_edges_between(limit=) - Bench: typed-find gate on a 1M-row shared node table Bench: sorted, limited typed find at 1M rows, GIN off: p95 9.2 -> 3.0 ms, one index scan and no sort (0.0.17 needed a BitmapAnd + sort). Co-Authored-By: Claude Opus 5 --- CHANGELOG.md | 29 ++ SPEC.md | 6 +- docs/bench/2026-09-hub-node-baseline.md | 47 ++- docs/bench/2026-09-hub-node-phase0.jsonl | 1 + docs/bench/2026-09-hub-node-phase3.jsonl | 4 + docs/md/environment-keys-reference.md | 1 + docs/md/postgres-guide.md | 57 +++- jvspatial/core/annotations.py | 91 +++++- jvspatial/core/context.py | 48 ++- jvspatial/core/entities/object.py | 35 ++ jvspatial/db/README.md | 2 + jvspatial/db/__init__.py | 2 + jvspatial/db/_postgres_translate.py | 45 ++- jvspatial/db/mongodb.py | 19 +- jvspatial/db/postgres.py | 198 +++++++++-- jvspatial/db/query.py | 64 +++- jvspatial/env_adapter.py | 1 + tests/benchmarks/hub_bench_report.py | 29 +- tests/benchmarks/test_hub_node_bench.py | 121 +++++++ tests/db/test_postgres_indexes_text.py | 397 +++++++++++++++++++++++ 20 files changed, 1140 insertions(+), 57 deletions(-) create mode 100644 docs/bench/2026-09-hub-node-phase3.jsonl create mode 100644 tests/db/test_postgres_indexes_text.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 5aff56f..ccd1f0a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -29,6 +29,17 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 `GraphContext.find_page` via `jvspatial.core.pager`), and `nodes_bulk(limit_per_source=...)` (one windowed query on Postgres). `count_neighbors` now delegates to `count_nodes`. +- **`$text` full-text search** (`jvspatial/db/_postgres_translate.py`, + `jvspatial/db/query.py`): `{"$text": {"$search": "...", "$fields": [...]}}` + becomes `to_tsvector('simple', …) @@ plainto_tsquery('simple', …)` on + Postgres, backed by a GIN index declared with `@fulltext_index([...])` or + `attribute(fulltext=True)`; other backends evaluate the same semantics in + memory. Previously `$text` fell back to a full scan and raised in memory. +- **`jvspatial.db.escape_regex()`** for building literal `$regex` patterns + from user input (`$regex` is never index-backed). +- **`find_edges_between(limit=...)`** pushes the limit to the database. +- **`gin_index="off"` / `JVSPATIAL_PG_GIN_INDEX=off`** skips the + whole-document `GIN (data jsonb_path_ops)` on new Postgres collections. - **`expand_node` keyset paging** (`jvspatial/core/graph_expansion.py`): `after=` / `pagination.next_after` (and the `after` query parameter on the graph expand endpoint) page incident edges by edge id in O(page). The @@ -56,6 +67,18 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 form took the join and everything else loaded every incident edge and hydrated every neighbour before slicing in Python. Backends without the pushdown keep the Python path, now filtering in the database `find`. +- **Postgres per-class indexes are `entity`-leading** (`PostgresDB.create_index`, + `GraphContext.ensure_indexes`). Indexes declared with + `attribute(indexed=True)` / `@compound_index` become `(entity, )` + (`_entity__idx`, shared by classes declaring the same fields), + or `WHERE entity = ''` with `index_partial_by_entity=True` / + `partial_by_entity=True`; the unscoped pre-0.0.18 index of the same fields + is dropped once its replacement exists. Descending keys are now + `DESC NULLS LAST` so sorted + limited finds walk the index, and `entity` / + `id` / `tenant_id` are indexed as the real columns the translator compares + (the edge `(source, target, entity)` unique index previously indexed + `data->entity`, which no query used). Indexes built by the old rules are + rebuilt in place on the next `ensure_indexes`. - **`expand_node` / `subgraph_bfs` source incident edges from the edge collection** in every mode (one query per node instead of one `get` per edge id); `total_edge_count` is a count of incident edges, and neighbours @@ -68,6 +91,12 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 `edge=[{"E": {...}}]` returned neighbours reached through *any* edge type, and `direction="both"` capped at 10 000 edges. Edge types and criteria are now always applied. +- **`observe=True` / caching wrappers turned `bulk_save_detailed` into one + round trip per record** (`jvspatial/db/_observable.py`, + `jvspatial/db/_cache.py`). Both wrappers subclass `Database`, whose default + `bulk_save_detailed` is a serial `save` loop; it shadowed `__getattr__` + forwarding, so a wrapped Postgres `COPY` (or Mongo `bulk_write`) never ran. + Both now forward to the backend (the cache also refreshes saved ids). - **`ObservableDatabase` advertised graph pushdowns its backend lacks** (`jvspatial/db/_observable.py`). `find_connected_nodes` / `traverse` were plain methods, so `getattr(db, "find_connected_nodes", None)` looked diff --git a/SPEC.md b/SPEC.md index 2ae5282..68df27b 100644 --- a/SPEC.md +++ b/SPEC.md @@ -291,11 +291,13 @@ No built-in migration framework. Adapters do not enforce schemas. Adding optiona | `$in`, `$nin` | Membership | | `$exists` | Field presence | | `$and`, `$or` | Logical combinators | -| `$regex` | Regex match (string fields) | +| `$regex` | Regex match (string fields). Never index-backed; build patterns from user input with `jvspatial.db.escape_regex` | +| `$text` | Top-level `{"$text": {"$search": "words", "$fields": ["context.a", ...]}}`: every search word (case-insensitive, `\w+` tokens, no stemming) occurs in the concatenated fields; without `$fields` every string value is searched | ### 5.2 Pushdown vs in-memory -- **MongoDB**: native pushdown; queries run server-side. +- **Postgres**: `translate_query` pushes the whole operator surface into JSONB SQL; `$text` becomes `to_tsvector('simple', …) @@ plainto_tsquery('simple', …)` (requires `$fields`; the GIN from `@fulltext_index` / `attribute(fulltext=True)` serves it when the fields match in order). Per-class indexes are `(entity, )` (or `WHERE entity = …` with `partial_by_entity`), descending keys `DESC NULLS LAST`, so typed `find(sort=…, limit=…)` walks an index; the whole-document GIN is optional (`JVSPATIAL_PG_GIN_INDEX`). Untranslatable queries fall back to a full scan + `QueryEngine.match`. +- **MongoDB**: native pushdown; queries run server-side (`$text` uses the collection's text index; `$fields` is stripped). - **SQLite**: translated to SQL via `SQLiteTranslator` (subset; complex `$or` chains may fall back). - **DynamoDB**: limited pushdown via `Select=COUNT` and key conditions; remainder filtered client-side. - **JSON**: full in-memory evaluation after loading matching collection. diff --git a/docs/bench/2026-09-hub-node-baseline.md b/docs/bench/2026-09-hub-node-baseline.md index 3a7d6a2..8b6c465 100644 --- a/docs/bench/2026-09-hub-node-baseline.md +++ b/docs/bench/2026-09-hub-node-baseline.md @@ -188,6 +188,49 @@ Code: `dc2a6cc` + the Phase 2 change set. Page instead: `nodes_page(...)`. - The Phase 1 numbers (connect, save, concurrency, sizes) are unchanged. -## After Phase 3 +## After Phase 3 — entity-leading indexes, optional GIN, `$text` -_Pending._ +Code: `905d122` + the Phase 3 change set. + +**Typed-find gate.** `test_typed_find_is_index_bound` runs +`find({"entity": "BenchEntry", "context.track_id": t}, sort=[("context.created_at", -1)], limit=20)` +50 times on a shared `node` table of 1M rows. The rows are spread over ten +entities, and `BenchEntry` holds 100k of them: 1000 tracks × 100. The class +declares `@compound_index([("track_id", 1), ("created_at", -1)])`. + +| typed find (sorted, limit 20) | rows | p50 / p95 ms | index walked | sort node | node_data_gin | +|---|---|---|---|---|---| +| 0.0.17 `8fc5138` | 1,000,000 | 4.44 / 9.24 | node_context_track_id_context_created_at_idx, node_entity_idx | yes | present | +| Phase 3 `905d122` | 1,000,000 | 0.72 / 2.96 | node_entity_context_track_id_context_created_at_idx | no | off | + +**Gate: passed.** p95 is 2.96 ms, well under the 10 ms bar, and the plan is +index-bound with the whole-document GIN off. One scan of +`(entity, track_id, created_at DESC NULLS LAST)` returns the first 20 rows +with no Sort node. On 0.0.17 the same class index skipped `entity` and +ordered descending keys `NULLS FIRST`. Postgres therefore had to AND it with +the entity index and sort every matching row. At this track size that still +stays under 10 ms locally, but the cost grows with the rows per track. + +Hub-node numbers on the same code: + +| operation | 1k p50 / p95 ms (trips) | 10k p50 / p95 ms (trips) | 100k p50 / p95 ms (trips) | +|---|---|---|---| +| `ctx.get(Hub)` (hydrate hub) | 0.6 / 2.1 | 0.8 / 2.3 | 0.8 / 2.9 | +| `hub.connect(leaf, edge=E)` | 1.7 / 3.0 (3) | 2.0 / 4.0 (3) | 2.7 / 4.2 (3) | +| `hub.save()` after scalar change | 0.8 / 2.1 | 0.7 / 1.3 | 1.2 / 2.4 | +| `hub.nodes(edge=[E], node=['Leaf'], limit=20)` | 1.1 / 2.4 (1) | 1.2 / 2.7 (1) | 1.5 / 3.8 (1) | +| `hub.nodes(edge=E, limit=20)` | 1.0 / 2.4 (1) | 1.3 / 2.5 (1) | 3.6 / 6.0 (1) | +| `sink.nodes(edge=[E], node=['Leaf'], direction='in', limit=20)` | 1.1 / 1.4 (1) | 1.3 / 3.7 (1) | 2.1 / 6.8 (1) | +| `sink.nodes(edge=[E], node=['Leaf'], direction='in')` | 19.2 / 47.8 (1) | 237.7 / 248.8 (1) | 2118.2 / 2203.4 (1) | +| `len(await hub.nodes(edge=[E]))` | 20.5 / 48.5 (1) | 239.0 / 254.5 (1) | 2086.3 / 2092.6 (1) | +| `hub.count_nodes(edge=[E])` | 3.5 / 4.4 (1) | 35.7 / 37.8 (1) | 89.2 / 141.5 (1) | +| 32× concurrent `connect()` (ms) | 16 wall / 16 max | 16 wall / 15 max | 15 wall / 14 max | + +The Phase 1 and 2 results hold. The Edge `(source, target, entity)` unique +index is now built on the real `entity` column. Its definition changed, but +traversal latencies stay within run-to-run noise. + +Seeding note: the Phase 3 hub tier seeds through `create_database(observe=True)`, +which now forwards `bulk_save_detailed` to COPY (fixed in `905d122`). Earlier +tiers seeded per record; that only affected setup time, not the measured +operations. diff --git a/docs/bench/2026-09-hub-node-phase0.jsonl b/docs/bench/2026-09-hub-node-phase0.jsonl index a730af5..810fd3d 100644 --- a/docs/bench/2026-09-hub-node-phase0.jsonl +++ b/docs/bench/2026-09-hub-node-phase0.jsonl @@ -1,3 +1,4 @@ {"degree": 1000, "jvspatial": "0.0.17", "git_sha": "8fc5138", "edge_ids_mode": "persist", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 2.25, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 26633, "node_gin_bytes": 196608, "node_total_bytes": 802816, "edge_total_bytes": 1892352}, "hub_get": {"p50_ms": 1.044, "p95_ms": 1.479, "max_ms": 2.615, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 40.666, "p95_ms": 64.47, "max_ms": 65.79, "n": 20, "round_trips": 3}, "nodes_class_out_limit20": {"p50_ms": 1.115, "p95_ms": 2.045, "max_ms": 2.392, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 45.583, "p95_ms": 72.069, "max_ms": 75.564, "n": 20, "round_trips": 3}, "nodes_list_in_unlimited": {"p50_ms": 39.259, "p95_ms": 70.509, "max_ms": 77.64, "n": 20, "round_trips": 3}, "count_via_len_nodes": {"p50_ms": 43.309, "p95_ms": 68.951, "max_ms": 71.525, "n": 20, "round_trips": 3}, "hub_save": {"p50_ms": 11.494, "p95_ms": 13.597, "max_ms": 24.191, "n": 50}, "hub_connect": {"p50_ms": 6.607, "p95_ms": 7.902, "max_ms": 10.426, "n": 50, "round_trips": 5}, "concurrent_connect_32": {"wall_ms": 133.819, "max_call_ms": 132.509, "p50_call_ms": 75.482}, "sizes_after_writes": {"hub_row_bytes": 28493, "node_gin_bytes": 3440640, "node_total_bytes": 7929856, "edge_total_bytes": 2015232}} {"degree": 10000, "jvspatial": "0.0.17", "git_sha": "8fc5138", "edge_ids_mode": "persist", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 23.26, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 263664, "node_gin_bytes": 1638400, "node_total_bytes": 6668288, "edge_total_bytes": 17047552}, "hub_get": {"p50_ms": 3.945, "p95_ms": 5.141, "max_ms": 5.279, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 439.637, "p95_ms": 485.561, "max_ms": 488.057, "n": 20, "round_trips": 21}, "nodes_class_out_limit20": {"p50_ms": 1.137, "p95_ms": 1.708, "max_ms": 2.802, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 474.404, "p95_ms": 511.823, "max_ms": 513.261, "n": 20, "round_trips": 21}, "nodes_list_in_unlimited": {"p50_ms": 470.006, "p95_ms": 523.042, "max_ms": 571.798, "n": 20, "round_trips": 21}, "count_via_len_nodes": {"p50_ms": 472.215, "p95_ms": 510.991, "max_ms": 527.017, "n": 20, "round_trips": 21}, "hub_save": {"p50_ms": 110.905, "p95_ms": 137.193, "max_ms": 231.894, "n": 50}, "hub_connect": {"p50_ms": 23.168, "p95_ms": 52.342, "max_ms": 336.508, "n": 50, "round_trips": 5}, "concurrent_connect_32": {"wall_ms": 813.037, "max_call_ms": 812.608, "p50_call_ms": 430.787}, "sizes_after_writes": {"hub_row_bytes": 256983, "node_gin_bytes": 8962048, "node_total_bytes": 49291264, "edge_total_bytes": 19529728}} {"degree": 100000, "jvspatial": "0.0.17", "git_sha": "8fc5138", "edge_ids_mode": "persist", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 262.07, "read_iterations": 5, "sizes_after_seed": {"hub_row_bytes": 2634003, "node_gin_bytes": 26787840, "node_total_bytes": 76308480, "edge_total_bytes": 171909120}, "hub_get": {"p50_ms": 32.39, "p95_ms": 47.729, "max_ms": 47.729, "n": 5}, "nodes_list_out_limit20": {"p50_ms": 4936.08, "p95_ms": 5184.805, "max_ms": 5184.805, "n": 5, "round_trips": 201}, "nodes_class_out_limit20": {"p50_ms": 1.344, "p95_ms": 3.033, "max_ms": 3.033, "n": 5, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 5317.014, "p95_ms": 5338.701, "max_ms": 5338.701, "n": 5, "round_trips": 201}, "nodes_list_in_unlimited": {"p50_ms": 5218.259, "p95_ms": 5379.298, "max_ms": 5379.298, "n": 5, "round_trips": 201}, "count_via_len_nodes": {"p50_ms": 4807.526, "p95_ms": 4881.467, "max_ms": 4881.467, "n": 5, "round_trips": 201}, "hub_save": {"p50_ms": 1160.649, "p95_ms": 1661.605, "max_ms": 1746.441, "n": 50}, "hub_connect": {"p50_ms": 160.187, "p95_ms": 581.505, "max_ms": 597.428, "n": 50, "round_trips": 5}, "concurrent_connect_32": {"wall_ms": 10276.447, "max_call_ms": 10244.43, "p50_call_ms": 4726.249}, "sizes_after_writes": {"hub_row_bytes": 2464513, "node_gin_bytes": 49561600, "node_total_bytes": 363249664, "edge_total_bytes": 171966464}} +{"kind": "typed_find", "rows": 1000000, "jvspatial": "0.0.17", "git_sha": "8fc5138", "seed_seconds": 27.58, "node_data_gin": true, "typed_find_sorted_limit20": {"p50_ms": 4.439, "p95_ms": 9.244, "max_ms": 10.754, "n": 50}, "plan_indexes": ["node_context_track_id_context_created_at_idx", "node_entity_idx"], "plan_sorts": true} diff --git a/docs/bench/2026-09-hub-node-phase3.jsonl b/docs/bench/2026-09-hub-node-phase3.jsonl new file mode 100644 index 0000000..7375873 --- /dev/null +++ b/docs/bench/2026-09-hub-node-phase3.jsonl @@ -0,0 +1,4 @@ +{"degree": 1000, "jvspatial": "0.0.17", "git_sha": "905d122", "edge_ids_mode": "derive", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 0.11, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 134, "node_gin_bytes": 106496, "node_total_bytes": 516096, "edge_total_bytes": 1908736}, "hub_get": {"p50_ms": 0.59, "p95_ms": 2.13, "max_ms": 2.229, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 1.056, "p95_ms": 2.356, "max_ms": 2.56, "n": 20, "round_trips": 1}, "nodes_class_out_limit20": {"p50_ms": 0.991, "p95_ms": 2.388, "max_ms": 3.348, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 1.1, "p95_ms": 1.449, "max_ms": 2.26, "n": 20, "round_trips": 1}, "nodes_list_in_unlimited": {"p50_ms": 19.172, "p95_ms": 47.839, "max_ms": 50.488, "n": 20, "round_trips": 1}, "count_via_len_nodes": {"p50_ms": 20.52, "p95_ms": 48.454, "max_ms": 48.999, "n": 20, "round_trips": 1}, "count_nodes": {"p50_ms": 3.509, "p95_ms": 4.355, "max_ms": 5.324, "n": 20, "round_trips": 1}, "hub_save": {"p50_ms": 0.762, "p95_ms": 2.085, "max_ms": 3.39, "n": 50}, "hub_connect": {"p50_ms": 1.689, "p95_ms": 2.961, "max_ms": 3.275, "n": 50, "round_trips": 3}, "concurrent_connect_32": {"wall_ms": 16.353, "max_call_ms": 15.639, "p50_call_ms": 14.351}, "sizes_after_writes": {"hub_row_bytes": 136, "node_gin_bytes": 114688, "node_total_bytes": 524288, "edge_total_bytes": 2015232}} +{"degree": 10000, "jvspatial": "0.0.17", "git_sha": "905d122", "edge_ids_mode": "derive", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 0.88, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 134, "node_gin_bytes": 827392, "node_total_bytes": 4120576, "edge_total_bytes": 17203200}, "hub_get": {"p50_ms": 0.833, "p95_ms": 2.251, "max_ms": 8.001, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 1.211, "p95_ms": 2.681, "max_ms": 6.026, "n": 20, "round_trips": 1}, "nodes_class_out_limit20": {"p50_ms": 1.303, "p95_ms": 2.492, "max_ms": 2.692, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 1.341, "p95_ms": 3.727, "max_ms": 4.264, "n": 20, "round_trips": 1}, "nodes_list_in_unlimited": {"p50_ms": 237.715, "p95_ms": 248.752, "max_ms": 256.754, "n": 20, "round_trips": 1}, "count_via_len_nodes": {"p50_ms": 238.995, "p95_ms": 254.522, "max_ms": 257.119, "n": 20, "round_trips": 1}, "count_nodes": {"p50_ms": 35.691, "p95_ms": 37.796, "max_ms": 43.161, "n": 20, "round_trips": 1}, "hub_save": {"p50_ms": 0.697, "p95_ms": 1.31, "max_ms": 5.451, "n": 50}, "hub_connect": {"p50_ms": 1.958, "p95_ms": 4.044, "max_ms": 6.656, "n": 50, "round_trips": 3}, "concurrent_connect_32": {"wall_ms": 15.534, "max_call_ms": 14.86, "p50_call_ms": 13.567}, "sizes_after_writes": {"hub_row_bytes": 136, "node_gin_bytes": 835584, "node_total_bytes": 4136960, "edge_total_bytes": 17252352}} +{"degree": 100000, "jvspatial": "0.0.17", "git_sha": "905d122", "edge_ids_mode": "derive", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 10.38, "read_iterations": 5, "sizes_after_seed": {"hub_row_bytes": 134, "node_gin_bytes": 10944512, "node_total_bytes": 43704320, "edge_total_bytes": 169295872}, "hub_get": {"p50_ms": 0.835, "p95_ms": 2.888, "max_ms": 2.888, "n": 5}, "nodes_list_out_limit20": {"p50_ms": 1.464, "p95_ms": 3.849, "max_ms": 3.849, "n": 5, "round_trips": 1}, "nodes_class_out_limit20": {"p50_ms": 3.551, "p95_ms": 5.971, "max_ms": 5.971, "n": 5, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 2.058, "p95_ms": 6.78, "max_ms": 6.78, "n": 5, "round_trips": 1}, "nodes_list_in_unlimited": {"p50_ms": 2118.187, "p95_ms": 2203.427, "max_ms": 2203.427, "n": 5, "round_trips": 1}, "count_via_len_nodes": {"p50_ms": 2086.288, "p95_ms": 2092.644, "max_ms": 2092.644, "n": 5, "round_trips": 1}, "count_nodes": {"p50_ms": 89.191, "p95_ms": 141.537, "max_ms": 141.537, "n": 5, "round_trips": 1}, "hub_save": {"p50_ms": 1.156, "p95_ms": 2.406, "max_ms": 3.036, "n": 50}, "hub_connect": {"p50_ms": 2.668, "p95_ms": 4.159, "max_ms": 4.714, "n": 50, "round_trips": 3}, "concurrent_connect_32": {"wall_ms": 14.692, "max_call_ms": 14.23, "p50_call_ms": 13.189}, "sizes_after_writes": {"hub_row_bytes": 136, "node_gin_bytes": 15622144, "node_total_bytes": 48398336, "edge_total_bytes": 169877504}} +{"kind": "typed_find", "rows": 1000000, "jvspatial": "0.0.17", "git_sha": "905d122", "seed_seconds": 24.82, "node_data_gin": false, "typed_find_sorted_limit20": {"p50_ms": 0.719, "p95_ms": 2.965, "max_ms": 3.626, "n": 50}, "plan_indexes": ["node_entity_context_track_id_context_created_at_idx"], "plan_sorts": false} diff --git a/docs/md/environment-keys-reference.md b/docs/md/environment-keys-reference.md index f6ff44d..d8cb69b 100644 --- a/docs/md/environment-keys-reference.md +++ b/docs/md/environment-keys-reference.md @@ -51,6 +51,7 @@ For full examples and default values, see: - `JVSPATIAL_DYNAMODB_ENDPOINT_URL` - DynamoDB endpoint (e.g. LocalStack). - `JVSPATIAL_DYNAMODB_WAIT_FOR_INDEX` - Wait-for-index toggle. - `JVSPATIAL_NODE_EDGE_IDS` - Node adjacency mode: `derive` (edge collection only; Postgres / MongoDB / SQLite default) or `persist` (node rows also store `edges`; JSON / DynamoDB default). An `edge_ids_mode` set on the adapter instance wins. +- `JVSPATIAL_PG_GIN_INDEX` - Postgres whole-document `GIN (data jsonb_path_ops)` index on new collections: `full` (default) or `off`. ### Auth and rate limit - `JVSPATIAL_AUTH_ENABLED` - Enables auth. diff --git a/docs/md/postgres-guide.md b/docs/md/postgres-guide.md index 95ed1bc..98b55dc 100644 --- a/docs/md/postgres-guide.md +++ b/docs/md/postgres-guide.md @@ -131,8 +131,60 @@ CREATE INDEX _tenant_idx ``` The full record lives in the `data` JSONB blob. `id`, `entity`, and `tenant_id` -are denormalized for indexing. The default GIN index on `data` accelerates -arbitrary JSONB containment / path queries. +are denormalized for indexing. The whole-document GIN index on `data` only +accelerates containment-shaped predicates (`$all`); see +[Index policy](#index-policy) for when to turn it off. + +### Index policy + +- **Per-class indexes lead with `entity`.** Every typed `find()` filters on + the `entity` column of a table shared by every class, so an index declared + with `attribute(indexed=True)` or `@compound_index(...)` is created as + `(entity, )` and named `_entity__idx`. Sibling classes + that declare the same fields share it. Pass `index_partial_by_entity=True` + (or `@compound_index(..., partial_by_entity=True)`) for a smaller + `WHERE entity = ''` index per class instead, or + `entity_leading=False` to keep the key unscoped. The pre-0.0.18 unscoped + index of the same fields is dropped when its replacement is created. +- **Descending keys are `DESC NULLS LAST`**, the order `find(sort=...)` + emits, so `sort=[("context.created_at", -1)], limit=20` walks the index + with no sort step. Indexes created by 0.0.17 and earlier (descending keys + without `NULLS LAST`, or `entity` / `id` indexed as JSONB paths — e.g. the + edge `(source, target, entity)` unique index) are rebuilt in place the next + time `ensure_indexes` runs. +- **The whole-document GIN is optional.** `JVSPATIAL_PG_GIN_INDEX=off` (or + `create_database("postgres", ..., gin_index="off")`) stops new + collections from creating `_data_gin`. Every node rewrite re-indexes + the whole document into that GIN, while equality / range / sort use the + functional B-trees above. Large deployments should turn it off, then drop + an existing one with `DROP INDEX CONCURRENTLY node_data_gin;` — + `find()` logs a one-time warning if an `$all` / `$elemMatch` query then + runs without it. + +### Full-text search + +```python +from jvspatial.core.annotations import attribute, fulltext_index + +@fulltext_index(["title", "body"]) # one GIN over both fields +class Article(Node): + title: str = "" + body: str = "" + summary: str = attribute(fulltext=True, default="") # or per-attribute + +await Article.find({"$text": {"$search": "graph database", + "$fields": ["context.title", "context.body"]}}) +``` + +`$text` becomes `to_tsvector('simple', …) @@ plainto_tsquery('simple', …)`: +every search word must appear (case-insensitive, no stemming — "database" +does not match "databases"). The index is used when `$fields` lists the same +fields in the same order as the declaration. Other backends evaluate the same +semantics in memory; MongoDB searches its own text index. + +`$regex` is never index-backed. When a pattern has to include user input, pass +the input through `jvspatial.db.escape_regex` so metacharacters match +literally (and a crafted pattern cannot make every scan slow). ### Custom indexes @@ -284,6 +336,7 @@ precedence): | `JVSPATIAL_POSTGRES_MAX_POOL_SIZE` | Pool max size override | | `JVSPATIAL_POSTGRES_POOLER_MODE` | `"session"` (default) or `"transaction"` | | `JVSPATIAL_NODE_EDGE_IDS` | `"derive"` (Postgres default) or `"persist"` | +| `JVSPATIAL_PG_GIN_INDEX` | `"full"` (default) or `"off"` — whole-document GIN | ## Operational tips diff --git a/jvspatial/core/annotations.py b/jvspatial/core/annotations.py index bab9a2c..dae0265 100644 --- a/jvspatial/core/annotations.py +++ b/jvspatial/core/annotations.py @@ -57,6 +57,9 @@ def attribute( index_unique: bool = False, index_direction: int = 1, index_partial_filter_expression: Optional[Dict[str, Any]] = None, + index_entity_leading: bool = True, + index_partial_by_entity: bool = False, + fulltext: bool = False, # Standard Pydantic Field parameters description: Optional[str] = None, title: Optional[str] = None, @@ -87,6 +90,17 @@ def attribute( ``context.`` DB path. Use this with ``index_unique=True`` to avoid null/empty-value conflicts in shared collections (e.g. ``{"context.name": {"$gt": ""}}``). Supersedes sparse behavior. + index_entity_leading: Postgres: lead the index key with the real + ``entity`` column (default) — every typed ``find()`` filters on + it in the collection shared by every class. + index_partial_by_entity: Postgres: scope the index to + ``WHERE entity = ''`` instead of leading with the + column (smaller; one index per class). + fulltext: Include this field in the class's full-text index (a + Postgres GIN over ``to_tsvector('simple', ...)`` of every + ``fulltext=True`` field, in declaration order), matched by + ``{"$text": {"$search": ..., "$fields": [...]}}`` with the + fields in that same order. description: Description for the attribute title: Title for the attribute examples: Example values for documentation @@ -147,7 +161,7 @@ def attribute( field_kwargs["default"] = default # Add protection/transient/index metadata to json_schema_extra - if protected or transient or indexed: + if protected or transient or indexed or fulltext: json_extra = field_kwargs.get("json_schema_extra", {}) if protected: json_extra["protected"] = True @@ -163,6 +177,12 @@ def attribute( json_extra["index_partial_filter_expression"] = ( index_partial_filter_expression ) + if not index_entity_leading: + json_extra["index_entity_leading"] = False + if index_partial_by_entity: + json_extra["index_partial_by_entity"] = True + if fulltext: + json_extra["fulltext"] = True field_kwargs["json_schema_extra"] = json_extra return Field(**field_kwargs) @@ -312,17 +332,41 @@ def get_indexed_fields(cls: Type) -> Dict[str, Dict[str, Any]]: field_config["partial_filter_expression"] = json_extra[ "index_partial_filter_expression" ] + if "index_entity_leading" in json_extra: + field_config["entity_leading"] = json_extra[ + "index_entity_leading" + ] + if "index_partial_by_entity" in json_extra: + field_config["partial_by_entity"] = json_extra[ + "index_partial_by_entity" + ] indexed_fields[field_name] = field_config return indexed_fields +def get_fulltext_fields(cls: Type) -> List[str]: + """Fields declared with ``attribute(fulltext=True)``, in declaration order.""" + names: List[str] = [] + for field_name, field_info in getattr(cls, "model_fields", {}).items(): + json_extra = getattr(field_info, "json_schema_extra", None) + if callable(json_extra): + schema: Dict[str, Any] = {} + json_extra(schema, cls) + json_extra = schema + if json_extra and json_extra.get("fulltext", False): + names.append(field_name) + return names + + def compound_index( fields: List[Tuple[str, int]], name: Optional[str] = None, unique: bool = False, sparse: bool = False, partial_filter_expression: Optional[Dict[str, Any]] = None, + entity_leading: bool = True, + partial_by_entity: bool = False, ): """Class decorator for declaring compound indexes. @@ -344,6 +388,11 @@ def compound_index( a unique index to a specific document sub-type within a shared collection. Field paths in the expression must use the full ``context.`` DB path (they are NOT auto-prefixed). + entity_leading: Postgres: lead the key with the real ``entity`` + column (default), so typed finds on a shared collection + seek straight to this class's rows. + partial_by_entity: Postgres: scope the index to + ``WHERE entity = ''`` instead (smaller; one per class). Returns: Class decorator function @@ -368,12 +417,52 @@ def decorator(cls: Type) -> Type: } if partial_filter_expression is not None: index_def["partialFilterExpression"] = partial_filter_expression + if not entity_leading: + index_def["entity_leading"] = False + if partial_by_entity: + index_def["partial_by_entity"] = True _COMPOUND_INDEXES[cls].append(index_def) return cls return decorator +_FULLTEXT_INDEXES: Dict[Type, List[Dict[str, Any]]] = {} + + +def fulltext_index(fields: List[str], name: Optional[str] = None): + """Class decorator declaring a full-text index over ``fields``. + + On Postgres this is a GIN index over ``to_tsvector('simple', ...)`` of the + concatenated field values, scoped to the class's rows. Queries use it via + ``{"$text": {"$search": "words", "$fields": ["context.title", ...]}}`` + with the same fields in the same order. Other backends evaluate ``$text`` + in memory (MongoDB searches its own text index). + + Example: + @fulltext_index(["title", "body"]) + class Article(Node): + title: str = "" + body: str = "" + """ + + def decorator(cls: Type) -> Type: + _FULLTEXT_INDEXES.setdefault(cls, []).append( + {"fields": list(fields), "name": name or f"fts_{'_'.join(fields)}"} + ) + return cls + + return decorator + + +def get_fulltext_indexes(cls: Type) -> List[Dict[str, Any]]: + """Full-text indexes declared with :func:`fulltext_index` across the MRO.""" + indexes: List[Dict[str, Any]] = [] + for klass in cls.__mro__: + indexes.extend(_FULLTEXT_INDEXES.get(klass, [])) + return indexes + + def get_compound_indexes(cls: Type) -> List[Dict[str, Any]]: """Get all compound indexes declared for a class. diff --git a/jvspatial/core/context.py b/jvspatial/core/context.py index 9f3dbea..227a650 100644 --- a/jvspatial/core/context.py +++ b/jvspatial/core/context.py @@ -1675,29 +1675,53 @@ async def ensure_indexes(self, entity_class: Type[T]) -> None: _ensured_indexes.add(collection_key) return # Database doesn't support indexing + # Per-class (annotation-declared) indexes are scoped to the class's + # entity on Postgres: entity-leading or entity-partial, replacing the + # unscoped pre-0.0.18 index of the same fields. The scoping keys are + # Postgres-only; other adapters receive the plain definition, and + # full-text definitions are skipped where ``$text`` runs in memory. + is_postgres = self._is_postgres(self.database) + scope_keys = ("per_class", "entity_leading", "partial_by_entity", "fulltext") + + def _scoped(index_def: Dict[str, Any], extra: Dict[str, Any]) -> Dict[str, Any]: + if is_postgres and index_def.get("per_class"): + extra.update( + entity=entity_class._entity_name(), + drop_legacy=True, + entity_leading=index_def.get("entity_leading", True), + partial_by_entity=index_def.get("partial_by_entity", False), + fulltext=index_def.get("fulltext", False), + ) + return extra + # Create each index for index_def in indexes: + if index_def.get("fulltext") and not is_postgres: + continue try: if "field" in index_def: # Single-field index; pass through name and extra kwargs extra = { k: v for k, v in index_def.items() - if k not in ("field", "unique", "direction") + if k not in ("field", "unique", "direction", *scope_keys) } await self.database.create_index( collection, index_def["field"], unique=index_def.get("unique", False), - **extra, + **_scoped(index_def, extra), ) elif "fields" in index_def: # Compound index; pass through name and other create_index kwargs - extra = { - k: v - for k, v in index_def.items() - if k not in ("fields", "unique") - } + extra = _scoped( + index_def, + { + k: v + for k, v in index_def.items() + if k not in ("fields", "unique", *scope_keys) + }, + ) await self.database.create_index( collection, index_def["fields"], @@ -1715,7 +1739,12 @@ async def ensure_indexes(self, entity_class: Type[T]) -> None: _ensured_indexes.add(collection_key) async def find_edges_between( - self, source_id: str, target_id: Optional[str] = None, edge_class=None, **kwargs + self, + source_id: str, + target_id: Optional[str] = None, + edge_class=None, + limit: Optional[int] = None, + **kwargs, ) -> List: """Find edges between nodes using database queries. @@ -1723,6 +1752,7 @@ async def find_edges_between( source_id: Source node ID target_id: Target node ID (optional) edge_class: Edge class to filter by + limit: Maximum number of edges to return (pushed to the database) **kwargs: Additional edge properties to match Returns: @@ -1756,7 +1786,7 @@ async def find_edges_between( collection = self._get_collection_name(self._get_entity_type_code(edge_cls)) db = self.database - results = await db.find(collection, query) + results = await db.find(collection, query, limit=limit) edges = [] for data in results: diff --git a/jvspatial/core/entities/object.py b/jvspatial/core/entities/object.py index f35ec87..acbef62 100644 --- a/jvspatial/core/entities/object.py +++ b/jvspatial/core/entities/object.py @@ -21,6 +21,8 @@ AttributeMixin, attribute, get_compound_indexes, + get_fulltext_fields, + get_fulltext_indexes, get_indexed_fields, ) from ..utils import generate_id @@ -784,8 +786,14 @@ def get_indexes(cls: Type["Object"]) -> List[Dict[str, Any]]: List of index definitions, each containing: - For single-field: {"field": "context.field_name", "unique": bool, "direction": int} - For compound: {"fields": [("context.field_name", direction), ...], "unique": bool, "name": str} + - For full-text: {"fields": [...], "fulltext": True, "name": str} + Annotation-declared definitions carry ``"per_class": True`` (plus + any ``entity_leading`` / ``partial_by_entity`` option) so + ``GraphContext.ensure_indexes`` can scope them to the class's + entity on Postgres. """ indexes: List[Dict[str, Any]] = [] + scope_options = ("entity_leading", "partial_by_entity") # Get single-field indexes from field annotations indexed_fields = get_indexed_fields(cls) @@ -796,11 +804,15 @@ def get_indexes(cls: Type["Object"]) -> List[Dict[str, Any]]: "field": db_field, "unique": index_config.get("unique", False), "direction": index_config.get("direction", 1), + "per_class": True, } if "partial_filter_expression" in index_config: single_entry["partialFilterExpression"] = index_config[ "partial_filter_expression" ] + for option in scope_options: + if option in index_config: + single_entry[option] = index_config[option] indexes.append(single_entry) # Get compound indexes from class decorators @@ -816,11 +828,34 @@ def get_indexes(cls: Type["Object"]) -> List[Dict[str, Any]]: "unique": comp_index.get("unique", False), "sparse": comp_index.get("sparse", False), "name": comp_index.get("name"), + "per_class": True, } if "partialFilterExpression" in comp_index: entry["partialFilterExpression"] = comp_index["partialFilterExpression"] + for option in scope_options: + if option in comp_index: + entry[option] = comp_index[option] indexes.append(entry) + # Full-text indexes: every ``attribute(fulltext=True)`` field in + # declaration order, plus each ``@fulltext_index([...])``. + fulltext_sets = [] + fulltext_fields = get_fulltext_fields(cls) + if fulltext_fields: + fulltext_sets.append( + {"fields": fulltext_fields, "name": f"fts_{'_'.join(fulltext_fields)}"} + ) + fulltext_sets.extend(get_fulltext_indexes(cls)) + for ft in fulltext_sets: + indexes.append( + { + "fields": [(f"context.{f}", 1) for f in ft["fields"]], + "fulltext": True, + "per_class": True, + "name": ft["name"], + } + ) + return indexes @classmethod diff --git a/jvspatial/db/README.md b/jvspatial/db/README.md index 911bbf2..98f42d5 100644 --- a/jvspatial/db/README.md +++ b/jvspatial/db/README.md @@ -63,6 +63,8 @@ db/ - **`find_many` and `bulk_save` are public and benefit from native overrides.** Defaults exist but are slow. (`database.py:176+`) - **`find_one_and_update` / `find_one_and_delete` are NOT atomic by default.** MongoDB and Postgres override with native atomic versions (`FOR UPDATE` on Postgres). - **Neighbour pushdown** (via `getattr`, not on the ABC): `find_connected_nodes` / `count_connected_nodes` on Postgres, MongoDB and SQLite take `edge_entities`, `node_entities`, `edge_query`, `node_query` (record paths), `sort` and `limit`, and raise `NotImplementedError` for anything they cannot translate so `Node.nodes()` falls back without dropping a filter. Postgres adds `find_connected_nodes_bulk(limit_per_source=...)`. +- **Postgres per-class indexes are entity-scoped.** `ensure_indexes` passes the class's entity to `create_index` for annotation-declared indexes: `(entity, )` by default, `WHERE entity = ...` with `partial_by_entity`. Descending keys are `DESC NULLS LAST`; stale pre-0.0.18 definitions are rebuilt. The whole-document GIN is optional (`gin_index="off"` / `JVSPATIAL_PG_GIN_INDEX=off`). +- **`$text` pushes down on Postgres only** (`to_tsvector('simple', …) @@ plainto_tsquery`, `$fields` required, GIN from `@fulltext_index` / `attribute(fulltext=True)`); SQLite / JsonDB / DynamoDB evaluate it in memory with the same semantics, MongoDB uses its text index. `escape_regex()` for literal `$regex` patterns. - **Postgres-only helpers** (via `getattr`, not on the ABC): `traverse`, `find_connected_nodes_bulk`, `save_with_edge_merge` (persist mode only). - **Atomic JSON writes use `temp + fsync + rename + fsync(dir)`.** No partial records survive a crash. (`_atomic.py`) - **Per-file locks serialize concurrent writes to the same record only.** Different files run in parallel. (`_path_locks.py`) diff --git a/jvspatial/db/__init__.py b/jvspatial/db/__init__.py index 86199dc..e77db8e 100644 --- a/jvspatial/db/__init__.py +++ b/jvspatial/db/__init__.py @@ -19,6 +19,7 @@ ) from .jsondb import JsonDB from .manager import DatabaseManager, get_database_manager, set_database_manager +from .query import escape_regex try: # Optional dependency (requires aiosqlite) from .sqlite import SQLiteDB # noqa: F401 @@ -69,6 +70,7 @@ "get_database_manager", "set_database_manager", "JsonDB", + "escape_regex", ] if _SQLITE_AVAILABLE: diff --git a/jvspatial/db/_postgres_translate.py b/jvspatial/db/_postgres_translate.py index e15611c..6d85bfc 100644 --- a/jvspatial/db/_postgres_translate.py +++ b/jvspatial/db/_postgres_translate.py @@ -507,6 +507,43 @@ def _to_jsonb_literal(value: Any) -> str: return json.dumps(value) +# ---- full-text search ------------------------------------------------------- + + +def tsvector_expression(fields: List[str], table: str = "") -> Optional[str]: + """``to_tsvector('simple', …)`` over the concatenated text of ``fields``. + + Shared by the ``$text`` translation and ``PostgresDB.create_index(..., + fulltext=True)`` so a query and its GIN index use the identical + expression (fields in the same order). ``'simple'`` lowercases and splits + words without stemming or stop words. Returns ``None`` for unsafe paths. + """ + if not fields or not all(_safe_field_path(f) for f in fields): + return None + prefix = f"{table}." if table else "" + parts = " || ' ' || ".join( + f"coalesce({prefix}data #>> '{_path_literal(f)}', '')" for f in fields + ) + return f"to_tsvector('simple'::regconfig, {parts})" + + +def _text_clause(spec: Any, pb: ParamBuilder, table: str) -> Optional[str]: + """``{"$search": str, "$fields": [paths]}`` → ``tsvector @@ plainto_tsquery``. + + ``$fields`` is required: without it there is no indexable expression + (the query falls back to in-memory evaluation). + """ + if not isinstance(spec, dict) or set(spec) - {"$search", "$fields"}: + return None + search, fields = spec.get("$search"), spec.get("$fields") + if not isinstance(search, str) or not isinstance(fields, (list, tuple)): + return None + expr = tsvector_expression(list(fields), table) + if expr is None: + return None + return f"{expr} @@ plainto_tsquery('simple'::regconfig, {pb.add(search)})" + + # ---- logical translation ---------------------------------------------------- @@ -559,6 +596,12 @@ def _translate_query_into( return None fragments.append(f"NOT ({inner})") continue + if key == "$text": + piece = _text_clause(value, pb, table) + if piece is None: + return None + fragments.append(piece) + continue if key.startswith("$"): return None piece = _translate_field_clause(key, value, pb, table) @@ -634,4 +677,4 @@ def translate_sort( return ", ".join(parts) -__all__ = ["ParamBuilder", "translate_query", "translate_sort"] +__all__ = ["ParamBuilder", "translate_query", "translate_sort", "tsvector_expression"] diff --git a/jvspatial/db/mongodb.py b/jvspatial/db/mongodb.py index 355b6a8..e83fc19 100644 --- a/jvspatial/db/mongodb.py +++ b/jvspatial/db/mongodb.py @@ -69,6 +69,19 @@ def _is_retryable_mongo_error(exc: BaseException) -> bool: return False +def _native_query(query: Dict[str, Any]) -> Dict[str, Any]: + """Drop the jvspatial-only ``$fields`` from ``$text``. + + MongoDB's ``$text`` searches the collection's text index and rejects + unknown keys; ``$fields`` names the searched fields for the other + backends (Postgres tsvector expression, in-memory evaluation). + """ + text = query.get("$text") if isinstance(query, dict) else None + if isinstance(text, dict) and "$fields" in text: + return {**query, "$text": {k: v for k, v in text.items() if k != "$fields"}} + return query + + class MongoDB(Database): """Simplified MongoDB-based database implementation.""" @@ -520,7 +533,7 @@ async def _find_op() -> List[Dict[str, Any]]: if self._db is None: raise DatabaseError("MongoDB database connection not established") collection_obj = self._db[collection] - cursor = collection_obj.find(query) + cursor = collection_obj.find(_native_query(query)) if sort: cursor = cursor.sort(sort) if limit is not None: @@ -686,7 +699,7 @@ async def count( if not q: # estimated_document_count is the fastest path for full counts. return await collection_obj.estimated_document_count() - return await collection_obj.count_documents(q) + return await collection_obj.count_documents(_native_query(q)) except PyMongoError as e: if _is_connection_error(e): logger.debug( @@ -703,7 +716,7 @@ async def count( collection_obj = self._db[collection] if not q: return await collection_obj.estimated_document_count() - return await collection_obj.count_documents(q) + return await collection_obj.count_documents(_native_query(q)) raise DatabaseError(f"MongoDB count error: {e}") from e async def find_one_and_update( diff --git a/jvspatial/db/postgres.py b/jvspatial/db/postgres.py index 0eeb991..f432548 100644 --- a/jvspatial/db/postgres.py +++ b/jvspatial/db/postgres.py @@ -84,7 +84,7 @@ Union, ) -from ._postgres_translate import translate_query, translate_sort +from ._postgres_translate import translate_query, translate_sort, tsvector_expression from .database import ( BulkSaveResult, Database, @@ -166,6 +166,27 @@ def _pg_string_literal(value: str) -> str: return "'" + value.replace("'", "''") + "'" +# An index definition written before 0.0.18: top-level columns indexed as JSONB +# paths (``data #>> '{entity}'``, which the translator never compares against) +# or a descending key without ``NULLS LAST`` (which ``find(sort=...)`` cannot +# walk). ``create_index`` rebuilds such an index in place. +_STALE_INDEX_DEF_RE = re.compile( + r"'\{(?:entity|id|tenant_id)\}'|\bDESC\b(?! NULLS LAST)" +) + + +def _uses_containment_ops(query: Any) -> bool: + """Whether a Mongo-style query uses operators the whole-document GIN serves.""" + if isinstance(query, dict): + return any( + key in ("$all", "$elemMatch") or _uses_containment_ops(value) + for key, value in query.items() + ) + if isinstance(query, list): + return any(_uses_containment_ops(item) for item in query) + return False + + def _pg_field_extract(field_path: str) -> Optional[str]: """Translate a dotted field path to ``(data #>> '{a,b,c}')``. @@ -286,6 +307,7 @@ def __init__( pooler_mode: str = "session", command_timeout: float = 60.0, schema_name: str = "public", + gin_index: Optional[str] = None, ) -> None: """Initialize the Postgres adapter. @@ -305,10 +327,17 @@ def __init__( command_timeout: Per-statement timeout in seconds. schema_name: Postgres schema to host the collection tables in. Defaults to ``public``. Must be an existing schema. + gin_index: ``"full"`` (default) creates the whole-document + ``GIN (data jsonb_path_ops)`` index on each new collection; + ``"off"`` skips it. Also read from ``JVSPATIAL_PG_GIN_INDEX``. + The GIN only serves containment-shaped predicates (``$all``); + typed ``find()`` equality / range / sort use functional + B-trees, so large deployments should turn it off. Raises: ImportError: ``asyncpg`` is not installed. - ValueError: ``pooler_mode`` is not ``"session"`` or ``"transaction"``. + ValueError: ``pooler_mode`` is not ``"session"`` or ``"transaction"``, + or ``gin_index`` is not ``"full"`` or ``"off"``. """ if asyncpg is None: # pragma: no cover raise ImportError( @@ -340,6 +369,14 @@ def __init__( self.pooler_mode = pooler_mode self.command_timeout = command_timeout self.schema_name = schema_name + self.gin_index = ( + (gin_index or env("JVSPATIAL_PG_GIN_INDEX", default="full")).strip().lower() + ) + if self.gin_index not in ("full", "off"): + raise ValueError( + f"gin_index must be 'full' or 'off', got {self.gin_index!r}" + ) + self._gin_off_warned = False self._pool: Optional["Pool"] = None self._pool_lock = asyncio.Lock() @@ -576,6 +613,12 @@ async def _bootstrap_collection(self, collection: str) -> None: return col = _safe_collection(collection) schema = _safe_collection(self.schema_name) + gin_sql = ( + f"CREATE INDEX IF NOT EXISTS {col}_data_gin " + f"ON {schema}.{col} USING GIN (data jsonb_path_ops);" + if self.gin_index == "full" + else "" + ) pool = await self._ensure_pool() async with pool.acquire() as conn: # Single round trip — CREATE IF NOT EXISTS is cheap when the @@ -591,8 +634,7 @@ async def _bootstrap_collection(self, collection: str) -> None: created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() ); - CREATE INDEX IF NOT EXISTS {col}_data_gin - ON {schema}.{col} USING GIN (data jsonb_path_ops); + {gin_sql} CREATE INDEX IF NOT EXISTS {col}_entity_idx ON {schema}.{col} (entity); CREATE INDEX IF NOT EXISTS {col}_tenant_idx @@ -898,6 +940,19 @@ async def find( where_sql, params = translated sort_sql = translate_sort(sort) if vec_field is None else None + if ( + self.gin_index == "off" + and not self._gin_off_warned + and _uses_containment_ops(query) + ): + self._gin_off_warned = True + logger.warning( + "PostgresDB: $all / $elemMatch query on %s.%s with gin_index='off' " + "— no whole-document GIN serves it; add a targeted index if it is " + "hot (logged once per adapter)", + schema, + col, + ) clauses: List[str] = [] if where_sql: @@ -1373,13 +1428,38 @@ async def create_index( unique: bool = False, **kwargs: Any, ) -> None: - """Create a functional B-tree index on a JSONB field path. - - Field paths use dot notation (``"context.user.id"``) and translate - to ``json_extract`` style ``(data #>> '{context,user,id}')``. Compound - index inputs (``[(field, dir)]``) are honored; ``unique=True`` adds - ``CREATE UNIQUE INDEX``. Postgres-specific kwargs (``where``, - ``method``) are accepted via ``**kwargs``. + """Create a functional index on JSONB field paths or top-level columns. + + Field paths use dot notation (``"context.user.id"``) and index + ``(data #>> '{context,user,id}')``; ``entity`` / ``id`` / + ``tenant_id`` index the real columns the query translator compares + against. A descending key is ``DESC NULLS LAST`` — the order + ``find(sort=...)`` emits — so sort + limit can walk the index. + ``unique=True`` adds ``CREATE UNIQUE INDEX``. + + Keyword options: + + * ``entity`` — the entity a per-class index serves (passed by + ``GraphContext.ensure_indexes``). The real ``entity`` column then + leads the key (``entity_leading=True``, default): every typed + ``find()`` filters on it, in a table shared by every class. Such an + index is named ``_entity__idx`` so sibling classes + declaring the same fields share it. ``partial_by_entity=True`` + instead restricts it to ``WHERE entity = ''`` + (``___idx``; smaller, one per class). + ``drop_legacy=True`` also drops the unscoped pre-0.0.18 + ``__idx`` / ``_uniq`` index it replaces. + * ``fulltext=True`` — a GIN index over ``to_tsvector('simple', ...)`` + of the fields in order (``___fts``, scoped to + ``entity`` when given), matching + ``{"$text": {"$search": ..., "$fields": []}}``. + * ``where`` / ``index_partial_filter_expression`` — partial index + predicate (raw SQL, or the Mongo-style dialect). + * ``method`` — ``btree`` (default), ``hash``, ``gin``, ``gist``, ``brin``. + + An existing index of the same name defined by the pre-0.0.18 rules + (JSONB paths for top-level columns, ``DESC`` without ``NULLS LAST``) + is rebuilt under a temporary name and swapped in. """ await self._bootstrap_collection(collection) col = _safe_collection(collection) @@ -1390,34 +1470,41 @@ async def create_index( fields: List[Tuple[str, int]] = [(field_or_fields, 1)] else: fields = list(field_or_fields) - - # Build the column expression list. - col_exprs: List[str] = [] - name_parts: List[str] = [] - for field_path, direction in fields: + for field_path, _direction in fields: for seg in field_path.split("."): if not _SAFE_IDENT_RE.match(seg): raise ValueError( f"PostgresDB.create_index rejects field path {field_path!r}: " f"segment {seg!r} must be a safe identifier" ) - path_literal = "{" + ",".join(field_path.split(".")) + "}" - direction_sql = "ASC" if direction == 1 else "DESC" - col_exprs.append(f"(data #>> '{path_literal}') {direction_sql}") - name_parts.append(field_path.replace(".", "_")) + field_names = "_".join(f.replace(".", "_") for f, _ in fields) - unique_sql = "UNIQUE " if unique else "" - index_name = f"{col}_{'_'.join(name_parts)}_idx" - if unique: - index_name = f"{col}_{'_'.join(name_parts)}_uniq" + entity = kwargs.get("entity") + if entity is not None and not _SAFE_IDENT_RE.match(str(entity)): + logger.warning( + "PostgresDB.create_index: entity %r is not a safe identifier; " + "index left unscoped", + entity, + ) + entity = None + fulltext = bool(kwargs.get("fulltext")) + partial_by_entity = bool(entity) and ( + fulltext or bool(kwargs.get("partial_by_entity")) + ) + entity_leading = ( + bool(entity) + and not partial_by_entity + and bool(kwargs.get("entity_leading", True)) + and all(f != "entity" for f, _ in fields) + ) - where_clause = "" + where_parts: List[str] = [] partial = kwargs.get("where") if partial: # We don't try to parse the partial expression — caller is # responsible for getting it right. We do require it to be a # simple Postgres predicate string. - where_clause = f" WHERE {partial}" + where_parts.append(str(partial)) else: # Translate Mongo-style ``index_partial_filter_expression`` # (the cross-backend kwarg used by ``attribute(index_unique=..., @@ -1451,19 +1538,64 @@ async def create_index( "(equality / $gt / $exists on safe field paths " "with scalar values)." ) - where_clause = f" WHERE {translated}" - - method = kwargs.get("method", "btree") + where_parts.append(translated) + if partial_by_entity: + where_parts.append(f"entity = {_pg_string_literal(str(entity))}") + where_clause = f" WHERE {' AND '.join(where_parts)}" if where_parts else "" + + if fulltext: + expr = tsvector_expression([f for f, _ in fields]) + if expr is None: # pragma: no cover - paths validated above + raise ValueError(f"unsafe full-text fields: {fields!r}") + keys = [expr] + method = "gin" + suffix = "fts" + else: + keys = ["entity"] if entity_leading else [] + for field_path, direction in fields: + extract = _pg_field_extract(field_path) + keys.append( + f"{extract} {'ASC' if direction == 1 else 'DESC NULLS LAST'}" + ) + method = kwargs.get("method", "btree") + suffix = "uniq" if unique else "idx" if method not in ("btree", "hash", "gin", "gist", "brin"): raise ValueError(f"Unsupported index method: {method!r}") + if partial_by_entity: + scope = f"{str(entity).lower()}_" + elif entity_leading: + scope = "entity_" + else: + scope = "" + index_name = f"{col}_{scope}{field_names}_{suffix}" + unique_sql = "UNIQUE " if unique and not fulltext else "" + body = f"ON {schema}.{col} USING {method} ({', '.join(keys)}){where_clause}" + pool = await self._ensure_pool() async with pool.acquire() as conn: - await conn.execute( - f"CREATE {unique_sql}INDEX IF NOT EXISTS {index_name} " - f"ON {schema}.{col} USING {method} ({', '.join(col_exprs)})" - f"{where_clause}" + existing = await conn.fetchval( + "SELECT indexdef FROM pg_indexes WHERE schemaname = $1 AND indexname = $2", + schema, + index_name, ) + if existing is not None and _STALE_INDEX_DEF_RE.search(existing): + # Build the corrected index first so a unique constraint is + # never absent, then swap it in under the canonical name. + temp = f"{index_name[:52]}_rebuild" + await conn.execute(f"DROP INDEX IF EXISTS {schema}.{temp}") + await conn.execute(f"CREATE {unique_sql}INDEX {temp} {body}") + await conn.execute(f"DROP INDEX {schema}.{index_name}") + await conn.execute( + f"ALTER INDEX {schema}.{temp} RENAME TO {index_name}" + ) + else: + await conn.execute( + f"CREATE {unique_sql}INDEX IF NOT EXISTS {index_name} {body}" + ) + if kwargs.get("drop_legacy") and scope: + legacy = f"{col}_{field_names}_{'uniq' if unique else 'idx'}" + await conn.execute(f"DROP INDEX IF EXISTS {schema}.{legacy}") # ---- atomic compound ops (C4) ------------------------------------------ diff --git a/jvspatial/db/query.py b/jvspatial/db/query.py index bc72582..76b1f82 100644 --- a/jvspatial/db/query.py +++ b/jvspatial/db/query.py @@ -8,7 +8,7 @@ import re import time from collections import OrderedDict -from typing import Any, Callable, Dict, List, Optional, Union +from typing import Any, Callable, Dict, Iterator, List, Optional, Union from jvspatial.exceptions import QueryError @@ -44,6 +44,34 @@ # Unified evaluation and builder in a single module +# Word splitter for in-memory ``$text`` (approximates Postgres' 'simple' parser). +_TEXT_WORD_RE = re.compile(r"\w+") + + +def _iter_strings(value: Any) -> Iterator[str]: + """Every string leaf of a nested dict / list document.""" + if isinstance(value, str): + yield value + elif isinstance(value, dict): + for item in value.values(): + yield from _iter_strings(item) + elif isinstance(value, (list, tuple)): + for item in value: + yield from _iter_strings(item) + + +def escape_regex(value: str) -> str: + """Escape ``value`` so it matches literally inside a ``$regex`` pattern. + + ``$regex`` is never index-backed (Postgres ``~``, SQLite and JsonDB + in-memory evaluation all scan), so prefer ``$text`` or an equality on an + indexed field. When a pattern must be built from user input, pass the + input through this helper: an unescaped ``.`` or ``(`` silently changes + the match, and a crafted pattern can make every scan pathologically slow. + """ + return re.escape(value) + + class QueryEngine: """Unified MongoDB-style query engine with built-in optimization for all backends.""" @@ -375,6 +403,9 @@ def match(document: Dict[str, Any], query: Optional[Dict[str, Any]]) -> bool: elif key == "$not": if QueryEngine.match(document, condition): return False + elif key == "$text": + if not QueryEngine._match_text(document, condition): + return False elif key in QueryEngine._IGNORED_TOP_LEVEL_MARKERS: # Optimizer hints — irrelevant to in-memory matching. continue @@ -386,7 +417,7 @@ def match(document: Dict[str, Any], query: Optional[Dict[str, Any]]) -> bool: query=str(query), reason=( f"unsupported top-level query operator: {key!r}. " - "Supported: $and, $or, $nor, $not. Field-level " + "Supported: $and, $or, $nor, $not, $text. Field-level " "operators (e.g. $regex, $mod, $type, $size) live " "inside a field condition dict." ), @@ -397,6 +428,35 @@ def match(document: Dict[str, Any], query: Optional[Dict[str, Any]]) -> bool: return False return True + @staticmethod + def _match_text(document: Dict[str, Any], spec: Any) -> bool: + """Evaluate ``{"$text": {"$search": ..., "$fields": [...]}}`` in memory. + + Mirrors the Postgres pushdown (``to_tsvector('simple', …) @@ + plainto_tsquery('simple', …)``): every word of ``$search`` must occur, + case-insensitively, as a word of the concatenated ``$fields`` values. + Without ``$fields`` every string value in the document is searched. + No stemming; a search with no words matches nothing. + """ + if not isinstance(spec, dict) or not isinstance(spec.get("$search"), str): + raise QueryError( + query=str(spec), + reason='$text needs {"$search": "", "$fields": []}', + ) + wanted = set(_TEXT_WORD_RE.findall(spec["$search"].lower())) + if not wanted: + return False + fields = spec.get("$fields") + if fields: + values = [QueryEngine.get_field_value(document, f) for f in fields] + else: + values = list(_iter_strings(document)) + words: set = set() + for value in values: + if isinstance(value, str): + words.update(_TEXT_WORD_RE.findall(value.lower())) + return wanted <= words + @staticmethod def _match_value(value: Any, condition: Any) -> bool: if not isinstance(condition, dict): diff --git a/jvspatial/env_adapter.py b/jvspatial/env_adapter.py index 32940ce..deaaf10 100644 --- a/jvspatial/env_adapter.py +++ b/jvspatial/env_adapter.py @@ -274,6 +274,7 @@ def server_config_overrides_from_env() -> Dict[str, Any]: "JVSPATIAL_DYNAMODB_WAIT_FOR_INDEX", "JVSPATIAL_AUTO_CREATE_INDEXES", "JVSPATIAL_NODE_EDGE_IDS", + "JVSPATIAL_PG_GIN_INDEX", # Auth "JVSPATIAL_AUTH_ENABLED", "JVSPATIAL_AUTH_STRICT_HASHING", diff --git a/tests/benchmarks/hub_bench_report.py b/tests/benchmarks/hub_bench_report.py index c0b93dd..7d37c88 100644 --- a/tests/benchmarks/hub_bench_report.py +++ b/tests/benchmarks/hub_bench_report.py @@ -52,8 +52,33 @@ def _label(degree: int) -> str: return f"{degree // 1000}k" if degree % 1000 == 0 else str(degree) -def render(records: List[Dict[str, Any]]) -> str: - records = sorted(records, key=lambda r: r["degree"]) +def render(all_records: List[Dict[str, Any]]) -> str: + typed = [r for r in all_records if r.get("kind") == "typed_find"] + records = sorted( + (r for r in all_records if "degree" in r), key=lambda r: r["degree"] + ) + if not records: + return _render_typed(typed) + return _render_hub(records) + ("\n\n" + _render_typed(typed) if typed else "") + + +def _render_typed(typed: List[Dict[str, Any]]) -> str: + out = [ + "| typed find (sorted, limit 20) | rows | p50 / p95 ms | " + "index walked | sort node | node_data_gin |", + "|---|---|---|---|---|---|", + ] + for r in typed: + m = r["typed_find_sorted_limit20"] + out.append( + f"| `{r['git_sha']}` | {r['rows']:,} | {m['p50_ms']:.2f} / {m['p95_ms']:.2f} | " + f"{', '.join(r['plan_indexes']) or '—'} | {'yes' if r['plan_sorts'] else 'no'} | " + f"{'present' if r['node_data_gin'] else 'off'} |" + ) + return "\n".join(out) + + +def _render_hub(records: List[Dict[str, Any]]) -> str: heads = [_label(r["degree"]) for r in records] out: List[str] = [] meta = records[0] diff --git a/tests/benchmarks/test_hub_node_bench.py b/tests/benchmarks/test_hub_node_bench.py index 1386d71..80bef7e 100644 --- a/tests/benchmarks/test_hub_node_bench.py +++ b/tests/benchmarks/test_hub_node_bench.py @@ -43,6 +43,7 @@ import jvspatial from jvspatial.core import context as context_module +from jvspatial.core.annotations import compound_index from jvspatial.core.context import GraphContext, set_default_context from jvspatial.core.entities import Edge, Node from jvspatial.core.utils import generate_id @@ -81,6 +82,14 @@ class BenchContains(Edge): """Typed containment edge (hub -> leaf, leaf -> sink).""" +@compound_index([("track_id", 1), ("created_at", -1)], name="bench_track_recent") +class BenchEntry(Node): + """Record-style node for the typed-find gate (sorted, limited per track).""" + + track_id: str = "" + created_at: str = "" + + # ---- helpers ---------------------------------------------------------------- @@ -439,3 +448,115 @@ async def _one(leaf: BenchLeaf) -> float: if _RESULTS_PATH: with open(_RESULTS_PATH, "a", encoding="utf-8") as fh: fh.write(json.dumps(record) + "\n") + + +@pytest.mark.bench_slow +async def test_typed_find_is_index_bound( + request: pytest.FixtureRequest, monkeypatch: pytest.MonkeyPatch +) -> None: + """Sorted, limited typed ``find`` on a shared ``node`` table of N rows. + + ``N`` defaults to 1M (``JVSPATIAL_BENCH_TYPED_ROWS``) spread over ten + entities; ``BenchEntry`` holds a tenth, 1000 tracks of ~100 rows each. + Measures ``find({"entity": "BenchEntry", "context.track_id": t}, + sort=created_at desc, limit=20)`` with the whole-document GIN off and + records whether the plan walks an index without a Sort node. + """ + if "bench_slow" not in (request.config.getoption("markexpr") or ""): + pytest.skip("typed-find gate is opt-in: -m 'bench or bench_slow'") + if not await _dsn_reachable(): + pytest.skip(f"Postgres unreachable at {_DSN}; set JVSPATIAL_POSTGRES_TEST_DSN") + monkeypatch.setenv("JVSPATIAL_PG_GIN_INDEX", "off") + rows = int(os.getenv("JVSPATIAL_BENCH_TYPED_ROWS", "1000000")) + entities = ["BenchEntry"] + [f"BenchOther{i}" for i in range(9)] + + async with _bench_graph() as (ctx, db, schema): + admin = await asyncpg.connect(dsn=_DSN) + try: + await ctx.ensure_indexes(BenchEntry) + # Seed through the adapter itself: before 0.0.18 the observable + # wrapper degraded bulk_save_detailed to per-record saves. + seed_db = getattr(db, "inner", db) + t0 = time.perf_counter() + batch: List[Dict[str, Any]] = [] + for i in range(rows): + entity = entities[i % len(entities)] + batch.append( + { + "id": generate_id("n", entity), + "entity": entity, + "context": { + "track_id": f"t{(i // len(entities)) % 1000}", + "created_at": f"2026-09-{i:08d}", + }, + } + ) + if len(batch) == _SEED_CHUNK: + await seed_db.bulk_save_detailed("node", batch) + batch = [] + if batch: + await seed_db.bulk_save_detailed("node", batch) + seed_s = time.perf_counter() - t0 + await admin.execute(f"ANALYZE {schema}.node") + + per_track = min(20, rows // len(entities) // 1000) + samples: List[float] = [] + for k in range(50): + query = { + "entity": "BenchEntry", + "context.track_id": f"t{(k * 37) % 1000}", + } + ms, _, out = await _timed( + functools.partial( + db.find, + "node", + query, + sort=[("context.created_at", -1)], + limit=20, + ) + ) + samples.append(ms) + assert len(out) == per_track + + from jvspatial.db._postgres_translate import translate_query, translate_sort + + where, params = translate_query( + {"entity": "BenchEntry", "context.track_id": "t7"} + ) + order = translate_sort([("context.created_at", -1)]) + raw = await admin.fetchval( + f"EXPLAIN (FORMAT JSON) SELECT data FROM {schema}.node " + f"WHERE {where} ORDER BY {order} LIMIT 20", + *params, + ) + plan = json.loads(raw) if isinstance(raw, str) else raw + + def _walk(node: Dict[str, Any]) -> List[Dict[str, Any]]: + return [node] + [n for c in node.get("Plans", []) for n in _walk(c)] + + steps = _walk(plan[0]["Plan"]) + index_names = sorted({s["Index Name"] for s in steps if "Index Name" in s}) + has_sort = any(s["Node Type"] == "Sort" for s in steps) + gin = await admin.fetchval( + "SELECT count(*) FROM pg_indexes WHERE schemaname = $1 " + "AND indexname = 'node_data_gin'", + schema, + ) + finally: + await admin.close() + + record = { + "kind": "typed_find", + "rows": rows, + "jvspatial": jvspatial.__version__, + "git_sha": _git_sha(), + "seed_seconds": round(seed_s, 2), + "node_data_gin": bool(gin), + "typed_find_sorted_limit20": _summary(samples), + "plan_indexes": index_names, + "plan_sorts": has_sort, + } + print("\n[typed-find]\n" + json.dumps(record, indent=2)) + if _RESULTS_PATH: + with open(_RESULTS_PATH, "a", encoding="utf-8") as fh: + fh.write(json.dumps(record) + "\n") diff --git a/tests/db/test_postgres_indexes_text.py b/tests/db/test_postgres_indexes_text.py new file mode 100644 index 0000000..ed4498e --- /dev/null +++ b/tests/db/test_postgres_indexes_text.py @@ -0,0 +1,397 @@ +"""Index hygiene and full-text search. + +* Per-class (annotation-declared) indexes on Postgres are entity-scoped: + ``entity``-leading by default, ``WHERE entity = ...`` on request, and they + replace the unscoped pre-0.0.18 index of the same fields. +* ``create_index`` indexes top-level columns as columns, emits + ``DESC NULLS LAST``, and rebuilds indexes defined by the old rules. +* The whole-document GIN is optional (``gin_index`` / ``JVSPATIAL_PG_GIN_INDEX``). +* ``$text`` pushes down to ``to_tsvector('simple', ...) @@ plainto_tsquery`` + on Postgres and matches the in-memory evaluation everywhere else. + +Postgres tests run when ``JVSPATIAL_POSTGRES_TEST_DSN`` is reachable. +""" + +from __future__ import annotations + +import asyncio +import contextlib +import json +import logging +import os +import uuid +from typing import Any, AsyncIterator, Dict, List, Tuple + +import pytest + +from jvspatial.core import context as context_module +from jvspatial.core.annotations import attribute, compound_index, fulltext_index +from jvspatial.core.context import GraphContext, set_default_context +from jvspatial.core.entities import Edge, Node +from jvspatial.core.utils import generate_id +from jvspatial.db import escape_regex +from jvspatial.db._postgres_translate import translate_query, translate_sort +from jvspatial.db.jsondb import JsonDB +from jvspatial.db.mongodb import _native_query +from jvspatial.db.query import QueryEngine +from jvspatial.db.sqlite import SQLiteDB +from jvspatial.exceptions import QueryError + +pytestmark = pytest.mark.asyncio(loop_scope="module") + +_PG_DSN = os.getenv( + "JVSPATIAL_POSTGRES_TEST_DSN", + "postgresql://jvspatial:jvspatial@localhost:5432/jvspatial", +) + + +@compound_index([("track_id", 1), ("created_at", -1)], name="track_recent") +class IdxEntry(Node): + track_id: str = attribute(indexed=True, default="") + created_at: str = "" + note: str = attribute(indexed=True, index_partial_by_entity=True, default="") + + +@fulltext_index(["title", "body"]) +class IdxArticle(Node): + title: str = "" + body: str = "" + summary: str = attribute(fulltext=True, default="") + + +class IdxLink(Edge): + pass + + +@pytest.fixture(autouse=True) +def _auto_indexes(monkeypatch): + monkeypatch.setenv("JVSPATIAL_AUTO_CREATE_INDEXES", "true") + + +@contextlib.asynccontextmanager +async def _pg(**db_kwargs: Any) -> AsyncIterator[Tuple[GraphContext, Any, Any, str]]: + try: + import asyncpg + + from jvspatial.db.postgres import PostgresDB + except ImportError: + pytest.skip("asyncpg not installed") + try: + admin = await asyncio.wait_for(asyncpg.connect(dsn=_PG_DSN), timeout=2.0) + except Exception: + pytest.skip(f"Postgres unreachable at {_PG_DSN}") + schema = f"t_idx_{uuid.uuid4().hex[:10]}" + await admin.execute(f"CREATE SCHEMA {schema}") + db = PostgresDB( + dsn=_PG_DSN, schema_name=schema, min_size=1, max_size=4, **db_kwargs + ) + ctx = GraphContext(database=db) + set_default_context(ctx) + context_module._ensured_indexes.clear() + try: + yield ctx, db, admin, schema + finally: + set_default_context(None) + context_module._ensured_indexes.clear() + await db.close() + await admin.execute(f"DROP SCHEMA {schema} CASCADE") + await admin.close() + + +async def _indexes(admin: Any, schema: str, table: str) -> Dict[str, str]: + rows = await admin.fetch( + "SELECT indexname, indexdef FROM pg_indexes " + "WHERE schemaname = $1 AND tablename = $2", + schema, + table, + ) + return {r["indexname"]: r["indexdef"] for r in rows} + + +def _plan_nodes(plan: Any) -> List[Dict[str, Any]]: + found = [plan] + for child in plan.get("Plans", []): + found += _plan_nodes(child) + return found + + +async def _explain(admin: Any, sql: str, params: List[Any]) -> List[Dict[str, Any]]: + raw = await admin.fetchval(f"EXPLAIN (FORMAT JSON) {sql}", *params) + plan = json.loads(raw) if isinstance(raw, str) else raw + return _plan_nodes(plan[0]["Plan"]) + + +# ---- entity-scoped per-class indexes ------------------------------------------ + + +async def test_per_class_indexes_are_entity_scoped(): + async with _pg() as (ctx, db, admin, schema): + await db.create_index("node", "context.track_id") # pre-0.0.18 unscoped + assert "node_context_track_id_idx" in await _indexes(admin, schema, "node") + + await ctx.ensure_indexes(IdxEntry) + idx = await _indexes(admin, schema, "node") + + single = idx["node_entity_context_track_id_idx"] + assert "(entity, ((data #>> '{context,track_id}'::text[])))" in single + compound = idx["node_entity_context_track_id_context_created_at_idx"] + assert compound.startswith("CREATE INDEX") + assert "DESC NULLS LAST" in compound + partial = idx["node_idxentry_context_note_idx"] + assert "WHERE (entity = 'IdxEntry'::text)" in partial + assert "node_context_track_id_idx" not in idx # legacy replaced + + +async def test_stale_edge_index_is_rebuilt_on_entity_column(): + async with _pg() as (ctx, db, admin, schema): + await ctx.save(IdxEntry(track_id="t")) # bootstraps the tables + await db._bootstrap_collection("edge") + await admin.execute( + f"CREATE UNIQUE INDEX edge_source_target_entity_uniq ON {schema}.edge " + "((data #>> '{source}'), (data #>> '{target}'), (data #>> '{entity}'))" + ) + await ctx.ensure_indexes(IdxLink) + fixed = (await _indexes(admin, schema, "edge"))[ + "edge_source_target_entity_uniq" + ] + assert fixed.startswith("CREATE UNIQUE INDEX") + assert "'{entity}'" not in fixed and ", entity)" in fixed + + +async def test_typed_sorted_find_walks_the_entity_leading_index(): + async with _pg(gin_index="off") as (ctx, db, admin, schema): + await ctx.ensure_indexes(IdxEntry) + records = [ + { + "id": generate_id("n", entity), + "entity": entity, + "context": {"track_id": f"t{i % 40}", "created_at": f"2026-01-{i:05d}"}, + } + for entity in ("IdxEntry", "Other1", "Other2", "Other3", "Other4") + for i in range(1500) + ] + await db.bulk_save_detailed("node", records) + await admin.execute(f"ANALYZE {schema}.node") + + where, params = translate_query( + {"entity": "IdxEntry", "context.track_id": "t7"} + ) + order = translate_sort([("context.created_at", -1)]) + await admin.execute(f"SET search_path TO {schema}") + plan = await _explain( + admin, + f"SELECT data FROM node WHERE {where} ORDER BY {order} LIMIT 20", + params, + ) + names = {p.get("Index Name") for p in plan} + assert "node_entity_context_track_id_context_created_at_idx" in names, plan + assert not any(p["Node Type"] == "Sort" for p in plan), plan + + found = await db.find( + "node", + {"entity": "IdxEntry", "context.track_id": "t7"}, + sort=[("context.created_at", -1)], + limit=20, + ) + stamps = [r["context"]["created_at"] for r in found] + assert len(found) == 20 and stamps == sorted(stamps, reverse=True) + + +# ---- optional whole-document GIN ---------------------------------------------- + + +async def test_gin_index_can_be_turned_off(monkeypatch, caplog): + async with _pg() as (ctx, db, admin, schema): + await ctx.save(IdxEntry(track_id="t")) + assert "node_data_gin" in await _indexes(admin, schema, "node") + + async with _pg(gin_index="off") as (ctx, db, admin, schema): + await ctx.save(IdxEntry(track_id="t")) + assert "node_data_gin" not in await _indexes(admin, schema, "node") + with caplog.at_level(logging.WARNING, logger="jvspatial.db.postgres"): + await db.find("node", {"context.tags": {"$all": ["a"]}}) + await db.find("node", {"context.tags": {"$all": ["b"]}}) + warnings = [r for r in caplog.records if "gin_index='off'" in r.getMessage()] + assert len(warnings) == 1 + + monkeypatch.setenv("JVSPATIAL_PG_GIN_INDEX", "off") + from jvspatial.db.postgres import PostgresDB + + assert PostgresDB(dsn=_PG_DSN).gin_index == "off" + with pytest.raises(ValueError): + PostgresDB(dsn=_PG_DSN, gin_index="partial") + + +# ---- full-text search ----------------------------------------------------------- + + +_ARTICLES = [ + ("Graph databases at scale", "Hub nodes and fan-out", "graph"), + ("The vector store", "Embeddings next to the graph", "vectors"), + ("Scaling Postgres", "DATABASE design for large tables", "postgres"), + ("Cooking notes", "Nothing about data here", "food"), +] +_TEXT_QUERIES = [ + {"$search": "graph", "$fields": ["context.title", "context.body"]}, + {"$search": "Graph DATABASE", "$fields": ["context.title", "context.body"]}, + {"$search": "database", "$fields": ["context.title"]}, + {"$search": "vector graph", "$fields": ["context.title", "context.body"]}, + {"$search": "missing", "$fields": ["context.title", "context.body"]}, + {"$search": "fan", "$fields": ["context.title", "context.body"]}, +] + + +def _article_records() -> List[Dict[str, Any]]: + return [ + { + "id": f"n.IdxArticle.a{i}", + "entity": "IdxArticle", + "context": {"title": title, "body": body, "summary": summary}, + } + for i, (title, body, summary) in enumerate(_ARTICLES) + ] + + +async def test_fulltext_indexes_and_text_pushdown(): + async with _pg(gin_index="off") as (ctx, db, admin, schema): + await ctx.ensure_indexes(IdxArticle) + idx = await _indexes(admin, schema, "node") + both = idx["node_idxarticle_context_title_context_body_fts"] + assert "USING gin (to_tsvector('simple'::regconfig" in both + assert "WHERE (entity = 'IdxArticle'::text)" in both + assert "node_idxarticle_context_summary_fts" in idx + + await db.bulk_save_detailed("node", _article_records()) + query = {"entity": "IdxArticle", "$text": _TEXT_QUERIES[0]} # "graph" + got = {r["id"] for r in await db.find("node", query)} + assert got == {"n.IdxArticle.a0", "n.IdxArticle.a1"} + # 'simple' does not stem: "database" does not match "databases". + miss = {"entity": "IdxArticle", "$text": _TEXT_QUERIES[1]} + assert await db.find("node", miss) == [] + + # Enough non-matching rows that the planner prefers the GIN over + # scanning every IdxArticle row through the plain entity index. + await db.bulk_save_detailed( + "node", + [ + { + "id": f"n.IdxArticle.f{i}", + "entity": "IdxArticle", + "context": {"title": f"filler {i}", "body": "lorem ipsum"}, + } + for i in range(3000) + ], + ) + await admin.execute(f"ANALYZE {schema}.node") + where, params = translate_query(query) + await admin.execute(f"SET search_path TO {schema}") + await admin.execute("SET enable_seqscan = off") + plan = await _explain(admin, f"SELECT data FROM node WHERE {where}", params) + assert any( + p.get("Index Name") == "node_idxarticle_context_title_context_body_fts" + for p in plan + ), plan + + +@pytest.mark.parametrize("backend", ["json", "sqlite", "postgres"]) +async def test_text_matches_in_memory_evaluation(backend, tmp_path): + records = _article_records() + expected = [ + {r["id"] for r in records if QueryEngine.match(r, {"$text": q})} + for q in _TEXT_QUERIES + ] + assert expected[0] == {"n.IdxArticle.a0", "n.IdxArticle.a1"} + + async def check(db: Any) -> None: + for record in records: + await db.save("node", dict(record)) + for q, want in zip(_TEXT_QUERIES, expected): + got = {r["id"] for r in await db.find("node", {"$text": q})} + assert got == want, (q, got, want) + assert await db.count("node", {"$text": q}) == len(want) + + if backend == "json": + await check(JsonDB(base_path=str(tmp_path / "j"))) + elif backend == "sqlite": + db = SQLiteDB(db_path=str(tmp_path / "s.db")) + await check(db) + await db.close() + else: + async with _pg() as (_ctx, db, _admin, _schema): + await check(db) + + +def test_text_translation_and_fallbacks(): + sql, params = translate_query( + {"$text": {"$search": "a b", "$fields": ["context.title", "context.body"]}}, + table="n", + ) + assert sql == ( + "to_tsvector('simple'::regconfig, coalesce(n.data #>> '{context,title}', '')" + " || ' ' || coalesce(n.data #>> '{context,body}', '')) " + "@@ plainto_tsquery('simple'::regconfig, $1)" + ) + assert params == ["a b"] + assert translate_query({"$text": {"$search": "a"}}) is None # no $fields + assert ( + translate_query({"$text": {"$search": "a", "$fields": ["x"], "$x": 1}}) is None + ) + assert _native_query({"$text": {"$search": "a", "$fields": ["x"]}}) == { + "$text": {"$search": "a"} + } + + +def test_in_memory_text_semantics(): + doc = {"context": {"title": "Graph Databases", "tags": ["alpha", "Beta"]}} + fields = ["context.title"] + assert QueryEngine.match(doc, {"$text": {"$search": "graph", "$fields": fields}}) + assert QueryEngine.match( + doc, {"$text": {"$search": "DATABASES graph", "$fields": fields}} + ) + assert not QueryEngine.match( + doc, {"$text": {"$search": "database", "$fields": fields}} + ) + assert not QueryEngine.match(doc, {"$text": {"$search": " ", "$fields": fields}}) + assert QueryEngine.match(doc, {"$text": {"$search": "beta graph"}}) # every string + with pytest.raises(QueryError): + QueryEngine.match(doc, {"$text": "graph"}) + + +# ---- $regex escaping, find_edges_between(limit) -------------------------------- + + +_METACHAR_VALUES = ["a.b(c)*", "1+1=2?", "[x]{y}^$|\\z", "plain words"] + + +@pytest.mark.parametrize("backend", ["memory", "postgres"]) +async def test_escape_regex_matches_literally(backend): + records = [ + {"id": f"n.Idx.{i}", "entity": "Idx", "context": {"v": v}} + for i, v in enumerate(_METACHAR_VALUES) + ] + if backend == "memory": + for rec in records: + q = {"context.v": {"$regex": f"^{escape_regex(rec['context']['v'])}$"}} + assert [r["id"] for r in records if QueryEngine.match(r, q)] == [rec["id"]] + return + async with _pg() as (_ctx, db, _admin, _schema): + for rec in records: + await db.save("node", rec) + for rec in records: + q = {"context.v": {"$regex": f"^{escape_regex(rec['context']['v'])}$"}} + assert [r["id"] for r in await db.find("node", q)] == [rec["id"]] + + +async def test_find_edges_between_limit(tmp_path): + ctx = GraphContext(database=JsonDB(base_path=str(tmp_path / "j"))) + set_default_context(ctx) + try: + hub = await IdxEntry.create(track_id="hub") + for i in range(5): + await hub.connect(await IdxEntry.create(track_id=f"l{i}"), edge=IdxLink) + assert len(await ctx.find_edges_between(hub.id, edge_class=IdxLink)) == 5 + assert ( + len(await ctx.find_edges_between(hub.id, edge_class=IdxLink, limit=2)) == 2 + ) + finally: + set_default_context(None) From 51186f1dedab7641e9e3b964be0e75e65b067d11 Mon Sep 17 00:00:00 2001 From: Eldon Marks Date: Fri, 11 Sep 2026 15:21:11 -0400 Subject: [PATCH 7/9] fix(postgres): honour tenant scope on every path; command timeout env Phase 5 of the hub-node scale remediation (verification of the tenancy and pooling paths). - traverse, find_one_and_update, find_one_and_delete and bulk_save_detailed took a raw pool connection without the app.tenant_id GUC, so under enable_rls they saw nothing (atomic ops returned None, traverse no hops) and COPY failed the policy's WITH CHECK before falling back to per-record saves. They now use the tenant-scoped connection (nested transactions become savepoints). - New RLS tests: two tenants sharing one hub id across find_connected_nodes / count_connected_nodes / _bulk, traverse, atomic ops and bulk_save_detailed (unprivileged role). - JVSPATIAL_POSTGRES_COMMAND_TIMEOUT for PostgresDB; postgres guide gains a pool-sizing rule of thumb and transaction-pooler notes. Co-Authored-By: Claude Opus 5 --- CHANGELOG.md | 10 ++ docs/md/environment-keys-reference.md | 1 + docs/md/multi-tenant-rls.md | 8 ++ docs/md/postgres-guide.md | 22 ++++ jvspatial/db/postgres.py | 26 ++-- jvspatial/env_adapter.py | 1 + tests/db/test_postgres_tenancy_graph.py | 167 ++++++++++++++++++++++++ 7 files changed, 224 insertions(+), 11 deletions(-) create mode 100644 tests/db/test_postgres_tenancy_graph.py diff --git a/CHANGELOG.md b/CHANGELOG.md index ccd1f0a..59956b5 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -38,6 +38,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - **`jvspatial.db.escape_regex()`** for building literal `$regex` patterns from user input (`$regex` is never index-backed). - **`find_edges_between(limit=...)`** pushes the limit to the database. +- **`JVSPATIAL_POSTGRES_COMMAND_TIMEOUT`** sets `PostgresDB`'s per-statement + timeout (default 60 s); the postgres guide gains a pool-sizing rule of thumb + and transaction-pooler notes for tenant scoping. - **`gin_index="off"` / `JVSPATIAL_PG_GIN_INDEX=off`** skips the whole-document `GIN (data jsonb_path_ops)` on new Postgres collections. - **`expand_node` keyset paging** (`jvspatial/core/graph_expansion.py`): @@ -91,6 +94,13 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 `edge=[{"E": {...}}]` returned neighbours reached through *any* edge type, and `direction="both"` capped at 10 000 edges. Edge types and criteria are now always applied. +- **Several Postgres paths ignored `PostgresDB.tenant(...)`** + (`jvspatial/db/postgres.py`). `traverse`, `find_one_and_update`, + `find_one_and_delete` and `bulk_save_detailed` took a raw pool connection + without the `app.tenant_id` GUC, so under `enable_rls` they saw nothing + (atomic ops returned `None`, `traverse` returned no hops) and the COPY path + failed the policy's `WITH CHECK` before falling back to per-record saves. + They now run on the tenant-scoped connection like every other operation. - **`observe=True` / caching wrappers turned `bulk_save_detailed` into one round trip per record** (`jvspatial/db/_observable.py`, `jvspatial/db/_cache.py`). Both wrappers subclass `Database`, whose default diff --git a/docs/md/environment-keys-reference.md b/docs/md/environment-keys-reference.md index d8cb69b..50a5992 100644 --- a/docs/md/environment-keys-reference.md +++ b/docs/md/environment-keys-reference.md @@ -51,6 +51,7 @@ For full examples and default values, see: - `JVSPATIAL_DYNAMODB_ENDPOINT_URL` - DynamoDB endpoint (e.g. LocalStack). - `JVSPATIAL_DYNAMODB_WAIT_FOR_INDEX` - Wait-for-index toggle. - `JVSPATIAL_NODE_EDGE_IDS` - Node adjacency mode: `derive` (edge collection only; Postgres / MongoDB / SQLite default) or `persist` (node rows also store `edges`; JSON / DynamoDB default). An `edge_ids_mode` set on the adapter instance wins. +- `JVSPATIAL_POSTGRES_COMMAND_TIMEOUT` - Postgres per-statement timeout in seconds (default 60). - `JVSPATIAL_PG_GIN_INDEX` - Postgres whole-document `GIN (data jsonb_path_ops)` index on new collections: `full` (default) or `off`. ### Auth and rate limit diff --git a/docs/md/multi-tenant-rls.md b/docs/md/multi-tenant-rls.md index 2f6828f..adc9180 100644 --- a/docs/md/multi-tenant-rls.md +++ b/docs/md/multi-tenant-rls.md @@ -48,6 +48,14 @@ async def handle_request(request, db): await Document.create(...) # tenant_id is stamped automatically ``` +Every `PostgresDB` read and write runs on a tenant-scoped connection inside +the block: `get` / `find` / `count`, `save` / `bulk_save_detailed` (the COPY +path), `find_one_and_update` / `find_one_and_delete`, the graph pushdowns +(`find_connected_nodes`, `count_connected_nodes`, `find_connected_nodes_bulk`) +and `traverse`. RLS applies per table, so a neighbour join sees only the +edges *and* the nodes the tenant may read: another tenant's edge that points +at a shared node id stays invisible. + The tenant scope uses `contextvars`, so: - Nested scopes shadow correctly: diff --git a/docs/md/postgres-guide.md b/docs/md/postgres-guide.md index 98b55dc..5f6e427 100644 --- a/docs/md/postgres-guide.md +++ b/docs/md/postgres-guide.md @@ -260,8 +260,23 @@ Or via env (see [environment-keys-reference.md](environment-keys-reference.md)): ```bash JVSPATIAL_POSTGRES_MIN_POOL_SIZE=5 JVSPATIAL_POSTGRES_MAX_POOL_SIZE=25 +JVSPATIAL_POSTGRES_COMMAND_TIMEOUT=30 # per-statement timeout, seconds (default 60) ``` +**Sizing rule of thumb.** Postgres does useful work on roughly two active +connections per database CPU core, and every open connection costs server +memory. Size the pools so their sum stays near that budget: + +```text +max_size ≈ (2 × DB host cores) ÷ number of app processes +``` + +Four app workers against an 8-core database: `max_size ≈ 16 ÷ 4 = 4`. A +bigger pool mostly moves the queue from the pool into Postgres. Keep +`max_size × processes` well below `max_connections`. For many short-lived +processes (Lambda, autoscaled containers) put PgBouncer / RDS Proxy in front +and size the pooler, not each process — see below. + ### Event loops The pool belongs to the event loop that created it. If a host bootstraps @@ -285,6 +300,12 @@ This disables asyncpg's statement cache (`statement_cache_size=0`) and binds parameters with the simple-query path, at the cost of some per-query overhead. Use it only when you're behind a transaction-pooling layer; direct or session-pooled connections should keep the default `pooler_mode="session"`. +Set it by environment with `JVSPATIAL_POSTGRES_POOLER_MODE=transaction`. + +Tenant scoping (`db.tenant(...)`) is transaction-pooler safe: every scoped +operation runs inside one transaction that begins with +`SELECT set_config('app.tenant_id', $1, true)` (a `SET LOCAL`), so the GUC +never leaks to the next client of a pooled server connection. ## Transactions @@ -335,6 +356,7 @@ precedence): | `JVSPATIAL_POSTGRES_MIN_POOL_SIZE` | Pool min size override | | `JVSPATIAL_POSTGRES_MAX_POOL_SIZE` | Pool max size override | | `JVSPATIAL_POSTGRES_POOLER_MODE` | `"session"` (default) or `"transaction"` | +| `JVSPATIAL_POSTGRES_COMMAND_TIMEOUT` | Per-statement timeout in seconds (default 60) | | `JVSPATIAL_NODE_EDGE_IDS` | `"derive"` (Postgres default) or `"persist"` | | `JVSPATIAL_PG_GIN_INDEX` | `"full"` (default) or `"off"` — whole-document GIN | diff --git a/jvspatial/db/postgres.py b/jvspatial/db/postgres.py index f432548..c50ddf3 100644 --- a/jvspatial/db/postgres.py +++ b/jvspatial/db/postgres.py @@ -305,7 +305,7 @@ def __init__( min_size: Optional[int] = None, max_size: Optional[int] = None, pooler_mode: str = "session", - command_timeout: float = 60.0, + command_timeout: Optional[float] = None, schema_name: str = "public", gin_index: Optional[str] = None, ) -> None: @@ -324,7 +324,8 @@ def __init__( direct or session-pooled connection) or ``"transaction"`` (compatible with PgBouncer / RDS Proxy in transaction-pool mode — disables statement cache, uses simple-query protocol). - command_timeout: Per-statement timeout in seconds. + command_timeout: Per-statement timeout in seconds. Also read from + ``JVSPATIAL_POSTGRES_COMMAND_TIMEOUT``; default 60. schema_name: Postgres schema to host the collection tables in. Defaults to ``public``. Must be an existing schema. gin_index: ``"full"`` (default) creates the whole-document @@ -367,7 +368,11 @@ def __init__( self.max_size = max_size if max_size is not None else (env_max or default_max) self.pooler_mode = pooler_mode - self.command_timeout = command_timeout + self.command_timeout = ( + command_timeout + if command_timeout is not None + else (env("JVSPATIAL_POSTGRES_COMMAND_TIMEOUT", parse=float) or 60.0) + ) self.schema_name = schema_name self.gin_index = ( (gin_index or env("JVSPATIAL_PG_GIN_INDEX", default="full")).strip().lower() @@ -1363,10 +1368,10 @@ async def bulk_save_detailed( self._split_payload(r) for r in records ] - pool = await self._ensure_pool() attempted = len(records) try: - async with pool.acquire() as conn: + # Tenant-scoped like every other write (RLS ``WITH CHECK``). + async with self._acquire_conn() as conn: async with conn.transaction(): await conn.execute( f""" @@ -1632,8 +1637,7 @@ async def find_one_and_delete( f"WHERE ctid = (SELECT ctid FROM {schema}.{col}{clause} LIMIT 1) " f"RETURNING data" ) - pool = await self._ensure_pool() - async with pool.acquire() as conn: + async with self._acquire_conn() as conn: row = await conn.fetchrow(sql, *params) return self._record_from_row(row) if row is not None else None @@ -1671,9 +1675,10 @@ async def find_one_and_update( where_sql, params = translated clause = f" WHERE {where_sql}" if where_sql else "" - pool = await self._ensure_pool() - async with pool.acquire() as conn: + # ``_acquire_conn`` applies the tenant scope; under it this + # transaction is a savepoint inside the tenant-scoped one. + async with self._acquire_conn() as conn: async with conn.transaction(): row = await conn.fetchrow( f"SELECT ctid, data FROM {schema}.{col}{clause} " @@ -2093,8 +2098,7 @@ async def traverse( ORDER BY node_id, depth ASC{limit_clause} """ - pool = await self._ensure_pool() - async with pool.acquire() as conn: + async with self._acquire_conn() as conn: rows = await conn.fetch(sql, *params) return [ { diff --git a/jvspatial/env_adapter.py b/jvspatial/env_adapter.py index deaaf10..dfbb229 100644 --- a/jvspatial/env_adapter.py +++ b/jvspatial/env_adapter.py @@ -268,6 +268,7 @@ def server_config_overrides_from_env() -> Dict[str, Any]: "JVSPATIAL_POSTGRES_MIN_POOL_SIZE", "JVSPATIAL_POSTGRES_MAX_POOL_SIZE", "JVSPATIAL_POSTGRES_POOLER_MODE", + "JVSPATIAL_POSTGRES_COMMAND_TIMEOUT", "JVSPATIAL_DYNAMODB_TABLE_NAME", "JVSPATIAL_DYNAMODB_REGION", "JVSPATIAL_DYNAMODB_ENDPOINT_URL", diff --git a/tests/db/test_postgres_tenancy_graph.py b/tests/db/test_postgres_tenancy_graph.py new file mode 100644 index 0000000..65498ed --- /dev/null +++ b/tests/db/test_postgres_tenancy_graph.py @@ -0,0 +1,167 @@ +"""Tenant isolation (row-level security) across the Postgres graph paths. + +RLS policies apply per table, so a join sees only the edges *and* the nodes +the active tenant may read. Two tenants share one hub id: ``acme`` owns the +hub and two leaves; ``beta`` owns an edge from that hub id to its own leaf. +Every read and write path must honour ``PostgresDB.tenant(...)`` — including +the ones that manage their own connection or transaction. + +Runs as an unprivileged role (superusers bypass RLS) when +``JVSPATIAL_POSTGRES_TEST_DSN`` is reachable. +""" + +from __future__ import annotations + +import asyncio +import os +import uuid +from typing import Any, AsyncIterator, Dict +from urllib.parse import urlparse, urlunparse + +import pytest + +pytestmark = pytest.mark.asyncio(loop_scope="module") + +_DSN = os.getenv( + "JVSPATIAL_POSTGRES_TEST_DSN", + "postgresql://jvspatial:jvspatial@localhost:5432/jvspatial", +) + +HUB = "n.Hub.h" + + +@pytest.fixture +async def rls_db() -> AsyncIterator[Any]: + """PostgresDB as a NOSUPERUSER NOBYPASSRLS role, RLS on node and edge.""" + try: + import asyncpg + + from jvspatial.db.postgres import PostgresDB + except ImportError: + pytest.skip("asyncpg not installed") + try: + admin = await asyncio.wait_for(asyncpg.connect(dsn=_DSN), timeout=2.0) + except Exception: + pytest.skip(f"Postgres unreachable at {_DSN}") + schema = f"jvs_tgraph_{uuid.uuid4().hex[:10]}" + role = f"jvs_tgraph_{uuid.uuid4().hex[:8]}" + await admin.execute(f'CREATE SCHEMA "{schema}"') + await admin.execute( + f"CREATE ROLE \"{role}\" WITH LOGIN PASSWORD 'tgraph-pw' " + "NOSUPERUSER NOBYPASSRLS" + ) + await admin.execute(f'GRANT USAGE, CREATE ON SCHEMA "{schema}" TO "{role}"') + parsed = urlparse(_DSN) + dsn = urlunparse( + parsed._replace( + netloc=f"{role}:tgraph-pw@{parsed.hostname}:{parsed.port or 5432}" + ) + ) + db = PostgresDB(dsn=dsn, schema_name=schema, min_size=1, max_size=4) + try: + await db.enable_rls("node") + await db.enable_rls("edge") + await _seed(db) + yield db + finally: + await db.close() + await admin.execute(f'DROP SCHEMA "{schema}" CASCADE') + await admin.execute(f'DROP ROLE IF EXISTS "{role}"') + await admin.close() + + +def _node(nid: str, tenant: str, **ctx: Any) -> Dict[str, Any]: + return {"id": nid, "entity": nid.split(".")[1], "tenant_id": tenant, "context": ctx} + + +def _edge(eid: str, src: str, tgt: str, tenant: str) -> Dict[str, Any]: + return { + "id": eid, + "entity": "Link", + "tenant_id": tenant, + "context": {}, + "source": src, + "target": tgt, + "bidirectional": False, + } + + +async def _seed(db: Any) -> None: + async with db.tenant("acme"): + for rec in ( + _node(HUB, "acme"), + _node("n.Leaf.a1", "acme"), + _node("n.Leaf.a2", "acme"), + ): + await db.save("node", rec) + for rec in ( + _edge("e.Link.a1", HUB, "n.Leaf.a1", "acme"), + _edge("e.Link.a2", HUB, "n.Leaf.a2", "acme"), + ): + await db.save("edge", rec) + async with db.tenant("beta"): + await db.save("node", _node("n.Leaf.b1", "beta")) + await db.save("edge", _edge("e.Link.b1", HUB, "n.Leaf.b1", "beta")) + + +async def _ids(rows: Any) -> set: + return {r["id"] for r in rows} + + +async def test_neighbour_joins_see_only_the_tenant(rls_db): + for tenant, expected in ( + ("acme", {"n.Leaf.a1", "n.Leaf.a2"}), + ("beta", {"n.Leaf.b1"}), + ): + async with rls_db.tenant(tenant): + for kwargs in ({}, {"edge_entities": ["Link"]}, {"direction": "both"}): + rows = await rls_db.find_connected_nodes("node", "edge", HUB, **kwargs) + assert await _ids(rows) == expected, (tenant, kwargs) + assert await rls_db.count_connected_nodes( + "node", "edge", HUB, **kwargs + ) == len(expected) + bulk = await rls_db.find_connected_nodes_bulk("node", "edge", [HUB]) + assert await _ids(bulk[HUB]) == expected + assert await rls_db.find_connected_nodes("node", "edge", HUB) == [] + + +async def test_traverse_sees_only_the_tenant(rls_db): + async with rls_db.tenant("acme"): + hops = await rls_db.traverse("edge", HUB, max_depth=2) + assert {h["node_id"] for h in hops} == {"n.Leaf.a1", "n.Leaf.a2"} + async with rls_db.tenant("beta"): + hops = await rls_db.traverse("edge", HUB, max_depth=2) + assert {h["node_id"] for h in hops} == {"n.Leaf.b1"} + + +async def test_atomic_ops_respect_the_tenant(rls_db): + async with rls_db.tenant("acme"): + doc = await rls_db.find_one_and_update( + "node", {"_id": "n.Leaf.a1"}, {"$set": {"context.seen": True}} + ) + assert doc is not None and doc["context"]["seen"] is True + async with rls_db.tenant("beta"): + assert ( + await rls_db.find_one_and_update( + "node", {"_id": "n.Leaf.a1"}, {"$set": {"context.seen": False}} + ) + is None + ) + assert await rls_db.find_one_and_delete("node", {"_id": "n.Leaf.a2"}) is None + async with rls_db.tenant("acme"): + assert (await rls_db.get("node", "n.Leaf.a1"))["context"]["seen"] is True + deleted = await rls_db.find_one_and_delete("node", {"_id": "n.Leaf.a2"}) + assert deleted is not None and deleted["id"] == "n.Leaf.a2" + + +async def test_bulk_save_writes_within_the_tenant(rls_db): + async with rls_db.tenant("acme"): + result = await rls_db.bulk_save_detailed( + "node", [_node(f"n.Leaf.bulk{i}", "acme") for i in range(3)] + ) + assert result.all_saved + assert {"n.Leaf.bulk0", "n.Leaf.bulk2"} <= await _ids( + await rls_db.find("node", {}) + ) + async with rls_db.tenant("beta"): + assert await rls_db.get("node", "n.Leaf.bulk0") is None From 380af2284128c69e62ba279e84ffabbb79b3fc0f Mon Sep 17 00:00:00 2001 From: Eldon Marks Date: Sun, 13 Sep 2026 15:33:15 -0400 Subject: [PATCH 8/9] =?UTF-8?q?Release=200.0.18=20=E2=80=94=20hub-node=20s?= =?UTF-8?q?cale=20remediation?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Derive adjacency from the edge table, push neighbour filters/limits/counts into SQL, entity-leading Postgres indexes, optional GIN, and $text search. Final pre-release bench at 51186f1 confirms gates still hold. Co-authored-by: Cursor --- CHANGELOG.md | 2 ++ docs/bench/2026-09-hub-node-baseline.md | 39 +++++++++++++++++++++++++ docs/bench/2026-09-hub-node-final.jsonl | 4 +++ jvspatial/version.py | 2 +- 4 files changed, 46 insertions(+), 1 deletion(-) create mode 100644 docs/bench/2026-09-hub-node-final.jsonl diff --git a/CHANGELOG.md b/CHANGELOG.md index 59956b5..f2f24d7 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,8 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +## [0.0.18] - 2026-09-13 + ### Added - **Hub-node scale recorder** (`tests/benchmarks/test_hub_node_bench.py`, diff --git a/docs/bench/2026-09-hub-node-baseline.md b/docs/bench/2026-09-hub-node-baseline.md index 8b6c465..9a77632 100644 --- a/docs/bench/2026-09-hub-node-baseline.md +++ b/docs/bench/2026-09-hub-node-baseline.md @@ -234,3 +234,42 @@ Seeding note: the Phase 3 hub tier seeds through `create_database(observe=True)` which now forwards `bulk_save_detailed` to COPY (fixed in `905d122`). Earlier tiers seeded per record; that only affected setup time, not the measured operations. + +## Final pre-release — `51186f1` (Phase 5 + wrappers/tenant fixes) + +Raw: [`2026-09-hub-node-final.jsonl`](2026-09-hub-node-final.jsonl). Same machine and +Postgres image as the Phase 0–3 runs. Confirms the full stack still holds the +gates before cutting 0.0.18. + +| operation | 1k p50 / p95 ms (trips) | 10k p50 / p95 ms (trips) | 100k p50 / p95 ms (trips) | +|---|---|---|---| +| `ctx.get(Hub)` (hydrate hub) | 0.7 / 1.0 | 0.6 / 1.2 | 0.9 / 1.9 | +| `hub.connect(leaf, edge=E)` | 2.1 / 3.3 (3) | 1.9 / 2.8 (3) | 2.1 / 3.4 (3) | +| `hub.save()` after scalar change | 0.8 / 1.3 | 0.8 / 4.2 | 0.8 / 1.4 | +| `hub.nodes(edge=[E], node=['Leaf'], limit=20)` | 1.2 / 1.5 (1) | 1.2 / 1.8 (1) | 1.7 / 2.6 (1) | +| `hub.nodes(edge=E, limit=20)` | 1.4 / 2.0 (1) | 1.4 / 1.8 (1) | 2.0 / 2.0 (1) | +| `sink.nodes(edge=[E], node=['Leaf'], direction='in', limit=20)` | 1.3 / 2.2 (1) | 1.3 / 1.9 (1) | 1.6 / 2.0 (1) | +| `sink.nodes(edge=[E], node=['Leaf'], direction='in')` | 23.1 / 48.7 (1) | 274.7 / 281.4 (1) | 2493.6 / 2600.1 (1) | +| `len(await hub.nodes(edge=[E]))` | 23.0 / 50.7 (1) | 276.8 / 344.6 (1) | 2367.4 / 2500.3 (1) | +| `hub.count_nodes(edge=[E])` | 3.8 / 4.2 (1) | 43.0 / 47.9 (1) | 103.3 / 110.7 (1) | +| 32× concurrent `connect()` (ms) | 18 wall / 18 max | 19 wall / 19 max | 17 wall / 16 max | + +| size (after seed → after 82 hub writes) | 1k | 10k | 100k | +|---|---|---|---| +| hub row `pg_column_size(data)` | 134 B → 136 B | 134 B → 136 B | 134 B → 136 B | +| `node_data_gin` size | 104.0 KB → 280.0 KB | 808.0 KB → 816.0 KB | 10.5 MB → 10.5 MB | +| `node` table total | 504.0 KB → 704.0 KB | 4.0 MB → 4.0 MB | 41.8 MB → 41.8 MB | +| `edge` table total | 1.8 MB → 2.1 MB | 16.5 MB → 16.5 MB | 161.7 MB → 161.8 MB | + +| typed find (sorted, limit 20) | rows | p50 / p95 ms | index walked | sort node | node_data_gin | +|---|---|---|---|---|---| +| `51186f1` | 1,000,000 | 1.03 / 1.81 | node_entity_context_track_id_context_created_at_idx | no | off | + +**Gates: still passed.** + +- A1: `connect()` 100k p95 is 3.4 ms vs 3.3 ms at 1k (≈ 1.0×; bar ≤ 1.5×). + `save()` p50 stays ~0.8 ms at every tier. +- A3: list-form `nodes(..., limit=20)` is 1 round trip and flat (~1–3 ms). +- A6: typed find p95 1.81 ms, index-bound, GIN off. +- Concurrent `connect()` wall ≈ 17–19 ms at every degree (no row-lock serialisation). +- Numbers sit within run-to-run noise of the Phase 3 table. diff --git a/docs/bench/2026-09-hub-node-final.jsonl b/docs/bench/2026-09-hub-node-final.jsonl new file mode 100644 index 0000000..259b974 --- /dev/null +++ b/docs/bench/2026-09-hub-node-final.jsonl @@ -0,0 +1,4 @@ +{"degree": 1000, "jvspatial": "0.0.17", "git_sha": "51186f1", "edge_ids_mode": "derive", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 0.09, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 134, "node_gin_bytes": 106496, "node_total_bytes": 516096, "edge_total_bytes": 1859584}, "hub_get": {"p50_ms": 0.65, "p95_ms": 0.984, "max_ms": 2.674, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 1.246, "p95_ms": 1.496, "max_ms": 2.882, "n": 20, "round_trips": 1}, "nodes_class_out_limit20": {"p50_ms": 1.358, "p95_ms": 2.01, "max_ms": 2.244, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 1.294, "p95_ms": 2.212, "max_ms": 2.227, "n": 20, "round_trips": 1}, "nodes_list_in_unlimited": {"p50_ms": 23.119, "p95_ms": 48.749, "max_ms": 52.302, "n": 20, "round_trips": 1}, "count_via_len_nodes": {"p50_ms": 23.008, "p95_ms": 50.687, "max_ms": 51.235, "n": 20, "round_trips": 1}, "count_nodes": {"p50_ms": 3.82, "p95_ms": 4.165, "max_ms": 4.476, "n": 20, "round_trips": 1}, "hub_save": {"p50_ms": 0.78, "p95_ms": 1.344, "max_ms": 3.236, "n": 50}, "hub_connect": {"p50_ms": 2.143, "p95_ms": 3.277, "max_ms": 4.345, "n": 50, "round_trips": 3}, "concurrent_connect_32": {"wall_ms": 18.152, "max_call_ms": 17.984, "p50_call_ms": 15.989}, "sizes_after_writes": {"hub_row_bytes": 136, "node_gin_bytes": 286720, "node_total_bytes": 720896, "edge_total_bytes": 2211840}} +{"degree": 10000, "jvspatial": "0.0.17", "git_sha": "51186f1", "edge_ids_mode": "derive", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 1.04, "read_iterations": 20, "sizes_after_seed": {"hub_row_bytes": 134, "node_gin_bytes": 827392, "node_total_bytes": 4153344, "edge_total_bytes": 17268736}, "hub_get": {"p50_ms": 0.64, "p95_ms": 1.178, "max_ms": 1.489, "n": 20}, "nodes_list_out_limit20": {"p50_ms": 1.23, "p95_ms": 1.801, "max_ms": 2.203, "n": 20, "round_trips": 1}, "nodes_class_out_limit20": {"p50_ms": 1.398, "p95_ms": 1.759, "max_ms": 2.027, "n": 20, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 1.295, "p95_ms": 1.944, "max_ms": 2.163, "n": 20, "round_trips": 1}, "nodes_list_in_unlimited": {"p50_ms": 274.702, "p95_ms": 281.38, "max_ms": 283.277, "n": 20, "round_trips": 1}, "count_via_len_nodes": {"p50_ms": 276.75, "p95_ms": 344.581, "max_ms": 443.018, "n": 20, "round_trips": 1}, "count_nodes": {"p50_ms": 42.965, "p95_ms": 47.948, "max_ms": 60.398, "n": 20, "round_trips": 1}, "hub_save": {"p50_ms": 0.785, "p95_ms": 4.25, "max_ms": 7.578, "n": 50}, "hub_connect": {"p50_ms": 1.929, "p95_ms": 2.833, "max_ms": 4.063, "n": 50, "round_trips": 3}, "concurrent_connect_32": {"wall_ms": 18.718, "max_call_ms": 18.567, "p50_call_ms": 16.747}, "sizes_after_writes": {"hub_row_bytes": 136, "node_gin_bytes": 835584, "node_total_bytes": 4169728, "edge_total_bytes": 17350656}} +{"degree": 100000, "jvspatial": "0.0.17", "git_sha": "51186f1", "edge_ids_mode": "derive", "postgres": "16.14 (Debian 16.14-1.pgdg12+1)", "seed_seconds": 12.22, "read_iterations": 5, "sizes_after_seed": {"hub_row_bytes": 134, "node_gin_bytes": 11042816, "node_total_bytes": 43835392, "edge_total_bytes": 169566208}, "hub_get": {"p50_ms": 0.906, "p95_ms": 1.862, "max_ms": 1.862, "n": 5}, "nodes_list_out_limit20": {"p50_ms": 1.682, "p95_ms": 2.644, "max_ms": 2.644, "n": 5, "round_trips": 1}, "nodes_class_out_limit20": {"p50_ms": 1.958, "p95_ms": 2.046, "max_ms": 2.046, "n": 5, "round_trips": 1}, "nodes_list_in_limit20": {"p50_ms": 1.628, "p95_ms": 2.047, "max_ms": 2.047, "n": 5, "round_trips": 1}, "nodes_list_in_unlimited": {"p50_ms": 2493.623, "p95_ms": 2600.06, "max_ms": 2600.06, "n": 5, "round_trips": 1}, "count_via_len_nodes": {"p50_ms": 2367.435, "p95_ms": 2500.263, "max_ms": 2500.263, "n": 5, "round_trips": 1}, "count_nodes": {"p50_ms": 103.293, "p95_ms": 110.708, "max_ms": 110.708, "n": 5, "round_trips": 1}, "hub_save": {"p50_ms": 0.804, "p95_ms": 1.355, "max_ms": 2.568, "n": 50}, "hub_connect": {"p50_ms": 2.133, "p95_ms": 3.376, "max_ms": 4.679, "n": 50, "round_trips": 3}, "concurrent_connect_32": {"wall_ms": 16.666, "max_call_ms": 16.285, "p50_call_ms": 15.309}, "sizes_after_writes": {"hub_row_bytes": 136, "node_gin_bytes": 11042816, "node_total_bytes": 43843584, "edge_total_bytes": 169672704}} +{"kind": "typed_find", "rows": 1000000, "jvspatial": "0.0.17", "git_sha": "51186f1", "seed_seconds": 29.75, "node_data_gin": false, "typed_find_sorted_limit20": {"p50_ms": 1.028, "p95_ms": 1.813, "max_ms": 3.038, "n": 50}, "plan_indexes": ["node_entity_context_track_id_context_created_at_idx"], "plan_sorts": false} diff --git a/jvspatial/version.py b/jvspatial/version.py index d2fe638..71ee54d 100644 --- a/jvspatial/version.py +++ b/jvspatial/version.py @@ -9,4 +9,4 @@ # - MAJOR: Breaking changes # - MINOR: New features, backward compatible # - PATCH: Bug fixes, backward compatible -__version__ = "0.0.17" +__version__ = "0.0.18" From b2e4cb1885e84e64785fbf521fc16ef4c4b8c10e Mon Sep 17 00:00:00 2001 From: Eldon Marks Date: Sun, 13 Sep 2026 15:45:45 -0400 Subject: [PATCH 9/9] chore: keep hub-node fixtures and docs domain-agnostic MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Rename bench/index fixture field track_id → group_id, neutralize consumer names in CHANGELOG, and use hub/Leaf in optimization examples so the remediation stays framework-generic. Co-authored-by: Cursor --- CHANGELOG.md | 2 +- docs/bench/2026-09-hub-node-baseline.md | 18 ++++++++----- docs/md/optimization.md | 12 ++++----- tests/benchmarks/test_hub_node_bench.py | 20 +++++++-------- tests/db/test_postgres_indexes_text.py | 34 ++++++++++++------------- 5 files changed, 46 insertions(+), 40 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index f2f24d7..3f2f756 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -594,7 +594,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Changed -- **BREAKING (limited):** `JsonDBTransaction(db).save/get/delete/find()` now raises `NotImplementedError` by default instead of silently no-op'ing. Pass `best_effort=True` to opt into the buffered-commit semantics, or check `Database.supports_transactions` and fall back to non-transactional writes. Audited downstream consumers (`jvagent`, `integral`) — neither uses this surface, so no coordinated change required. +- **BREAKING (limited):** `JsonDBTransaction(db).save/get/delete/find()` now raises `NotImplementedError` by default instead of silently no-op'ing. Pass `best_effort=True` to opt into the buffered-commit semantics, or check `Database.supports_transactions` and fall back to non-transactional writes. Audited known adopters — none use this surface, so no coordinated change required. - **`QueryEngine` optimization cache** is now bounded by an LRU (default 1024 entries, configurable via `QueryEngine(cache_size=...)`). Was unbounded. - **MongoDB retries** now go through the shared retry helper (`utils/retry.py`); behavior preserved (one retry on connection-error with reset). - CI coverage gate raised from 55 → 60% to reflect the new tested code. diff --git a/docs/bench/2026-09-hub-node-baseline.md b/docs/bench/2026-09-hub-node-baseline.md index 9a77632..83d009d 100644 --- a/docs/bench/2026-09-hub-node-baseline.md +++ b/docs/bench/2026-09-hub-node-baseline.md @@ -193,10 +193,13 @@ Code: `dc2a6cc` + the Phase 2 change set. Code: `905d122` + the Phase 3 change set. **Typed-find gate.** `test_typed_find_is_index_bound` runs -`find({"entity": "BenchEntry", "context.track_id": t}, sort=[("context.created_at", -1)], limit=20)` +`find({"entity": "BenchEntry", "context.group_id": g}, sort=[("context.created_at", -1)], limit=20)` 50 times on a shared `node` table of 1M rows. The rows are spread over ten -entities, and `BenchEntry` holds 100k of them: 1000 tracks × 100. The class -declares `@compound_index([("track_id", 1), ("created_at", -1)])`. +entities, and `BenchEntry` holds 100k of them: 1000 groups × 100. The class +declares `@compound_index([("group_id", 1), ("created_at", -1)])`. +(Archived runs below still show the earlier fixture field name `track_id` in +index identifiers; the field was renamed to `group_id` so the harness stays +domain-agnostic.) | typed find (sorted, limit 20) | rows | p50 / p95 ms | index walked | sort node | node_data_gin | |---|---|---|---|---|---| @@ -205,11 +208,11 @@ declares `@compound_index([("track_id", 1), ("created_at", -1)])`. **Gate: passed.** p95 is 2.96 ms, well under the 10 ms bar, and the plan is index-bound with the whole-document GIN off. One scan of -`(entity, track_id, created_at DESC NULLS LAST)` returns the first 20 rows +`(entity, , created_at DESC NULLS LAST)` returns the first 20 rows with no Sort node. On 0.0.17 the same class index skipped `entity` and ordered descending keys `NULLS FIRST`. Postgres therefore had to AND it with -the entity index and sort every matching row. At this track size that still -stays under 10 ms locally, but the cost grows with the rows per track. +the entity index and sort every matching row. At this group size that still +stays under 10 ms locally, but the cost grows with the rows per group. Hub-node numbers on the same code: @@ -265,6 +268,9 @@ gates before cutting 0.0.18. |---|---|---|---|---|---| | `51186f1` | 1,000,000 | 1.03 / 1.81 | node_entity_context_track_id_context_created_at_idx | no | off | +(Field renamed `track_id` → `group_id` after this run; same compound shape, +new index name is `…_group_id_…`.) + **Gates: still passed.** - A1: `connect()` 100k p95 is 3.4 ms vs 3.3 ms at 1k (≈ 1.0×; bar ≤ 1.5×). diff --git a/docs/md/optimization.md b/docs/md/optimization.md index 2118cdf..3637f08 100644 --- a/docs/md/optimization.md +++ b/docs/md/optimization.md @@ -61,21 +61,21 @@ user = await current_node.node(node=User, direction="out") # Efficient On Postgres, MongoDB and SQLite every `nodes()` filter shape — classes, names, lists, `{Name: criteria}` dicts, property kwargs — plus `limit` runs as -one query, so `nodes(edge=[Contains], node=["Entry"], limit=20)` costs the same +one query, so `nodes(edge=[Contains], node=["Leaf"], limit=20)` costs the same on a node with 20 neighbours as on one with 100 000. Don't count or page by listing: ```python # Bad: hydrates every neighbour just to count or slice them -total = len(await track.nodes(edge=[Contains])) -recent = (await track.nodes(edge=[Contains]))[:20] +total = len(await hub.nodes(edge=[Contains])) +recent = (await hub.nodes(edge=[Contains]))[:20] # Good: one COUNT, and keyset pages that stay O(page) -total = await track.count_nodes(edge=[Contains]) -page, cursor = await track.nodes_page( +total = await hub.count_nodes(edge=[Contains]) +page, cursor = await hub.nodes_page( edge=[Contains], sort=[("context.created_at", -1)], limit=20 ) -more, cursor = await track.nodes_page( +more, cursor = await hub.nodes_page( edge=[Contains], sort=[("context.created_at", -1)], limit=20, cursor=cursor ) ``` diff --git a/tests/benchmarks/test_hub_node_bench.py b/tests/benchmarks/test_hub_node_bench.py index 80bef7e..befb8f2 100644 --- a/tests/benchmarks/test_hub_node_bench.py +++ b/tests/benchmarks/test_hub_node_bench.py @@ -82,11 +82,11 @@ class BenchContains(Edge): """Typed containment edge (hub -> leaf, leaf -> sink).""" -@compound_index([("track_id", 1), ("created_at", -1)], name="bench_track_recent") +@compound_index([("group_id", 1), ("created_at", -1)], name="bench_group_recent") class BenchEntry(Node): - """Record-style node for the typed-find gate (sorted, limited per track).""" + """Record-style node for the typed-find gate (sorted, limited per group).""" - track_id: str = "" + group_id: str = "" created_at: str = "" @@ -457,8 +457,8 @@ async def test_typed_find_is_index_bound( """Sorted, limited typed ``find`` on a shared ``node`` table of N rows. ``N`` defaults to 1M (``JVSPATIAL_BENCH_TYPED_ROWS``) spread over ten - entities; ``BenchEntry`` holds a tenth, 1000 tracks of ~100 rows each. - Measures ``find({"entity": "BenchEntry", "context.track_id": t}, + entities; ``BenchEntry`` holds a tenth, 1000 groups of ~100 rows each. + Measures ``find({"entity": "BenchEntry", "context.group_id": t}, sort=created_at desc, limit=20)`` with the whole-document GIN off and records whether the plan walks an index without a Sort node. """ @@ -486,7 +486,7 @@ async def test_typed_find_is_index_bound( "id": generate_id("n", entity), "entity": entity, "context": { - "track_id": f"t{(i // len(entities)) % 1000}", + "group_id": f"t{(i // len(entities)) % 1000}", "created_at": f"2026-09-{i:08d}", }, } @@ -499,12 +499,12 @@ async def test_typed_find_is_index_bound( seed_s = time.perf_counter() - t0 await admin.execute(f"ANALYZE {schema}.node") - per_track = min(20, rows // len(entities) // 1000) + per_group = min(20, rows // len(entities) // 1000) samples: List[float] = [] for k in range(50): query = { "entity": "BenchEntry", - "context.track_id": f"t{(k * 37) % 1000}", + "context.group_id": f"t{(k * 37) % 1000}", } ms, _, out = await _timed( functools.partial( @@ -516,12 +516,12 @@ async def test_typed_find_is_index_bound( ) ) samples.append(ms) - assert len(out) == per_track + assert len(out) == per_group from jvspatial.db._postgres_translate import translate_query, translate_sort where, params = translate_query( - {"entity": "BenchEntry", "context.track_id": "t7"} + {"entity": "BenchEntry", "context.group_id": "t7"} ) order = translate_sort([("context.created_at", -1)]) raw = await admin.fetchval( diff --git a/tests/db/test_postgres_indexes_text.py b/tests/db/test_postgres_indexes_text.py index ed4498e..0a4658b 100644 --- a/tests/db/test_postgres_indexes_text.py +++ b/tests/db/test_postgres_indexes_text.py @@ -45,9 +45,9 @@ ) -@compound_index([("track_id", 1), ("created_at", -1)], name="track_recent") +@compound_index([("group_id", 1), ("created_at", -1)], name="group_recent") class IdxEntry(Node): - track_id: str = attribute(indexed=True, default="") + group_id: str = attribute(indexed=True, default="") created_at: str = "" note: str = attribute(indexed=True, index_partial_by_entity=True, default="") @@ -126,25 +126,25 @@ async def _explain(admin: Any, sql: str, params: List[Any]) -> List[Dict[str, An async def test_per_class_indexes_are_entity_scoped(): async with _pg() as (ctx, db, admin, schema): - await db.create_index("node", "context.track_id") # pre-0.0.18 unscoped - assert "node_context_track_id_idx" in await _indexes(admin, schema, "node") + await db.create_index("node", "context.group_id") # pre-0.0.18 unscoped + assert "node_context_group_id_idx" in await _indexes(admin, schema, "node") await ctx.ensure_indexes(IdxEntry) idx = await _indexes(admin, schema, "node") - single = idx["node_entity_context_track_id_idx"] - assert "(entity, ((data #>> '{context,track_id}'::text[])))" in single - compound = idx["node_entity_context_track_id_context_created_at_idx"] + single = idx["node_entity_context_group_id_idx"] + assert "(entity, ((data #>> '{context,group_id}'::text[])))" in single + compound = idx["node_entity_context_group_id_context_created_at_idx"] assert compound.startswith("CREATE INDEX") assert "DESC NULLS LAST" in compound partial = idx["node_idxentry_context_note_idx"] assert "WHERE (entity = 'IdxEntry'::text)" in partial - assert "node_context_track_id_idx" not in idx # legacy replaced + assert "node_context_group_id_idx" not in idx # legacy replaced async def test_stale_edge_index_is_rebuilt_on_entity_column(): async with _pg() as (ctx, db, admin, schema): - await ctx.save(IdxEntry(track_id="t")) # bootstraps the tables + await ctx.save(IdxEntry(group_id="t")) # bootstraps the tables await db._bootstrap_collection("edge") await admin.execute( f"CREATE UNIQUE INDEX edge_source_target_entity_uniq ON {schema}.edge " @@ -165,7 +165,7 @@ async def test_typed_sorted_find_walks_the_entity_leading_index(): { "id": generate_id("n", entity), "entity": entity, - "context": {"track_id": f"t{i % 40}", "created_at": f"2026-01-{i:05d}"}, + "context": {"group_id": f"t{i % 40}", "created_at": f"2026-01-{i:05d}"}, } for entity in ("IdxEntry", "Other1", "Other2", "Other3", "Other4") for i in range(1500) @@ -174,7 +174,7 @@ async def test_typed_sorted_find_walks_the_entity_leading_index(): await admin.execute(f"ANALYZE {schema}.node") where, params = translate_query( - {"entity": "IdxEntry", "context.track_id": "t7"} + {"entity": "IdxEntry", "context.group_id": "t7"} ) order = translate_sort([("context.created_at", -1)]) await admin.execute(f"SET search_path TO {schema}") @@ -184,12 +184,12 @@ async def test_typed_sorted_find_walks_the_entity_leading_index(): params, ) names = {p.get("Index Name") for p in plan} - assert "node_entity_context_track_id_context_created_at_idx" in names, plan + assert "node_entity_context_group_id_context_created_at_idx" in names, plan assert not any(p["Node Type"] == "Sort" for p in plan), plan found = await db.find( "node", - {"entity": "IdxEntry", "context.track_id": "t7"}, + {"entity": "IdxEntry", "context.group_id": "t7"}, sort=[("context.created_at", -1)], limit=20, ) @@ -202,11 +202,11 @@ async def test_typed_sorted_find_walks_the_entity_leading_index(): async def test_gin_index_can_be_turned_off(monkeypatch, caplog): async with _pg() as (ctx, db, admin, schema): - await ctx.save(IdxEntry(track_id="t")) + await ctx.save(IdxEntry(group_id="t")) assert "node_data_gin" in await _indexes(admin, schema, "node") async with _pg(gin_index="off") as (ctx, db, admin, schema): - await ctx.save(IdxEntry(track_id="t")) + await ctx.save(IdxEntry(group_id="t")) assert "node_data_gin" not in await _indexes(admin, schema, "node") with caplog.at_level(logging.WARNING, logger="jvspatial.db.postgres"): await db.find("node", {"context.tags": {"$all": ["a"]}}) @@ -386,9 +386,9 @@ async def test_find_edges_between_limit(tmp_path): ctx = GraphContext(database=JsonDB(base_path=str(tmp_path / "j"))) set_default_context(ctx) try: - hub = await IdxEntry.create(track_id="hub") + hub = await IdxEntry.create(group_id="hub") for i in range(5): - await hub.connect(await IdxEntry.create(track_id=f"l{i}"), edge=IdxLink) + await hub.connect(await IdxEntry.create(group_id=f"l{i}"), edge=IdxLink) assert len(await ctx.find_edges_between(hub.id, edge_class=IdxLink)) == 5 assert ( len(await ctx.find_edges_between(hub.id, edge_class=IdxLink, limit=2)) == 2