diff --git a/CHANGELOG.md b/CHANGELOG.md index 2dc35d6..3f2f756 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,128 @@ 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`, + `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`. +- **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). +- **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`. +- **`$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. +- **`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`): + `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"`. +- **`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`. +- **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 + are fetched with a single `find_many`. + +### 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. +- **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 + `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 + 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 + 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 + `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 @@ -472,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/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/SPEC.md b/SPEC.md index 5a74fa5..68df27b 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. @@ -290,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. @@ -364,6 +367,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 new file mode 100644 index 0000000..83d009d --- /dev/null +++ b/docs/bench/2026-09-hub-node-baseline.md @@ -0,0 +1,281 @@ +# 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`. +- **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`. +- **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-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 +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 / 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)` | 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. 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 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. + - `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 + (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 ~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 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 — adjacency derived from the edge table (`edge_ids_mode=derive`) + +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 — neighbour filters, limits and counts pushed into SQL + +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 — entity-leading indexes, optional GIN, `$text` + +Code: `905d122` + the Phase 3 change set. + +**Typed-find gate.** `test_typed_find_is_index_bound` runs +`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 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 | +|---|---|---|---|---|---| +| 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, , 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 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: + +| 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. + +## 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 | + +(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×). + `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/docs/bench/2026-09-hub-node-phase0.jsonl b/docs/bench/2026-09-hub-node-phase0.jsonl new file mode 100644 index 0000000..810fd3d --- /dev/null +++ b/docs/bench/2026-09-hub-node-phase0.jsonl @@ -0,0 +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-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/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/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/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/docs/md/entity-reference.md b/docs/md/entity-reference.md index 6527a05..efe721f 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,10 +64,17 @@ 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 +- **`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/environment-keys-reference.md b/docs/md/environment-keys-reference.md index 932db0f..50a5992 100644 --- a/docs/md/environment-keys-reference.md +++ b/docs/md/environment-keys-reference.md @@ -50,6 +50,9 @@ 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. +- `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 - `JVSPATIAL_AUTH_ENABLED` - Enables auth. @@ -125,6 +128,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/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/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/optimization.md b/docs/md/optimization.md index 5c7857f..3637f08 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=["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 hub.nodes(edge=[Contains])) +recent = (await hub.nodes(edge=[Contains]))[:20] + +# Good: one COUNT, and keyset pages that stay O(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 hub.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/docs/md/postgres-guide.md b/docs/md/postgres-guide.md index 06b7fb5..5f6e427 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 @@ -159,6 +211,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 @@ -181,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 @@ -206,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 @@ -256,6 +356,9 @@ 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 | ## 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/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 b201043..227a650 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,7 +21,7 @@ cast, ) -from jvspatial.db.database import Database, 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 @@ -426,6 +424,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 +902,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 +971,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 +1097,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 +1114,7 @@ async def expand_node( direction=direction, limit=limit, cursor=cursor, + after=after, detail_level=detail_level, # type: ignore[arg-type] ) @@ -1228,8 +1255,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 +1301,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: @@ -1411,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( @@ -1481,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( @@ -1508,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 @@ -1532,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} @@ -1598,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: @@ -1605,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 @@ -1662,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"], @@ -1702,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. @@ -1710,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: @@ -1743,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: @@ -1839,9 +1882,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 +1922,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..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. @@ -196,8 +316,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 +352,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 +366,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 +428,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, @@ -369,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 @@ -386,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( @@ -491,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, @@ -628,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 @@ -639,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, @@ -1132,16 +1316,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 +1368,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 +1431,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 +1502,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 +1520,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 +1561,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 +1669,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/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/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/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 2f80a94..98f42d5 100644 --- a/jvspatial/db/README.md +++ b/jvspatial/db/README.md @@ -59,9 +59,13 @@ 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`. +- **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`) - **`QueryEngine` LRU is bounded.** Default 1024; configurable. Unbounded query construction will not leak memory. (`query.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/_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 964a4dc..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]]: @@ -335,60 +351,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..6d85bfc 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})") @@ -503,17 +507,56 @@ 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 ---------------------------------------------------- -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 +569,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 +581,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 +591,20 @@ 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 == "$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) + piece = _translate_field_clause(key, value, pb, table) if piece is None: return None fragments.append(piece) @@ -561,8 +612,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 +632,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 +644,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 +660,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 +669,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: @@ -611,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/_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/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..e83fc19 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 ( @@ -58,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.""" @@ -70,6 +94,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 +324,200 @@ 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)) + + @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, @@ -310,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: @@ -476,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( @@ -493,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 8f01d70..c50ddf3 100644 --- a/jvspatial/db/postgres.py +++ b/jvspatial/db/postgres.py @@ -78,12 +78,13 @@ Dict, List, Optional, + Sequence, Set, Tuple, Union, ) -from ._postgres_translate import translate_query, translate_sort +from ._postgres_translate import translate_query, translate_sort, tsvector_expression from .database import ( BulkSaveResult, Database, @@ -165,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}')``. @@ -271,6 +293,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, @@ -278,8 +305,9 @@ 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: """Initialize the Postgres adapter. @@ -296,13 +324,21 @@ 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 + ``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( @@ -332,8 +368,20 @@ 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() + ) + 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() @@ -570,6 +618,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 @@ -585,8 +639,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 @@ -721,7 +774,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 +824,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) @@ -831,6 +945,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: @@ -928,55 +1055,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: @@ -1009,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""" @@ -1074,13 +1433,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) @@ -1091,34 +1475,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=..., @@ -1152,19 +1543,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) ------------------------------------------ @@ -1201,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 @@ -1240,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} " @@ -1662,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/db/query.py b/jvspatial/db/query.py index 93f7d73..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): @@ -579,6 +639,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..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, @@ -67,6 +77,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, @@ -729,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/jvspatial/env_adapter.py b/jvspatial/env_adapter.py index c65523d..dfbb229 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 @@ -256,11 +268,14 @@ 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", "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", @@ -322,6 +337,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/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" 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/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 diff --git a/tests/benchmarks/hub_bench_report.py b/tests/benchmarks/hub_bench_report.py new file mode 100644 index 0000000..7d37c88 --- /dev/null +++ b/tests/benchmarks/hub_bench_report.py @@ -0,0 +1,137 @@ +"""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(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] + 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..befb8f2 --- /dev/null +++ b/tests/benchmarks/test_hub_node_bench.py @@ -0,0 +1,562 @@ +"""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.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 +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).""" + + +@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 group).""" + + group_id: str = "" + created_at: str = "" + + +# ---- 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, + # 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, + ) + 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, + ), + } + 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 + for _ in range(read_iters): + await ctx.clear_cache() + ms, trips, out = await _timed(fn) + samples.append(ms) + 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 -- + 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 + + # 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 + 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") + + +@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 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. + """ + 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": { + "group_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_group = min(20, rows // len(entities) // 1000) + samples: List[float] = [] + for k in range(50): + query = { + "entity": "BenchEntry", + "context.group_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_group + + from jvspatial.db._postgres_translate import translate_query, translate_sort + + where, params = translate_query( + {"entity": "BenchEntry", "context.group_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/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_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_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]} 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_postgres_indexes_text.py b/tests/db/test_postgres_indexes_text.py new file mode 100644 index 0000000..0a4658b --- /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([("group_id", 1), ("created_at", -1)], name="group_recent") +class IdxEntry(Node): + group_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.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_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_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(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 " + "((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": {"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) + ] + await db.bulk_save_detailed("node", records) + await admin.execute(f"ANALYZE {schema}.node") + + where, params = translate_query( + {"entity": "IdxEntry", "context.group_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_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.group_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(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(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"]}}) + 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(group_id="hub") + for i in range(5): + 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 + ) + finally: + set_default_context(None) 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 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/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") 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)