Compare commits
3 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 1ed97fb075 | |||
| e44341e47c | |||
| ccc54e8f18 |
@@ -71,5 +71,6 @@ src/barabadb/storage/lsm
|
|||||||
src/barabadb/storage/wal
|
src/barabadb/storage/wal
|
||||||
src/barabadb/storage/btree
|
src/barabadb/storage/btree
|
||||||
src/barabadb/storage/gate
|
src/barabadb/storage/gate
|
||||||
|
src/barabadb/protocol/scram
|
||||||
clients/nim/tests/test_pool
|
clients/nim/tests/test_pool
|
||||||
clients/nim/tests/test_wire
|
clients/nim/tests/test_wire
|
||||||
|
|||||||
@@ -0,0 +1,80 @@
|
|||||||
|
# BaraDB — Deep Audit (август 2026)
|
||||||
|
|
||||||
|
> Дата: 2026-08-02
|
||||||
|
> Метод: 4 паралелни одит-агента по слоеве (Storage / Query / Core / Protocol), всеки чете всички файлове в обхвата си и проверява находките срещу реалния код.
|
||||||
|
> Обхват: **само нови дефекти** — 80-те вече оправени в `BUGS.md` / `BUG_AUDIT.md` / `BARADB_CLIENT_BUGS.md` са изключени.
|
||||||
|
> **Общо: ~28 находки | Поправени: 28 (батч 1: 5 + батч 2: 12 + батч 3: 10 + батч 4: 2) | Остават: 0**
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Поправени — батч 1 (5)
|
||||||
|
|
||||||
|
| # | Severity | Проблем | Файл | Fix |
|
||||||
|
|---|----------|---------|------|-----|
|
||||||
|
| C1 | 🔴 CRITICAL | **MIGRATE handler без auth gate** — неавтентикиран клиент пишеше произволни key/value в базата (`handleMigrationMessage → applyMigrationBatch → storeKeys → db.put`). Открито независимо от 2 агента. | `core/server.nim:682` | Добавен `if not authenticated: ... continue` (като REP/DISTTXN блоковете) |
|
||||||
|
| C2 | 🔴 CRITICAL | **Raft commit quorum off-by-one за even-N** — `(N+1) div 2` commit-ваше с малцинство при четен брой възли (N=4 → 2/4). Election-ът ползваше коректното strict majority. *GA обхватът е 3-node (нечетно), където формулите съвпадат.* | `core/raft.nim:653` | `let majority = (node.peers.len + 1) div 2 + 1` (съвпада с election); регресионен тест за 4-node |
|
||||||
|
| H1 | 🟠 HIGH | **Pre-auth memory-exhaustion DoS** — `parseHeader` не ограничаваше `length` (uint32 до ~4 GiB); `recvExactWithTimeout` пре-алокира преди auth check. | `core/server.nim:168` | Reject `length > uint32(MaxWireStringLen)` (64 MB) преди алокация |
|
||||||
|
| H4 | 🟠 HIGH | **`**` и `++` се lower-ваха към equality** — `bkPow`/`bkConcat` липсваха в op-mapping case-а и попадаха в `else: irEq` (`2 ** 3` → `false`, `'a' ++ 'b'` → `false`). | `query/exec/lower.nim:79` | `of bkPow: irOp = irPow`, `of bkConcat: irOp = irAdd`; 2 регресионни теста |
|
||||||
|
| H5 | 🟠 HIGH | **`!=` не е отрицание на `=`** — `irNeq` short-circuit-ваше на string inequality, така че `5 != 5.0` → true, но `5 = 5.0` → true. | `query/exec/eval.nim:438` | `irNeq` numeric-first (точно допълнение на `irEq`); регресионен тест |
|
||||||
|
|
||||||
|
## Поправени — батч 2 (12)
|
||||||
|
|
||||||
|
| # | Severity | Проблем | Файл | Fix |
|
||||||
|
|---|----------|---------|------|-----|
|
||||||
|
| H3 | 🟠 HIGH | **Semi-sync partial/zero ack** — връщаше LSN дори при 0 acks | `core/replication.nim` | `return 0` когато connected replicas < `syncReplicaCount` acks; 0 connected → local-only (като sync) |
|
||||||
|
| H6 | 🟠 HIGH | **`COUNT/SUM/AVG(DISTINCT)` игнорира DISTINCT** | `query/exec/lower.nim`, `plan_exec.nim` | `aggDistinct = node.funcDistinct`; dedup с `HashSet` в agg пътищата |
|
||||||
|
| H7 | 🟠 HIGH | **`UNION/INTERSECT/EXCEPT` KeyError** | `query/executor.nim` | Dedup fingerprint от projected cols, не `row["$value"]` |
|
||||||
|
| H8 | 🟠 HIGH | **`MERGE … THEN DELETE` / matched condition no-op** | `query/executor.nim` | Honor `mergeMatchedDelete` + `mergeMatchedCondition` |
|
||||||
|
| H9 | 🟠 HIGH | **WAL recovery crash на torn record** | `storage/lsm.nim`, `wal.nim`, `recovery.nim` | Bound key/val ≤ 64 MB; validate kind преди enum cast |
|
||||||
|
| M1 | 🟡 MEDIUM | **MVCC `write` delete-during-iteration** | `core/mvcc.nim` | Collect-then-delete stale txn ids |
|
||||||
|
| M3 | 🟡 MEDIUM | **`checkpoint` lock leak** | `storage/lsm.nim` | try/finally около write lock + walLock |
|
||||||
|
| M4 | 🟡 MEDIUM | **`flushUnsafe` clear-before-write** | `storage/lsm.nim` | Clear memtable едва след успешен `writeSSTable` |
|
||||||
|
| M5 | 🟡 MEDIUM | **Compaction empty-key skip** | `storage/compaction.nim` | `haveLast` флаг вместо `lastKey = ""` sentinel |
|
||||||
|
| M6 | 🟡 MEDIUM | **`rewriteLive` remove-before-move** | `storage/wal.nim` | Само атомен `moveFile` (rename replace) |
|
||||||
|
| L3 | 🟢 LOW | **mmap `offset+size` overflow** | `storage/mmap.nim` | Overflow-safe: `offset > size - length` |
|
||||||
|
| — | hygiene | **Stray ELF `protocol/scram`** | `.gitignore` | Премахнат binary + ignore entry |
|
||||||
|
|
||||||
|
**Верификация (батч 2):** `baradadb` build чист; `tests/bugfix_test.nim` (вкл. batch-2 suite) и `tests/test_all.nim` (501 OK) минават без `[FAILED]`. `tests/prop_test.nim` B-Tree suite OK (H10 *не* е в този батч — naive left-max fix чупи interleaved remove).
|
||||||
|
|
||||||
|
## Поправени — батч 3 (10)
|
||||||
|
|
||||||
|
| # | Severity | Проблем | Файл | Fix |
|
||||||
|
|---|----------|---------|------|-----|
|
||||||
|
| H2 | 🟠 HIGH | **TLS client връзките не верифицираха сертификата** | `core/server.nim`, `core/config.nim` | Отделен `tlsClient` контекст; CA auto-enable verify; production fail-closed |
|
||||||
|
| M2 | 🟡 MEDIUM | **disttxn refused-connect wedge** | `core/disttxn.nim` | `getsockopt(SO_ERROR)` + try/except около RPC |
|
||||||
|
| M7 | 🟡 MEDIUM | **Compaction unlink преди catalog load** | `storage/compaction.nim`, `baradadb.nim` | Unlink след load + MANIFEST |
|
||||||
|
| M8 | 🟡 MEDIUM | **`OFFSET` без `LIMIT` → 0 реда** | `query/exec/lower.nim`, `plan_exec.nim` | `limitCount = -1` unlimited; clamp negative |
|
||||||
|
| M9 | 🟡 MEDIUM | **Window SUM/AVG/COUNT/MIN/MAX → NULL** | `query/exec/window.nim` | Frame aggregates |
|
||||||
|
| M10 | 🟡 MEDIUM | **WebSocket unmasked client frames** | `core/websocket.nim` | Protocol error / close |
|
||||||
|
| M11 | 🟡 MEDIUM | **WebSocket unbounded buffer** | `core/websocket.nim` | 1 MiB frame / 4 MiB message / 125-byte control |
|
||||||
|
| M12 | 🟡 MEDIUM | **SUBSCRIBE без table auth** | `core/websocket.nim`, `httpserver.nim` | `canSubscribe` + `hasPrivilegeFor` SELECT |
|
||||||
|
| L1 | 🟢 LOW | **SCRAM timing user enumeration** | `protocol/auth.nim` | Dummy nonce+encode work за unknown users |
|
||||||
|
| L2 | 🟢 LOW | **SCRAM `c=` не се верифицира** | `protocol/auth.nim`, `scram.nim` | `c=` must match gs2 header (`biws` за `n,,`) |
|
||||||
|
|
||||||
|
## Поправени — батч 4 (2)
|
||||||
|
|
||||||
|
| # | Severity | Проблем | Файл | Fix |
|
||||||
|
|---|----------|---------|------|-----|
|
||||||
|
| H10 | 🟠 HIGH | **B-tree `remove` separator convention** — `splitChild`/search са left-max (`key > sep → right`); `removeRec` копираше first key на дясното дете (right-min). Naive left-max върху *internal* ключове чупи `prop_test`. | `storage/btree.nim` | Leaf borrow/remove пишат left-max; merge underflow се качва нагоре; invariant `max(left) <= sep` (boundary duplicates са позволени, `next` ги събира) |
|
||||||
|
| L4 | 🟢 LOW | **NULL equality** — `NULL = NULL` / `col = NULL` бяха true през string sentinel | `query/exec/eval.nim` | Сравнения, LIKE, IN, NOT, AND/OR: NULL operand → unknown (`\N`); `IS NULL` непроменен |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Остават (0)
|
||||||
|
|
||||||
|
Няма отворени находки от този одит.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Проверени и чисти (не са бъгове)
|
||||||
|
|
||||||
|
- JWT `exp`/alg-confusion: pinned `jwt-nim-baraba#fbe084b` `verify()` enforce-ва alg-match, reject-ва `NONE`, проверява `exp/nbf/iat`, constant-time compare.
|
||||||
|
- `auth.nim` `constantTimeCompare` и SCRAM `verifyClientProof` са constant-time; празен JWT secret fail-closed (`server.nim:59`).
|
||||||
|
- `wire.nim` deserialize bounds/depth caps са sound.
|
||||||
|
- CRC byte ranges / `headerSize = 40` са консистентни между write/verify/load (format *коментарът* още казва "36" — само коментар).
|
||||||
|
- Lock ordering (`walLock` в `db.lock`; gate преди `db.lock`), mmap negative offset / `close()` recursion / fd handling (BUG-036/046) — непокътнати.
|
||||||
|
- `COUNT(col)` изключва NULLs (`v.kind != vkNull`); `LIMIT 0` → празно е коректно; IN-list се lower-ва към OR/AND вериги.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
*Виж също: `PLAN.md` (Сесия 13), `docs/en/known-limitations.md`.*
|
||||||
@@ -2,6 +2,53 @@
|
|||||||
|
|
||||||
All notable changes to BaraDB are documented in this file.
|
All notable changes to BaraDB are documented in this file.
|
||||||
|
|
||||||
|
## [Unreleased]
|
||||||
|
|
||||||
|
### Security
|
||||||
|
|
||||||
|
- **MIGRATE auth bypass (CRITICAL)** — the internal `MIGRATE` text-protocol handler now requires authentication (matching the `REP`/`DISTTXN` handlers); previously an unauthenticated client could inject arbitrary key/value rows (`core/server.nim`)
|
||||||
|
- **Pre-auth wire-length DoS (HIGH)** — `parseHeader` rejects messages larger than the 64 MB wire cap before allocating the receive buffer (`core/server.nim`)
|
||||||
|
- **TLS peer verify on cluster forwarding (HIGH)** — follower→leader SQL forwarding uses a dedicated client TLS context; `BARADB_TLS_CA_FILE` auto-enables `BARADB_TLS_VERIFY_PEER`; production fails closed if TLS is on without CA+verify (`core/config.nim`, `core/server.nim`)
|
||||||
|
- **WebSocket unmasked frames / unbounded buffers (MEDIUM)** — unmasked client frames are rejected (RFC 6455 §5.1); frame/message size and control-frame caps bound memory (`core/websocket.nim`)
|
||||||
|
- **WebSocket SUBSCRIBE table auth (MEDIUM)** — with auth enabled, SUBSCRIBE requires table SELECT privilege (`core/websocket.nim`, `core/httpserver.nim`)
|
||||||
|
- **SCRAM timing + channel-binding (LOW)** — unknown users do equivalent dummy work; `c=` must match the gs2 header (`protocol/auth.nim`, `protocol/scram.nim`)
|
||||||
|
|
||||||
|
### Fixed
|
||||||
|
|
||||||
|
- **Raft commit quorum (CRITICAL)** — commit now requires a strict majority (`N div 2 + 1`), matching the election check; the previous `(N+1) div 2` formula committed at a minority for even-sized clusters (`core/raft.nim`)
|
||||||
|
- **`**` / `++` operators (HIGH)** — power and concat are no longer lowered to equality: `2 ** 3` → 8, `'a' ++ 'b'` → `'ab'` (`query/exec/lower.nim`)
|
||||||
|
- **`!=` semantics (HIGH)** — `!=` is now the exact complement of `=` for numerically-equal values (`5 != 5.0` is false) (`query/exec/eval.nim`)
|
||||||
|
- **Semi-sync partial ack (HIGH)** — `writeLsn` in `rmSemiSync` returns `0` when connected replicas fail to meet `syncReplicaCount`; zero connected peers still succeed local-only (like sync) (`core/replication.nim`)
|
||||||
|
- **`COUNT/SUM/AVG(DISTINCT …)` (HIGH)** — `funcDistinct` is copied to `aggDistinct` and applied via `HashSet` dedup in aggregate paths (`query/exec/lower.nim`, `plan_exec.nim`)
|
||||||
|
- **`UNION` / `INTERSECT` / `EXCEPT` (HIGH)** — set-op dedup fingerprints projected columns instead of missing `row["$value"]` (KeyError crash) (`query/executor.nim`)
|
||||||
|
- **`MERGE … WHEN MATCHED THEN DELETE` (HIGH)** — executor honors `mergeMatchedDelete` and optional `mergeMatchedCondition` (`query/executor.nim`)
|
||||||
|
- **WAL recovery torn records (HIGH)** — recovery bounds key/value to 64 MB and rejects out-of-range entry kinds before enum cast (avoids multi-GiB alloc / `CaseStmtError` Defect) (`storage/lsm.nim`, `wal.nim`, `recovery.nim`)
|
||||||
|
- **MVCC `write` timeout cleanup** — stale active transactions are collected then deleted (no mutation during `activeTxns` iteration) (`core/mvcc.nim`)
|
||||||
|
- **`checkpoint` lock leak** — write lock and `walLock` released in `try/finally` (`storage/lsm.nim`)
|
||||||
|
- **`flushUnsafe` data-loss window** — memtable is cleared only after a successful SSTable write (`storage/lsm.nim`)
|
||||||
|
- **Compaction empty-string key** — dedup uses a `haveLast` flag so key `""` is not skipped (`storage/compaction.nim`)
|
||||||
|
- **`rewriteLive` crash window** — atomic `moveFile` replace only (no `removeFile` before rename) (`storage/wal.nim`)
|
||||||
|
- **mmap OOB on overflow** — length checks use `offset > size - length` instead of wrapping `offset + size` (`storage/mmap.nim`)
|
||||||
|
- **REP replication put/delete encoding** — the legacy REP payload carries an explicit op tag so PK-only inserts (empty value) replicate as puts instead of vanishing as deletes (`core/replication.nim`, `core/server.nim`)
|
||||||
|
- **REP receiver secondary indexes** — the legacy REP receiver applies via `applyReplicatedPut/Delete` under the storage gate, keeping B-tree/FTS/HNSW/graph indexes consistent on the replica (`core/server.nim`)
|
||||||
|
- **Snapshot send stall (partial)** — the leader's snapshot send runs gzip off the event loop on a worker thread (`gzipFileAsync`), so heartbeats keep firing during compression; tar (send) and the restore path still run on the loop (`core/backup.nim`, `core/raft.nim`)
|
||||||
|
- **disttxn refused-connect wedge (MEDIUM)** — `connectWithTimeout` checks `SO_ERROR`; `sendDistTxnRpc` catches `CatchableError` so a refused peer cannot leave 2PC stuck (`core/disttxn.nim`)
|
||||||
|
- **Compaction catalog order (MEDIUM)** — inputs stay on disk until the output is loaded and MANIFEST is written; then they are unlinked (`storage/compaction.nim`, `baradadb.nim`)
|
||||||
|
- **`OFFSET n` without `LIMIT` (MEDIUM)** — no longer treated as `LIMIT 0`; negative LIMIT/OFFSET are clamped (`query/exec/lower.nim`, `plan_exec.nim`)
|
||||||
|
- **Window `SUM`/`AVG`/`COUNT`/`MIN`/`MAX` (MEDIUM)** — aggregate window functions compute over the frame instead of returning NULL (`query/exec/window.nim`)
|
||||||
|
- **B-tree leaf separators (HIGH)** — leaf borrow/remove now keep the left-max convention used by `splitChild` and search (`key > sep → right`); merge underflow rebalances up the tree; `checkInvariants` guards `max(left) <= sep` (`storage/btree.nim`)
|
||||||
|
- **NULL comparison three-valued logic (LOW)** — `NULL = NULL` / `col = NULL` / `col != x` yield unknown (`\N`) so `WHERE` excludes them; `IS NULL` is unchanged (`query/exec/eval.nim`)
|
||||||
|
|
||||||
|
### Removed
|
||||||
|
|
||||||
|
- Stray compiled ELF `src/barabadb/protocol/scram` from the source tree (added to `.gitignore`)
|
||||||
|
|
||||||
|
### Added
|
||||||
|
|
||||||
|
- Deep audit report `BUG_AUDIT_2026-08.md` (~28 findings; all fixed across batches 1–4)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
## [1.3.0] — 2026-07-30
|
## [1.3.0] — 2026-07-30
|
||||||
|
|
||||||
### Raft cluster — Supported (single `default` DB scope)
|
### Raft cluster — Supported (single `default` DB scope)
|
||||||
|
|||||||
@@ -145,6 +145,34 @@
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
## Сесия 13: Stabilization & Deep Audit (август 2026)
|
||||||
|
|
||||||
|
> **Цел**: Подобряване на надеждността и коректността върху вече завършеното ядро — целеви поправки + системен паралелен одит по слоеве (Storage / Query / Core / Protocol).
|
||||||
|
|
||||||
|
### Целеви поправки (завършени)
|
||||||
|
|
||||||
|
| # | Поправка | Файлове | Статус |
|
||||||
|
|---|----------|---------|--------|
|
||||||
|
| 1 | **REP delete-from-empty** — legacy REP payload носи явен put/delete таг (`encodeRepPayload`/`decodeRepPayload`); PK-only редове вече не изчезват при репликация | `core/replication.nim`, `core/server.nim` | ✅ + тестове |
|
||||||
|
| 2 | **Snapshot stall (частична mitigation)** — gzip при leader snapshot send се изнася извън event loop-а през worker thread (`gzipFileAsync`); heartbeats текат по време на компресия | `core/backup.nim`, `core/raft.nim`, `baradadb.nim` | ✅ + e2e |
|
||||||
|
| 3 | **REP receiver индекси** — receiver-ът минава през `applyReplicatedPut/Delete` под storage gate; вторичните индекси (B-tree/FTS/HNSW/graph) се поддържат на репликата | `core/server.nim` | ✅ + тест |
|
||||||
|
|
||||||
|
### Deep Audit (2026-08)
|
||||||
|
|
||||||
|
4 паралелни одит-агента по слоеве; ~28 нови находки (без дублиране на 80-те вече оправени в `BUGS.md`/`BUG_AUDIT.md`). Пълен отчет: [`BUG_AUDIT_2026-08.md`](BUG_AUDIT_2026-08.md).
|
||||||
|
|
||||||
|
**Батч 1 — поправени (5):** MIGRATE auth bypass (CRITICAL), raft commit strict-majority за even-N (CRITICAL), pre-auth wire-length DoS (HIGH), `**`/`++` lowering към equality (HIGH), `!=` не е отрицание на `=` (HIGH).
|
||||||
|
|
||||||
|
**Батч 2 — поправени (12):** semi-sync partial-ack (H3), COUNT/SUM/AVG(DISTINCT) (H6), UNION/INTERSECT/EXCEPT (H7), MERGE THEN DELETE (H8), WAL torn-record recovery (H9), MVCC write iteration (M1), checkpoint lock leak (M3), flushUnsafe order (M4), compaction empty-key (M5), rewriteLive atomic replace (M6), mmap overflow (L3), stray `protocol/scram` ELF.
|
||||||
|
|
||||||
|
**Батч 3 — поправени (10):** TLS peer verify (H2), disttxn SO_ERROR (M2), compaction catalog order (M7), OFFSET-без-LIMIT (M8), window агрегати (M9), WebSocket (M10–M12), SCRAM (L1–L2).
|
||||||
|
|
||||||
|
**Батч 4 — поправени (2):** B-tree leaf left-max separators (H10), NULL three-valued comparisons (L4).
|
||||||
|
|
||||||
|
**Остават (0):** вж. `BUG_AUDIT_2026-08.md`.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
## Какво остава от старите планове
|
## Какво остава от старите планове
|
||||||
|
|
||||||
| Стар план | Статус |
|
| Стар план | Статус |
|
||||||
@@ -157,6 +185,7 @@
|
|||||||
| **Този план** — Сесии 10, 11, 12 | ✅ Завършен |
|
| **Този план** — Сесии 10, 11, 12 | ✅ Завършен |
|
||||||
| Raft C3a/C3b + DDL/forward/compact/metrics (2026-07-30) | ✅ Завършен на `main` — `docs/superpowers/specs/2026-07-30-raft-cluster-status.md` |
|
| Raft C3a/C3b + DDL/forward/compact/metrics (2026-07-30) | ✅ Завършен на `main` — `docs/superpowers/specs/2026-07-30-raft-cluster-status.md` |
|
||||||
| **Production GA v1.2.0** (single-node) | ✅ `docs/superpowers/plans/2026-07-30-production-ga.md` |
|
| **Production GA v1.2.0** (single-node) | ✅ `docs/superpowers/plans/2026-07-30-production-ga.md` |
|
||||||
|
| **Сесия 13** — Stabilization & Deep Audit (2026-08) | ✅ Батч 1–4 (28 поправки); `BUG_AUDIT_2026-08.md` |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -170,4 +199,4 @@
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
*План версия: 2026-05-17*
|
*План версия: 2026-08-02*
|
||||||
|
|||||||
@@ -46,11 +46,11 @@ Documented in [distributed.md](distributed.md). Supported scope:
|
|||||||
|
|
||||||
## Newly documented limitations
|
## Newly documented limitations
|
||||||
|
|
||||||
- **Legacy non-raft REP replication infers delete from empty value** — the non-raft replication path still treats an empty value as a delete, so inserts into a PK-only table are misapplied over that path (the row vanishes). Use raft replication instead.
|
- **Legacy non-raft REP replication delete inference** — resolved: the non-raft REP payload now carries an explicit put/delete op tag (`encodeRepPayload`/`decodeRepPayload` in `core/replication.nim`), so PK-only inserts (empty LSM value) replicate as puts instead of being misapplied as deletes.
|
||||||
- **Snapshot-restore ctx staleness** — after an InstallSnapshot restore, HTTP endpoints using the startup-captured ctx may serve stale data until the node is restarted; the `/query` path is fresh per-request. Pre-existing client connections likewise see pre-restore state — reconnect after a restore.
|
- **Snapshot-restore ctx staleness** — after an InstallSnapshot restore, HTTP endpoints using the startup-captured ctx may serve stale data until the node is restarted; the `/query` path is fresh per-request. Pre-existing client connections likewise see pre-restore state — reconnect after a restore.
|
||||||
- **FK-cascade divergence under raft** — `ON DELETE/UPDATE CASCADE` (and `SET NULL`) effects are not raft-replicated: followers only apply the parent row's KV change, so cascaded child rows persist on followers. Avoid FK actions on raft-replicated tables, or accept periodic snapshot resync.
|
- **FK-cascade divergence under raft** — `ON DELETE/UPDATE CASCADE` (and `SET NULL`) effects are not raft-replicated: followers only apply the parent row's KV change, so cascaded child rows persist on followers. Avoid FK actions on raft-replicated tables, or accept periodic snapshot resync.
|
||||||
- **Uncommitted writes in snapshots** — the leader applies writes locally before raft majority commit; a snapshot taken in that window can include writes that never commit (phantom rows after restore + leadership change). Narrow window; fix tracked for a later release.
|
- **Uncommitted writes in snapshots** — the leader applies writes locally before raft majority commit; a snapshot taken in that window can include writes that never commit (phantom rows after restore + leadership change). Narrow window; fix tracked for a later release.
|
||||||
- **Event-loop stall during snapshot build/restore** — snapshot build/restore performs blocking tar/gzip on the node's event loop; large data dirs can stall heartbeats and trigger an election mid-transfer.
|
- **Event-loop stall during snapshot build/restore** — partially mitigated: the leader's snapshot *send* now tars under the storage gate but runs the CPU-heavy gzip off the event loop on a worker thread (`gzipFileAsync` in `core/backup.nim`), so heartbeats keep firing during compression. The tar itself (send path) and the whole *restore* path (tar extract + DB reopen) still run on the event loop, so very large data dirs can still stall heartbeats during those phases; a full fix (an async/try-lock storage gate so the loop never blocks) is tracked for a later release.
|
||||||
|
|
||||||
## Operational requirements
|
## Operational requirements
|
||||||
|
|
||||||
|
|||||||
@@ -26,6 +26,8 @@ import std/strutils
|
|||||||
import std/times
|
import std/times
|
||||||
import std/algorithm
|
import std/algorithm
|
||||||
import std/json
|
import std/json
|
||||||
|
import std/asyncdispatch
|
||||||
|
import std/threadpool
|
||||||
import barabadb/storage/lsm
|
import barabadb/storage/lsm
|
||||||
|
|
||||||
type
|
type
|
||||||
@@ -260,6 +262,80 @@ proc backupDataDir*(dataDir: string, output: string, excludes: seq[string] = @[]
|
|||||||
echo " Source: ", dataDir
|
echo " Source: ", dataDir
|
||||||
return true
|
return true
|
||||||
|
|
||||||
|
proc tarDataDir*(dataDir: string, output: string, excludes: seq[string] = @[]): bool =
|
||||||
|
## Create an UNCOMPRESSED tar of `dataDir` at `output` (no gzip). The raft
|
||||||
|
## snapshot sender runs this under the storage gate for a consistent file
|
||||||
|
## capture, then compresses off the event loop via gzipFileAsync.
|
||||||
|
if not dirExists(dataDir):
|
||||||
|
echo "ERROR: Data directory not found: ", dataDir
|
||||||
|
return false
|
||||||
|
|
||||||
|
let parent = parentDir(dataDir)
|
||||||
|
let name = lastPathPart(dataDir)
|
||||||
|
var excludeArgs = ""
|
||||||
|
for pattern in excludes:
|
||||||
|
excludeArgs.add(" --exclude=" & quoteShell(pattern))
|
||||||
|
|
||||||
|
let cmd = "tar -cf " & quoteShell(output) & excludeArgs &
|
||||||
|
" -C " & quoteShell(parent) & " " & quoteShell(name)
|
||||||
|
let (outputStr, exitCode) = execCmdEx(cmd)
|
||||||
|
if exitCode != 0:
|
||||||
|
echo "ERROR: tar command failed with exit code ", exitCode
|
||||||
|
if outputStr.len > 0:
|
||||||
|
echo outputStr
|
||||||
|
return false
|
||||||
|
return true
|
||||||
|
|
||||||
|
proc gzipFile*(input: string, output: string,
|
||||||
|
compression: int = DEFAULT_COMPRESSION): bool =
|
||||||
|
## gzip a single file `input` -> `output`. Pure CPU over an already-captured
|
||||||
|
## file: no shared storage state, so it is safe to run on a worker thread
|
||||||
|
## outside the storage gate and off the raft event loop.
|
||||||
|
if not fileExists(input):
|
||||||
|
echo "ERROR: File not found: ", input
|
||||||
|
return false
|
||||||
|
|
||||||
|
let cmd = "gzip -" & $compression & " -c " & quoteShell(input) &
|
||||||
|
" > " & quoteShell(output)
|
||||||
|
let (outputStr, exitCode) = execCmdEx("bash -c " & quoteShell(cmd))
|
||||||
|
if exitCode != 0:
|
||||||
|
echo "ERROR: gzip command failed with exit code ", exitCode
|
||||||
|
if outputStr.len > 0:
|
||||||
|
echo outputStr
|
||||||
|
return false
|
||||||
|
return true
|
||||||
|
|
||||||
|
proc gunzipFile*(input: string, output: string): bool =
|
||||||
|
## Decompress a gzip file `input` -> `output`. Inverse of gzipFile.
|
||||||
|
if not fileExists(input):
|
||||||
|
echo "ERROR: File not found: ", input
|
||||||
|
return false
|
||||||
|
|
||||||
|
let cmd = "gzip -dc " & quoteShell(input) & " > " & quoteShell(output)
|
||||||
|
let (outputStr, exitCode) = execCmdEx("bash -c " & quoteShell(cmd))
|
||||||
|
if exitCode != 0:
|
||||||
|
echo "ERROR: gunzip command failed with exit code ", exitCode
|
||||||
|
if outputStr.len > 0:
|
||||||
|
echo outputStr
|
||||||
|
return false
|
||||||
|
return true
|
||||||
|
|
||||||
|
proc gzipFileWorker(input: string, output: string, compression: int): bool {.gcsafe.} =
|
||||||
|
## Thread entry point: touches only its own (copied) string args + execCmdEx,
|
||||||
|
## so it is safe to run off the main/event-loop thread under ARC/ORC.
|
||||||
|
gzipFile(input, output, compression)
|
||||||
|
|
||||||
|
proc gzipFileAsync*(input: string, output: string,
|
||||||
|
compression: int = DEFAULT_COMPRESSION): Future[bool] {.async.} =
|
||||||
|
## Run gzipFile on a threadpool worker and await completion WITHOUT blocking
|
||||||
|
## the calling async event loop — heartbeats/election timers keep firing
|
||||||
|
## during the CPU-heavy compression. Polls the FlowVar via sleepAsync so the
|
||||||
|
## dispatcher stays responsive instead of stalling on a blocking join.
|
||||||
|
var fv = spawn gzipFileWorker(input, output, compression)
|
||||||
|
while not fv.isReady:
|
||||||
|
await sleepAsync(20)
|
||||||
|
result = ^fv
|
||||||
|
|
||||||
proc restoreDataDir*(input: string, dataDir: string, verbose: bool = false, dryRun: bool = false): bool =
|
proc restoreDataDir*(input: string, dataDir: string, verbose: bool = false, dryRun: bool = false): bool =
|
||||||
## Restore from a tar.gz backup.
|
## Restore from a tar.gz backup.
|
||||||
## When dryRun is true, only prints what would be done.
|
## When dryRun is true, only prints what would be done.
|
||||||
|
|||||||
@@ -14,6 +14,8 @@ type
|
|||||||
tlsEnabled*: bool
|
tlsEnabled*: bool
|
||||||
certFile*: string
|
certFile*: string
|
||||||
keyFile*: string
|
keyFile*: string
|
||||||
|
tlsCaFile*: string
|
||||||
|
tlsVerifyPeer*: bool
|
||||||
idleTimeoutMs*: int
|
idleTimeoutMs*: int
|
||||||
queryTimeoutMs*: int
|
queryTimeoutMs*: int
|
||||||
slowQueryThresholdMs*: int
|
slowQueryThresholdMs*: int
|
||||||
@@ -67,6 +69,8 @@ proc defaultConfig*(): BaraConfig =
|
|||||||
tlsEnabled: false,
|
tlsEnabled: false,
|
||||||
certFile: "",
|
certFile: "",
|
||||||
keyFile: "",
|
keyFile: "",
|
||||||
|
tlsCaFile: "",
|
||||||
|
tlsVerifyPeer: false,
|
||||||
idleTimeoutMs: 300_000,
|
idleTimeoutMs: 300_000,
|
||||||
queryTimeoutMs: 30_000,
|
queryTimeoutMs: 30_000,
|
||||||
slowQueryThresholdMs: 1_000,
|
slowQueryThresholdMs: 1_000,
|
||||||
@@ -134,6 +138,8 @@ proc loadConfigFromJson*(path: string, cfg: var BaraConfig) =
|
|||||||
if s.hasKey("enabled"): cfg.tlsEnabled = s["enabled"].getBool()
|
if s.hasKey("enabled"): cfg.tlsEnabled = s["enabled"].getBool()
|
||||||
if s.hasKey("cert_file"): cfg.certFile = s["cert_file"].getStr()
|
if s.hasKey("cert_file"): cfg.certFile = s["cert_file"].getStr()
|
||||||
if s.hasKey("key_file"): cfg.keyFile = s["key_file"].getStr()
|
if s.hasKey("key_file"): cfg.keyFile = s["key_file"].getStr()
|
||||||
|
if s.hasKey("ca_file"): cfg.tlsCaFile = s["ca_file"].getStr()
|
||||||
|
if s.hasKey("verify_peer"): cfg.tlsVerifyPeer = s["verify_peer"].getBool()
|
||||||
if j.hasKey("auth"):
|
if j.hasKey("auth"):
|
||||||
let s = j["auth"]
|
let s = j["auth"]
|
||||||
if s.hasKey("enabled"): cfg.authEnabled = s["enabled"].getBool()
|
if s.hasKey("enabled"): cfg.authEnabled = s["enabled"].getBool()
|
||||||
@@ -177,6 +183,13 @@ proc loadConfigFromEnv*(cfg: var BaraConfig) =
|
|||||||
cfg.tlsEnabled = parseEnvBool(getEnv("BARADB_TLS_ENABLED", ""), cfg.tlsEnabled)
|
cfg.tlsEnabled = parseEnvBool(getEnv("BARADB_TLS_ENABLED", ""), cfg.tlsEnabled)
|
||||||
cfg.certFile = getEnv("BARADB_CERT_FILE", cfg.certFile)
|
cfg.certFile = getEnv("BARADB_CERT_FILE", cfg.certFile)
|
||||||
cfg.keyFile = getEnv("BARADB_KEY_FILE", cfg.keyFile)
|
cfg.keyFile = getEnv("BARADB_KEY_FILE", cfg.keyFile)
|
||||||
|
cfg.tlsCaFile = getEnv("BARADB_TLS_CA_FILE", cfg.tlsCaFile)
|
||||||
|
let tlsVerifyEnv = getEnv("BARADB_TLS_VERIFY_PEER", "")
|
||||||
|
if tlsVerifyEnv.len > 0:
|
||||||
|
cfg.tlsVerifyPeer = parseEnvBool(tlsVerifyEnv, cfg.tlsVerifyPeer)
|
||||||
|
elif cfg.tlsCaFile.len > 0:
|
||||||
|
# CA present and verify flag unset → verify (fail-closed for MITM).
|
||||||
|
cfg.tlsVerifyPeer = true
|
||||||
cfg.idleTimeoutMs = parseEnvInt(getEnv("BARADB_IDLE_TIMEOUT_MS", ""), cfg.idleTimeoutMs)
|
cfg.idleTimeoutMs = parseEnvInt(getEnv("BARADB_IDLE_TIMEOUT_MS", ""), cfg.idleTimeoutMs)
|
||||||
cfg.queryTimeoutMs = parseEnvInt(getEnv("BARADB_QUERY_TIMEOUT_MS", ""), cfg.queryTimeoutMs)
|
cfg.queryTimeoutMs = parseEnvInt(getEnv("BARADB_QUERY_TIMEOUT_MS", ""), cfg.queryTimeoutMs)
|
||||||
cfg.slowQueryThresholdMs = parseEnvInt(getEnv("BARADB_SLOW_QUERY_THRESHOLD_MS", ""), cfg.slowQueryThresholdMs)
|
cfg.slowQueryThresholdMs = parseEnvInt(getEnv("BARADB_SLOW_QUERY_THRESHOLD_MS", ""), cfg.slowQueryThresholdMs)
|
||||||
@@ -233,7 +246,11 @@ proc loadConfigFromEnv*(cfg: var BaraConfig) =
|
|||||||
cfg.raftTlsCertFile = getEnv("BARADB_RAFT_TLS_CERT_FILE", cfg.raftTlsCertFile)
|
cfg.raftTlsCertFile = getEnv("BARADB_RAFT_TLS_CERT_FILE", cfg.raftTlsCertFile)
|
||||||
cfg.raftTlsKeyFile = getEnv("BARADB_RAFT_TLS_KEY_FILE", cfg.raftTlsKeyFile)
|
cfg.raftTlsKeyFile = getEnv("BARADB_RAFT_TLS_KEY_FILE", cfg.raftTlsKeyFile)
|
||||||
cfg.raftTlsCaFile = getEnv("BARADB_RAFT_TLS_CA_FILE", cfg.raftTlsCaFile)
|
cfg.raftTlsCaFile = getEnv("BARADB_RAFT_TLS_CA_FILE", cfg.raftTlsCaFile)
|
||||||
cfg.raftTlsVerifyPeer = parseEnvBool(getEnv("BARADB_RAFT_TLS_VERIFY_PEER", ""), cfg.raftTlsVerifyPeer)
|
let raftVerifyEnv = getEnv("BARADB_RAFT_TLS_VERIFY_PEER", "")
|
||||||
|
if raftVerifyEnv.len > 0:
|
||||||
|
cfg.raftTlsVerifyPeer = parseEnvBool(raftVerifyEnv, cfg.raftTlsVerifyPeer)
|
||||||
|
elif cfg.raftTlsCaFile.len > 0:
|
||||||
|
cfg.raftTlsVerifyPeer = true
|
||||||
# Optional: client (SQL) addresses for leader write forwarding.
|
# Optional: client (SQL) addresses for leader write forwarding.
|
||||||
# Same id@host:port shape as BARADB_RAFT_PEERS, but ports are BARADB_PORT values.
|
# Same id@host:port shape as BARADB_RAFT_PEERS, but ports are BARADB_PORT values.
|
||||||
let clientPeersEnv = getEnv("BARADB_RAFT_CLIENT_PEERS", "")
|
let clientPeersEnv = getEnv("BARADB_RAFT_CLIENT_PEERS", "")
|
||||||
@@ -292,6 +309,14 @@ proc validateProductionConfig*(cfg: BaraConfig) =
|
|||||||
if cfg.jwtSecret in ["change-me", "change-me-to-random-32-char-string", "secret", "default"]:
|
if cfg.jwtSecret in ["change-me", "change-me-to-random-32-char-string", "secret", "default"]:
|
||||||
raise newException(ValueError,
|
raise newException(ValueError,
|
||||||
"Production refuses insecure JWT secret placeholder. Set a strong BARADB_JWT_SECRET.")
|
"Production refuses insecure JWT secret placeholder. Set a strong BARADB_JWT_SECRET.")
|
||||||
|
if cfg.tlsEnabled:
|
||||||
|
if not cfg.tlsVerifyPeer or cfg.tlsCaFile.len == 0:
|
||||||
|
raise newException(ValueError,
|
||||||
|
"Production TLS requires peer verification. Set BARADB_TLS_VERIFY_PEER=true and BARADB_TLS_CA_FILE.")
|
||||||
|
if cfg.raftTlsEnabled:
|
||||||
|
if not cfg.raftTlsVerifyPeer or cfg.raftTlsCaFile.len == 0:
|
||||||
|
raise newException(ValueError,
|
||||||
|
"Production raft TLS requires peer verification. Set BARADB_RAFT_TLS_VERIFY_PEER=true and BARADB_RAFT_TLS_CA_FILE.")
|
||||||
|
|
||||||
proc getEffectiveJwtSecret*(cfg: BaraConfig): string =
|
proc getEffectiveJwtSecret*(cfg: BaraConfig): string =
|
||||||
if cfg.jwtSecret.len > 0:
|
if cfg.jwtSecret.len > 0:
|
||||||
|
|||||||
@@ -5,6 +5,8 @@ import std/monotimes
|
|||||||
import std/net
|
import std/net
|
||||||
import std/strutils
|
import std/strutils
|
||||||
import std/nativesockets
|
import std/nativesockets
|
||||||
|
when defined(posix):
|
||||||
|
import std/posix
|
||||||
|
|
||||||
type
|
type
|
||||||
DistTxnState* = enum
|
DistTxnState* = enum
|
||||||
@@ -89,6 +91,13 @@ proc connectWithTimeout(sock: Socket, host: string, port: Port, timeoutMs: int):
|
|||||||
var fds = @[sock.getFd]
|
var fds = @[sock.getFd]
|
||||||
if selectWrite(fds, timeoutMs) <= 0:
|
if selectWrite(fds, timeoutMs) <= 0:
|
||||||
return false
|
return false
|
||||||
|
when defined(posix):
|
||||||
|
# selectWrite reports a refused connect as writable; SO_ERROR tells the truth.
|
||||||
|
var err: cint = 0
|
||||||
|
var errLen = SockLen(sizeof(err))
|
||||||
|
discard posix.getsockopt(sock.getFd, 1'i32, 4'i32, addr err, addr errLen)
|
||||||
|
if err != 0:
|
||||||
|
return false
|
||||||
sock.getFd.setBlocking(true)
|
sock.getFd.setBlocking(true)
|
||||||
return true
|
return true
|
||||||
|
|
||||||
@@ -96,15 +105,18 @@ proc sendDistTxnRpc(host: string, port: int, txnId: uint64, action: string, time
|
|||||||
## Send 2PC RPC to participant node via TCP text protocol.
|
## Send 2PC RPC to participant node via TCP text protocol.
|
||||||
## Protocol: "DISTTXN <txnId> <action>\n" where action = PREPARE|COMMIT|ROLLBACK
|
## Protocol: "DISTTXN <txnId> <action>\n" where action = PREPARE|COMMIT|ROLLBACK
|
||||||
## Response: "OK\n" or "ERR <msg>\n"
|
## Response: "OK\n" or "ERR <msg>\n"
|
||||||
var sock = newSocket()
|
try:
|
||||||
defer: sock.close()
|
var sock = newSocket()
|
||||||
if not connectWithTimeout(sock, host, Port(port), timeoutMs):
|
defer: sock.close()
|
||||||
|
if not connectWithTimeout(sock, host, Port(port), timeoutMs):
|
||||||
|
return false
|
||||||
|
let msg = "DISTTXN " & $txnId & " " & action & "\n"
|
||||||
|
sock.send(msg)
|
||||||
|
var response = ""
|
||||||
|
sock.readLine(response)
|
||||||
|
return response.strip() == "OK"
|
||||||
|
except CatchableError:
|
||||||
return false
|
return false
|
||||||
let msg = "DISTTXN " & $txnId & " " & action & "\n"
|
|
||||||
sock.send(msg)
|
|
||||||
var response = ""
|
|
||||||
sock.readLine(response)
|
|
||||||
return response.strip() == "OK"
|
|
||||||
|
|
||||||
type
|
type
|
||||||
ParticipantInfo = object
|
ParticipantInfo = object
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ import ../storage/gate
|
|||||||
import ../core/mvcc
|
import ../core/mvcc
|
||||||
import ../protocol/wire
|
import ../protocol/wire
|
||||||
import ../core/websocket
|
import ../core/websocket
|
||||||
|
import ../query/exec/rls
|
||||||
import jwt as jwtlib
|
import jwt as jwtlib
|
||||||
import ../protocol/auth
|
import ../protocol/auth
|
||||||
import ../protocol/ratelimit
|
import ../protocol/ratelimit
|
||||||
@@ -55,6 +56,14 @@ proc newHttpServerWithRegistry*(config: BaraConfig, registry: DatabaseRegistry):
|
|||||||
ctx.txnManager = newTxnManager()
|
ctx.txnManager = newTxnManager()
|
||||||
let secret = config.getEffectiveJwtSecret()
|
let secret = config.getEffectiveJwtSecret()
|
||||||
let ws = newWsServer(config, secret)
|
let ws = newWsServer(config, secret)
|
||||||
|
block:
|
||||||
|
let wsRef {.cursor.} = ws
|
||||||
|
let ctxRef {.cursor.} = ctx
|
||||||
|
wsRef.canSubscribe = proc(username, table: string): bool {.gcsafe.} =
|
||||||
|
if username.len == 0:
|
||||||
|
return false
|
||||||
|
{.cast(gcsafe).}:
|
||||||
|
return hasPrivilegeFor(ctxRef, username, table, "SELECT")
|
||||||
let rl = newRateLimiter(rlaTokenBucket, config.rateLimitGlobal, config.rateLimitPerClient)
|
let rl = newRateLimiter(rlaTokenBucket, config.rateLimitGlobal, config.rateLimitPerClient)
|
||||||
ctx.onChange = proc(ev: ChangeEvent) =
|
ctx.onChange = proc(ev: ChangeEvent) =
|
||||||
let msg = $ev.kind & " " & ev.table
|
let msg = $ev.kind & " " & ev.table
|
||||||
|
|||||||
@@ -178,12 +178,16 @@ proc write*(tm: TxnManager, txn: Transaction, key: string, value: seq[byte]): bo
|
|||||||
return false
|
return false
|
||||||
|
|
||||||
# Timeout-based deadlock detection: abort stale transactions
|
# Timeout-based deadlock detection: abort stale transactions
|
||||||
|
# Collect then delete — never mutate activeTxns while iterating it.
|
||||||
let now = getMonoTime().ticks()
|
let now = getMonoTime().ticks()
|
||||||
|
var staleIds: seq[TxnId] = @[]
|
||||||
for otherId, otherTxn in tm.activeTxns:
|
for otherId, otherTxn in tm.activeTxns:
|
||||||
if otherId != txn.id and otherTxn.state == tsActive:
|
if otherId != txn.id and otherTxn.state == tsActive:
|
||||||
if now - otherTxn.startTime > tm.txnTimeoutMs * 1_000_000:
|
if now - otherTxn.startTime > tm.txnTimeoutMs * 1_000_000:
|
||||||
otherTxn.state = tsAborted
|
otherTxn.state = tsAborted
|
||||||
tm.activeTxns.del(otherId)
|
staleIds.add(otherId)
|
||||||
|
for id in staleIds:
|
||||||
|
tm.activeTxns.del(id)
|
||||||
|
|
||||||
# Check for write-write conflict against other active transactions' write sets
|
# Check for write-write conflict against other active transactions' write sets
|
||||||
for otherId, otherTxn in tm.activeTxns:
|
for otherId, otherTxn in tm.activeTxns:
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ import std/os
|
|||||||
import logging
|
import logging
|
||||||
import ../protocol/wire
|
import ../protocol/wire
|
||||||
import ../protocol/ssl
|
import ../protocol/ssl
|
||||||
|
import backup
|
||||||
|
|
||||||
type
|
type
|
||||||
RaftState* = enum
|
RaftState* = enum
|
||||||
@@ -95,8 +96,10 @@ type
|
|||||||
baseTerm: uint64): bool {.gcsafe.}
|
baseTerm: uint64): bool {.gcsafe.}
|
||||||
snapIncomingId*: uint64
|
snapIncomingId*: uint64
|
||||||
snapIncomingFile*: string
|
snapIncomingFile*: string
|
||||||
## Leader InstallSnapshot send. buildSnapshot archives the current data
|
## Leader InstallSnapshot send. buildSnapshot writes an UNCOMPRESSED tar
|
||||||
## dir into destPath (wired in baradadb.nim via backupDataDir).
|
## of the data dir to destPath (wired in baradadb.nim via tarDataDir, under
|
||||||
|
## the storage gate); sendSnapshot then compresses it off the event loop
|
||||||
|
## (gzipFileAsync) and streams the resulting .tar.gz to the follower.
|
||||||
## snapRejectStreak counts consecutive floor-level AppendEntries rejects
|
## snapRejectStreak counts consecutive floor-level AppendEntries rejects
|
||||||
## per peer; at 2 the peer is queued in snapPending and the network layer
|
## per peer; at 2 the peer is queued in snapPending and the network layer
|
||||||
## (processMessage) kicks off sendSnapshot. snapSending is the
|
## (processMessage) kicks off sendSnapshot. snapSending is the
|
||||||
@@ -646,8 +649,11 @@ proc handleAppendReply*(node: RaftNode, peerId: string, reply: RaftMessage) =
|
|||||||
node.snapRejectStreak.del(peerId)
|
node.snapRejectStreak.del(peerId)
|
||||||
node.snapPending.excl(peerId)
|
node.snapPending.excl(peerId)
|
||||||
|
|
||||||
# Update commit index using true majority calculation
|
# Update commit index using strict majority — the same form as the election
|
||||||
let majority = (node.peers.len + 1 + 1) div 2 # majority of cluster (peers + leader)
|
# check in handleVoteReply. Cluster size N = peers.len + 1; a strict
|
||||||
|
# majority is N div 2 + 1. The previous (N + 1) div 2 under-counted for
|
||||||
|
# even-sized clusters (e.g. N=4 committed at 2/4, a minority).
|
||||||
|
let majority = (node.peers.len + 1) div 2 + 1 # strict majority of cluster
|
||||||
var newCommitIdx = node.commitIndex
|
var newCommitIdx = node.commitIndex
|
||||||
|
|
||||||
# Walk logical indices high→low via findLogEntryByIndex (log may be compacted).
|
# Walk logical indices high→low via findLogEntryByIndex (log may be compacted).
|
||||||
@@ -1006,9 +1012,10 @@ proc sendSnapshot*(net: RaftNetwork, peerId: string) {.async.} =
|
|||||||
## when handleAppendReply queues the peer in snapPending after consecutive
|
## when handleAppendReply queues the peer in snapPending after consecutive
|
||||||
## floor-level rejects. Single-flight per peer via node.snapSending.
|
## floor-level rejects. Single-flight per peer via node.snapSending.
|
||||||
##
|
##
|
||||||
## Runs on the raft event loop; buildSnapshot performs blocking disk I/O
|
## Runs on the raft event loop. buildSnapshot performs the tar on the loop
|
||||||
## (tar+gzip). Snapshot sends are rare, so we accept the stall rather than
|
## under the storage gate (consistent capture); the CPU-heavy gzip then runs
|
||||||
## adding a worker round-trip (same trade-off as restoreSnapshot).
|
## on a worker thread via gzipFileAsync, awaited here, so heartbeats and the
|
||||||
|
## election timer keep firing during compression instead of stalling.
|
||||||
let node = net.node
|
let node = net.node
|
||||||
if peerId in node.snapSending:
|
if peerId in node.snapSending:
|
||||||
return
|
return
|
||||||
@@ -1026,15 +1033,25 @@ proc sendSnapshot*(net: RaftNetwork, peerId: string) {.async.} =
|
|||||||
|
|
||||||
let baseIndex = node.lastSnapshotIndex
|
let baseIndex = node.lastSnapshotIndex
|
||||||
let baseTerm = node.lastSnapshotTerm
|
let baseTerm = node.lastSnapshotTerm
|
||||||
let destPath = node.dataDir / ("snap_out_" & $snapId & ".tar.gz")
|
# buildSnapshot writes an uncompressed tar (under the storage gate, on this
|
||||||
|
# loop); gzipFileAsync then compresses it on a worker thread off the loop.
|
||||||
|
# The follower still receives a normal .tar.gz byte stream.
|
||||||
|
let tarPath = node.dataDir / ("snap_out_" & $snapId & ".tar")
|
||||||
|
let destPath = tarPath & ".gz"
|
||||||
defer:
|
defer:
|
||||||
|
if fileExists(tarPath):
|
||||||
|
removeFile(tarPath)
|
||||||
if fileExists(destPath):
|
if fileExists(destPath):
|
||||||
removeFile(destPath)
|
removeFile(destPath)
|
||||||
|
|
||||||
if not node.buildSnapshot(destPath):
|
if not node.buildSnapshot(tarPath):
|
||||||
warn("sendSnapshot: buildSnapshot failed; aborting snapshot send to " & peerId)
|
warn("sendSnapshot: buildSnapshot failed; aborting snapshot send to " & peerId)
|
||||||
return
|
return
|
||||||
|
|
||||||
|
if not await gzipFileAsync(tarPath, destPath):
|
||||||
|
warn("sendSnapshot: snapshot compression failed; aborting send to " & peerId)
|
||||||
|
return
|
||||||
|
|
||||||
var f: File
|
var f: File
|
||||||
if not open(f, destPath, fmRead):
|
if not open(f, destPath, fmRead):
|
||||||
warn("sendSnapshot: cannot open built archive " & destPath)
|
warn("sendSnapshot: cannot open built archive " & destPath)
|
||||||
|
|||||||
@@ -96,6 +96,41 @@ proc connectWithTimeout(sock: Socket, host: string, port: Port, timeoutMs: int):
|
|||||||
sock.getFd.setBlocking(true)
|
sock.getFd.setBlocking(true)
|
||||||
return err == 0
|
return err == 0
|
||||||
|
|
||||||
|
type
|
||||||
|
RepOp* = enum
|
||||||
|
ropPut
|
||||||
|
ropDelete
|
||||||
|
ropInvalid
|
||||||
|
|
||||||
|
proc encodeRepPayload*(deleted: bool, key: string, value: seq[byte]): seq[byte] =
|
||||||
|
## Tagged legacy-REP payload. A leading op tag makes put/delete explicit so
|
||||||
|
## an empty put value (PK-only rows store an empty LSM value) is never
|
||||||
|
## mistaken for a delete on the receiver:
|
||||||
|
## put -> 'P' & key & "\x00" & value (value may be empty)
|
||||||
|
## delete -> 'D' & key
|
||||||
|
## Mirrors the raft convention (explicit "put"/"delete" commands).
|
||||||
|
if deleted:
|
||||||
|
cast[seq[byte]]("D" & key)
|
||||||
|
else:
|
||||||
|
cast[seq[byte]]("P" & key & "\x00" & cast[string](value))
|
||||||
|
|
||||||
|
proc decodeRepPayload*(data: seq[byte]): tuple[op: RepOp, key: string, value: seq[byte]] =
|
||||||
|
## Inverse of encodeRepPayload. Returns ropInvalid for empty or untagged
|
||||||
|
## payloads rather than guessing the operation from the value length.
|
||||||
|
if data.len == 0:
|
||||||
|
return (ropInvalid, "", @[])
|
||||||
|
case char(data[0])
|
||||||
|
of 'P':
|
||||||
|
let body = data[1 ..< data.len]
|
||||||
|
let nullPos = find(body, byte(0))
|
||||||
|
if nullPos < 0:
|
||||||
|
return (ropInvalid, "", @[])
|
||||||
|
return (ropPut, cast[string](body[0 ..< nullPos]), body[nullPos + 1 ..< body.len])
|
||||||
|
of 'D':
|
||||||
|
return (ropDelete, cast[string](data[1 ..< data.len]), @[])
|
||||||
|
else:
|
||||||
|
return (ropInvalid, "", @[])
|
||||||
|
|
||||||
proc shipToReplica(replica: Replica, lsn: uint64, data: seq[byte]): bool =
|
proc shipToReplica(replica: Replica, lsn: uint64, data: seq[byte]): bool =
|
||||||
## Send replication data to a replica via TCP.
|
## Send replication data to a replica via TCP.
|
||||||
## Protocol: "REP <lsn> <dataLen>\n<data>"
|
## Protocol: "REP <lsn> <dataLen>\n<data>"
|
||||||
@@ -178,10 +213,18 @@ proc writeLsn*(rm: ReplicationManager, data: seq[byte]): uint64 =
|
|||||||
rm.pendingAcks[lsn].excl(id)
|
rm.pendingAcks[lsn].excl(id)
|
||||||
if rm.pendingAcks[lsn].len == 0:
|
if rm.pendingAcks[lsn].len == 0:
|
||||||
rm.pendingAcks.del(lsn)
|
rm.pendingAcks.del(lsn)
|
||||||
|
# Semi-sync requires at least syncReplicaCount acks when replicas are
|
||||||
|
# connected. With zero connected peers (nothing to ship) the write is
|
||||||
|
# local-only — same as sync mode with an empty replica set.
|
||||||
|
if rm.syncReplicaCount > 0 and replicasToShip.len > 0 and
|
||||||
|
ackCount < rm.syncReplicaCount:
|
||||||
|
# Drop the LSN from pendingAcks — write is not durable
|
||||||
|
rm.pendingAcks.del(lsn)
|
||||||
|
release(rm.lock)
|
||||||
|
echo "[ERROR] Semi-sync replication failed: only ", ackCount, "/",
|
||||||
|
rm.syncReplicaCount, " replicas acked for LSN ", lsn
|
||||||
|
return 0
|
||||||
release(rm.lock)
|
release(rm.lock)
|
||||||
if replicasToShip.len > 0 and ackCount == 0 and rm.syncReplicaCount > 0:
|
|
||||||
when defined(debug):
|
|
||||||
echo "Replication semi-sync: no replicas acked for LSN ", lsn
|
|
||||||
return lsn
|
return lsn
|
||||||
|
|
||||||
proc ackLsn*(rm: ReplicationManager, replicaId: string, lsn: uint64) =
|
proc ackLsn*(rm: ReplicationManager, replicaId: string, lsn: uint64) =
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
## BaraDB Server — async TCP server with wire protocol
|
## BaraDB Server — async TCP server with wire protocol
|
||||||
import std/asyncdispatch
|
import std/asyncdispatch
|
||||||
import std/asyncnet
|
import std/asyncnet
|
||||||
|
import std/os
|
||||||
import std/strutils
|
import std/strutils
|
||||||
import std/sequtils
|
import std/sequtils
|
||||||
import std/tables
|
import std/tables
|
||||||
@@ -22,6 +23,7 @@ import ../query/parser
|
|||||||
import ../query/ast
|
import ../query/ast
|
||||||
import ../query/executor
|
import ../query/executor
|
||||||
import ../query/exec/params
|
import ../query/exec/params
|
||||||
|
import ../query/exec/dml
|
||||||
import ../storage/lsm
|
import ../storage/lsm
|
||||||
import ../storage/gate
|
import ../storage/gate
|
||||||
import ../core/mvcc
|
import ../core/mvcc
|
||||||
@@ -49,6 +51,10 @@ type
|
|||||||
clusterMembership*: ClusterMembership
|
clusterMembership*: ClusterMembership
|
||||||
gossipProtocol*: GossipProtocol
|
gossipProtocol*: GossipProtocol
|
||||||
tls*: TLSContext
|
tls*: TLSContext
|
||||||
|
## Dedicated client-role TLS context for follower→leader forwarding.
|
||||||
|
## Must not reuse `tls` with verifyPeer — OpenSSL contexts are role-agnostic
|
||||||
|
## and enabling verify on the server context would break inbound handshakes.
|
||||||
|
tlsClient*: TLSContext
|
||||||
rateLimiter*: RateLimiter
|
rateLimiter*: RateLimiter
|
||||||
activeConnections*: int
|
activeConnections*: int
|
||||||
activeConnectionsLock*: Lock
|
activeConnectionsLock*: Lock
|
||||||
@@ -65,9 +71,22 @@ proc newServerWithRegistry*(config: BaraConfig, registry: DatabaseRegistry): Ser
|
|||||||
let ctx = cast[ExecutionContext](cast[pointer](dbInfo.ctx))
|
let ctx = cast[ExecutionContext](cast[pointer](dbInfo.ctx))
|
||||||
ctx.txnManager = newTxnManager()
|
ctx.txnManager = newTxnManager()
|
||||||
var tls: TLSContext = nil
|
var tls: TLSContext = nil
|
||||||
|
var tlsClient: TLSContext = nil
|
||||||
if config.tlsEnabled and config.certFile.len > 0 and config.keyFile.len > 0:
|
if config.tlsEnabled and config.certFile.len > 0 and config.keyFile.len > 0:
|
||||||
let tlsConfig = newTLSConfig(config.certFile, config.keyFile)
|
if config.tlsVerifyPeer and config.tlsCaFile.len == 0:
|
||||||
tls = newTLSContext(tlsConfig)
|
raise newException(ValueError,
|
||||||
|
"tlsVerifyPeer is true but CA file is missing. Set BARADB_TLS_CA_FILE")
|
||||||
|
if config.tlsVerifyPeer and config.tlsCaFile.len > 0 and
|
||||||
|
not fileExists(config.tlsCaFile):
|
||||||
|
raise newException(ValueError,
|
||||||
|
"BARADB_TLS_VERIFY_PEER=true but CA file missing: " & config.tlsCaFile)
|
||||||
|
tls = newTLSContext(newTLSConfig(config.certFile, config.keyFile))
|
||||||
|
if config.tlsVerifyPeer:
|
||||||
|
tlsClient = newTLSContext(newTLSConfig(
|
||||||
|
config.certFile, config.keyFile,
|
||||||
|
caFile = config.tlsCaFile, verifyPeer = true))
|
||||||
|
else:
|
||||||
|
tlsClient = tls
|
||||||
|
|
||||||
# Initialize sharding / gossip. Server fields own the refs; locals used inside
|
# Initialize sharding / gossip. Server fields own the refs; locals used inside
|
||||||
# callback closures are {.cursor.} so ARC does not form uncollectable cycles
|
# callback closures are {.cursor.} so ARC does not form uncollectable cycles
|
||||||
@@ -84,6 +103,7 @@ proc newServerWithRegistry*(config: BaraConfig, registry: DatabaseRegistry): Ser
|
|||||||
clusterMembership: nil,
|
clusterMembership: nil,
|
||||||
gossipProtocol: newGossipProtocol(localId, config.address, config.port, gossipPort = gossipPort),
|
gossipProtocol: newGossipProtocol(localId, config.address, config.port, gossipPort = gossipPort),
|
||||||
tls: tls,
|
tls: tls,
|
||||||
|
tlsClient: tlsClient,
|
||||||
rateLimiter: rl)
|
rateLimiter: rl)
|
||||||
result.clusterMembership = newClusterMembership(result.shardRouter, localId)
|
result.clusterMembership = newClusterMembership(result.shardRouter, localId)
|
||||||
initLock(result.activeConnectionsLock)
|
initLock(result.activeConnectionsLock)
|
||||||
@@ -165,6 +185,11 @@ proc parseHeader(data: string): (bool, MessageHeader) =
|
|||||||
return (false, MessageHeader())
|
return (false, MessageHeader())
|
||||||
let kind = cast[MsgKind](rawKind)
|
let kind = cast[MsgKind](rawKind)
|
||||||
let length = readUint32BE(data, 4)
|
let length = readUint32BE(data, 4)
|
||||||
|
# Reject oversized messages before any buffer allocation: recvExactWithTimeout
|
||||||
|
# pre-allocates `length` bytes before the auth check, so an unbounded uint32
|
||||||
|
# (up to ~4 GiB) is a pre-auth memory-exhaustion DoS. Cap at the wire max.
|
||||||
|
if length > uint32(MaxWireStringLen):
|
||||||
|
return (false, MessageHeader())
|
||||||
let requestId = readUint32BE(data, 8)
|
let requestId = readUint32BE(data, 8)
|
||||||
return (true, MessageHeader(kind: kind, length: length, requestId: requestId))
|
return (true, MessageHeader(kind: kind, length: length, requestId: requestId))
|
||||||
|
|
||||||
@@ -221,12 +246,12 @@ proc forwardQueryToLeader*(host: string, port: int, query: string,
|
|||||||
timeoutMs: int = 5000): Future[(bool, QueryResult, string)] {.async.} =
|
timeoutMs: int = 5000): Future[(bool, QueryResult, string)] {.async.} =
|
||||||
## Proxy a write/DDL to the known leader's SQL port. Used by followers when
|
## Proxy a write/DDL to the known leader's SQL port. Used by followers when
|
||||||
## BARADB_RAFT_CLIENT_PEERS maps leader id → host:clientPort.
|
## BARADB_RAFT_CLIENT_PEERS maps leader id → host:clientPort.
|
||||||
## `tls` is the local server's client-port TLS context: when the wire port
|
## `tls` is a *client-role* context (Server.tlsClient). When the wire port
|
||||||
## serves TLS, the leader's does too, so the forwarding dial must complete a
|
## serves TLS, the leader's does too, so the forwarding dial must complete a
|
||||||
## client handshake. The context is reused as-is (verifyMode stays
|
## client handshake. Peer verification is honoured when that context was
|
||||||
## CVerifyNone — do NOT enable verifyPeer on the reused context); OpenSSL
|
## built with verifyPeer (BARADB_TLS_VERIFY_PEER + BARADB_TLS_CA_FILE).
|
||||||
## contexts are role-agnostic in Nim's stdlib, wrapConnectedSocket with
|
## Do NOT pass Server.tls (the inbound/server context) with verifyPeer
|
||||||
## handshakeAsClient sets the role.
|
## flipped on — OpenSSL contexts are role-agnostic in Nim's stdlib.
|
||||||
var sock: AsyncSocket = nil
|
var sock: AsyncSocket = nil
|
||||||
try:
|
try:
|
||||||
sock = newAsyncSocket()
|
sock = newAsyncSocket()
|
||||||
@@ -413,14 +438,12 @@ proc executeQuery(db: LSMTree, ctx: ExecutionContext, query: string, params: seq
|
|||||||
# the raft path below handles the statement).
|
# the raft path below handles the statement).
|
||||||
if raftNode == nil and replication != nil and res.keyValuePairs.len > 0:
|
if raftNode == nil and replication != nil and res.keyValuePairs.len > 0:
|
||||||
for pair in res.keyValuePairs:
|
for pair in res.keyValuePairs:
|
||||||
# Legacy REP wire format: key \x00 value, empty value = delete
|
# Legacy REP wire format: explicit 'P'/'D' op tag (see
|
||||||
# on the receiver. Deletes ship an empty value as before.
|
# encodeRepPayload). The tag — not an empty value — distinguishes
|
||||||
let value = if pair.deleted: @[] else: pair.value
|
# a put from a delete, so PK-only rows (empty value) replicate as
|
||||||
var data = newSeq[byte](pair.key.len + 1 + value.len)
|
# puts instead of vanishing as deletes.
|
||||||
for i, c in pair.key: data[i] = byte(c)
|
discard replication.writeLsn(
|
||||||
data[pair.key.len] = byte(0)
|
encodeRepPayload(pair.deleted, pair.key, pair.value))
|
||||||
for i, c in value: data[pair.key.len + 1 + i] = c
|
|
||||||
discard replication.writeLsn(data)
|
|
||||||
qr = QueryResult(affectedRows: res.affectedRows, rowCount: res.rows.len)
|
qr = QueryResult(affectedRows: res.affectedRows, rowCount: res.rows.len)
|
||||||
qr.columns = res.columns
|
qr.columns = res.columns
|
||||||
|
|
||||||
@@ -655,14 +678,25 @@ proc handleClient(server: Server, client: AsyncSocket, clientId: int) {.async.}
|
|||||||
if chunk.len == 0: break
|
if chunk.len == 0: break
|
||||||
data.add(chunk)
|
data.add(chunk)
|
||||||
if data.len > 0:
|
if data.len > 0:
|
||||||
let nullPos = data.find('\0')
|
# Op tag — not value length — decides put vs delete, so a PK-only
|
||||||
if nullPos >= 0:
|
# put (empty value) is applied as a put and the row survives.
|
||||||
let key = data[0..<nullPos]
|
let decoded = decodeRepPayload(cast[seq[byte]](data))
|
||||||
let value = data[nullPos+1..^1]
|
case decoded.op
|
||||||
if value.len > 0:
|
of ropPut, ropDelete:
|
||||||
server.db.put(key, stringToBytes(value))
|
# Apply through applyReplicatedPut/Delete (not raw db.put/delete)
|
||||||
else:
|
# so secondary B-tree/FTS/HNSW/graph indexes stay consistent on
|
||||||
server.db.delete(key)
|
# the replica — the same path raft uses. server.ctx is the
|
||||||
|
# canonical default ctx whose index structures the per-connection
|
||||||
|
# query clones share. Under the storage gate: those structures
|
||||||
|
# are shared with hunos HTTP workers and are only safe to mutate
|
||||||
|
# under it.
|
||||||
|
withStorageGate:
|
||||||
|
if decoded.op == ropPut:
|
||||||
|
applyReplicatedPut(server.ctx, decoded.key, decoded.value)
|
||||||
|
else:
|
||||||
|
applyReplicatedDelete(server.ctx, decoded.key)
|
||||||
|
of ropInvalid:
|
||||||
|
discard
|
||||||
await client.send("ACK " & $lsn & "\n")
|
await client.send("ACK " & $lsn & "\n")
|
||||||
else:
|
else:
|
||||||
await client.send("ERR\n")
|
await client.send("ERR\n")
|
||||||
@@ -670,6 +704,9 @@ proc handleClient(server: Server, client: AsyncSocket, clientId: int) {.async.}
|
|||||||
|
|
||||||
# Detect shard migration data (starts with "MIGRATE ")
|
# Detect shard migration data (starts with "MIGRATE ")
|
||||||
if headerData.len >= 8 and headerData[0..7] == "MIGRATE ":
|
if headerData.len >= 8 and headerData[0..7] == "MIGRATE ":
|
||||||
|
if not authenticated:
|
||||||
|
await client.send("ERR auth required\n")
|
||||||
|
continue
|
||||||
var rest = headerData[8..^1]
|
var rest = headerData[8..^1]
|
||||||
while '\n' notin rest:
|
while '\n' notin rest:
|
||||||
let more = await client.recvWithTimeout(1024, idleTimeout)
|
let more = await client.recvWithTimeout(1024, idleTimeout)
|
||||||
@@ -788,7 +825,7 @@ proc handleClient(server: Server, client: AsyncSocket, clientId: int) {.async.}
|
|||||||
replication=server.replicationManager, raftNode=server.raftNode,
|
replication=server.replicationManager, raftNode=server.raftNode,
|
||||||
raftWriteTimeoutMs=server.config.raftWriteTimeoutMs,
|
raftWriteTimeoutMs=server.config.raftWriteTimeoutMs,
|
||||||
raftPeerClientAddrs=server.config.raftPeerClientAddrs,
|
raftPeerClientAddrs=server.config.raftPeerClientAddrs,
|
||||||
forwardTls=server.tls)
|
forwardTls=server.tlsClient)
|
||||||
let durationMs = int((getMonoTime().ticks() - startTicks) div 1_000_000)
|
let durationMs = int((getMonoTime().ticks() - startTicks) div 1_000_000)
|
||||||
|
|
||||||
if durationMs >= slowThreshold:
|
if durationMs >= slowThreshold:
|
||||||
@@ -812,7 +849,7 @@ proc handleClient(server: Server, client: AsyncSocket, clientId: int) {.async.}
|
|||||||
replication=server.replicationManager, raftNode=server.raftNode,
|
replication=server.replicationManager, raftNode=server.raftNode,
|
||||||
raftWriteTimeoutMs=server.config.raftWriteTimeoutMs,
|
raftWriteTimeoutMs=server.config.raftWriteTimeoutMs,
|
||||||
raftPeerClientAddrs=server.config.raftPeerClientAddrs,
|
raftPeerClientAddrs=server.config.raftPeerClientAddrs,
|
||||||
forwardTls=server.tls)
|
forwardTls=server.tlsClient)
|
||||||
let durationMs = int((getMonoTime().ticks() - startTicks) div 1_000_000)
|
let durationMs = int((getMonoTime().ticks() - startTicks) div 1_000_000)
|
||||||
|
|
||||||
if durationMs >= slowThreshold:
|
if durationMs >= slowThreshold:
|
||||||
|
|||||||
@@ -15,18 +15,25 @@ else:
|
|||||||
import config
|
import config
|
||||||
import jwt as jwtlib
|
import jwt as jwtlib
|
||||||
|
|
||||||
|
const
|
||||||
|
## RFC 6455: reject oversized frames/messages to bound memory.
|
||||||
|
MaxWsFrameBytes* = 1 * 1024 * 1024
|
||||||
|
MaxWsMessageBytes* = 4 * 1024 * 1024
|
||||||
|
MaxWsControlPayload* = 125
|
||||||
|
|
||||||
type
|
type
|
||||||
WsFrame = object
|
WsFrame* = object
|
||||||
fin: bool
|
fin*: bool
|
||||||
opcode: uint8
|
opcode*: uint8
|
||||||
masked: bool
|
masked*: bool
|
||||||
payloadLen: uint64
|
payloadLen*: uint64
|
||||||
maskKey: array[4, byte]
|
maskKey*: array[4, byte]
|
||||||
payload: string
|
payload*: string
|
||||||
|
|
||||||
WsClient* = ref object
|
WsClient* = ref object
|
||||||
socket: AsyncSocket
|
socket: AsyncSocket
|
||||||
id: int
|
id: int
|
||||||
|
username: string
|
||||||
subscriptions: HashSet[string]
|
subscriptions: HashSet[string]
|
||||||
|
|
||||||
WsServer* = ref object
|
WsServer* = ref object
|
||||||
@@ -35,6 +42,8 @@ type
|
|||||||
running: bool
|
running: bool
|
||||||
config*: BaraConfig
|
config*: BaraConfig
|
||||||
secretKey*: string
|
secretKey*: string
|
||||||
|
## Table-level read authorization for SUBSCRIBE. Nil + authEnabled → deny.
|
||||||
|
canSubscribe*: proc(username, table: string): bool {.closure, gcsafe.}
|
||||||
onInsert*: proc (table, key, value: string) {.closure.}
|
onInsert*: proc (table, key, value: string) {.closure.}
|
||||||
onDelete*: proc (table, key: string) {.closure.}
|
onDelete*: proc (table, key: string) {.closure.}
|
||||||
|
|
||||||
@@ -46,20 +55,19 @@ proc newWsServer*(cfg: BaraConfig = defaultConfig(), secret: string = ""): WsSer
|
|||||||
# WebSocket frame encoding/decoding (RFC 6455)
|
# WebSocket frame encoding/decoding (RFC 6455)
|
||||||
# ----------------------------------------------------------------------
|
# ----------------------------------------------------------------------
|
||||||
|
|
||||||
proc encodeFrame(opcode: uint8, payload: string): string =
|
proc encodeFrame*(opcode: uint8, payload: string, masked = false,
|
||||||
|
maskKey: array[4, byte] = [0'u8, 0, 0, 0]): string =
|
||||||
result = ""
|
result = ""
|
||||||
let isMasked = false
|
|
||||||
var b0 = 0x80'u8 or opcode
|
var b0 = 0x80'u8 or opcode
|
||||||
result.add(char(b0))
|
result.add(char(b0))
|
||||||
|
|
||||||
var b1 = 0'u8
|
var b1 = if masked: 0x80'u8 else: 0'u8
|
||||||
if not isMasked:
|
if payload.len < 126:
|
||||||
if payload.len < 126:
|
b1 = b1 or uint8(payload.len)
|
||||||
b1 = uint8(payload.len)
|
elif payload.len <= 65535:
|
||||||
elif payload.len <= 65535:
|
b1 = b1 or 126
|
||||||
b1 = 126
|
else:
|
||||||
else:
|
b1 = b1 or 127
|
||||||
b1 = 127
|
|
||||||
result.add(char(b1))
|
result.add(char(b1))
|
||||||
|
|
||||||
if payload.len >= 126 and payload.len <= 65535:
|
if payload.len >= 126 and payload.len <= 65535:
|
||||||
@@ -71,9 +79,17 @@ proc encodeFrame(opcode: uint8, payload: string): string =
|
|||||||
for i in countdown(7, 0):
|
for i in countdown(7, 0):
|
||||||
result.add(char((len64 shr (i * 8)) and 0xFF))
|
result.add(char((len64 shr (i * 8)) and 0xFF))
|
||||||
|
|
||||||
result.add(payload)
|
if masked:
|
||||||
|
for i in 0..3:
|
||||||
|
result.add(char(maskKey[i]))
|
||||||
|
for i, c in payload:
|
||||||
|
result.add(char(byte(c) xor maskKey[i mod 4]))
|
||||||
|
else:
|
||||||
|
result.add(payload)
|
||||||
|
|
||||||
proc decodeFrame(data: string): (WsFrame, int) =
|
proc decodeFrame*(data: string): (WsFrame, int) =
|
||||||
|
## Returns (frame, consumed). consumed == 0 → need more bytes;
|
||||||
|
## consumed < 0 → protocol error (close the connection).
|
||||||
if data.len < 2:
|
if data.len < 2:
|
||||||
return (WsFrame(), 0)
|
return (WsFrame(), 0)
|
||||||
|
|
||||||
@@ -84,6 +100,10 @@ proc decodeFrame(data: string): (WsFrame, int) =
|
|||||||
frame.opcode = b0 and 0x0F
|
frame.opcode = b0 and 0x0F
|
||||||
frame.masked = (b1 and 0x80) != 0
|
frame.masked = (b1 and 0x80) != 0
|
||||||
|
|
||||||
|
# RFC 6455 §5.1 — client-to-server frames MUST be masked.
|
||||||
|
if not frame.masked:
|
||||||
|
return (WsFrame(), -1)
|
||||||
|
|
||||||
var len = uint64(b1 and 0x7F)
|
var len = uint64(b1 and 0x7F)
|
||||||
var offset = 2
|
var offset = 2
|
||||||
|
|
||||||
@@ -98,23 +118,24 @@ proc decodeFrame(data: string): (WsFrame, int) =
|
|||||||
len = (len shl 8) or uint64(uint8(data[2 + i]))
|
len = (len shl 8) or uint64(uint8(data[2 + i]))
|
||||||
offset = 10
|
offset = 10
|
||||||
|
|
||||||
if frame.masked:
|
let isControl = frame.opcode == 0x8 or frame.opcode == 0x9 or frame.opcode == 0xA
|
||||||
if data.len < offset + 4: return (WsFrame(), 0)
|
if isControl and (not frame.fin or len > uint64(MaxWsControlPayload)):
|
||||||
for i in 0..3:
|
return (WsFrame(), -1)
|
||||||
frame.maskKey[i] = byte(data[offset + i])
|
if len > uint64(MaxWsFrameBytes):
|
||||||
offset += 4
|
return (WsFrame(), -1)
|
||||||
|
|
||||||
|
if data.len < offset + 4: return (WsFrame(), 0)
|
||||||
|
for i in 0..3:
|
||||||
|
frame.maskKey[i] = byte(data[offset + i])
|
||||||
|
offset += 4
|
||||||
|
|
||||||
if uint64(data.len) < uint64(offset) + len:
|
if uint64(data.len) < uint64(offset) + len:
|
||||||
return (Wsframe(), 0)
|
return (WsFrame(), 0)
|
||||||
|
|
||||||
if len > uint64(high(int) - 1):
|
|
||||||
return (Wsframe(), 0)
|
|
||||||
let plen = int(len)
|
let plen = int(len)
|
||||||
if frame.masked:
|
frame.payloadLen = len
|
||||||
for i in 0..<plen:
|
for i in 0..<plen:
|
||||||
frame.payload.add(char(byte(data[offset + i]) xor frame.maskKey[i mod 4]))
|
frame.payload.add(char(byte(data[offset + i]) xor frame.maskKey[i mod 4]))
|
||||||
else:
|
|
||||||
frame.payload = data[offset..offset + plen - 1]
|
|
||||||
|
|
||||||
return (frame, offset + plen)
|
return (frame, offset + plen)
|
||||||
|
|
||||||
@@ -179,6 +200,14 @@ proc computeAcceptKey(key: string): string =
|
|||||||
# Subscription management
|
# Subscription management
|
||||||
# ----------------------------------------------------------------------
|
# ----------------------------------------------------------------------
|
||||||
|
|
||||||
|
proc validSubscribeTable(table: string): bool =
|
||||||
|
if table.len == 0 or table.len > 128:
|
||||||
|
return false
|
||||||
|
for c in table:
|
||||||
|
if c notin {'a'..'z', 'A'..'Z', '0'..'9', '_'}:
|
||||||
|
return false
|
||||||
|
return true
|
||||||
|
|
||||||
proc subscribe*(client: WsClient, table: string) =
|
proc subscribe*(client: WsClient, table: string) =
|
||||||
client.subscriptions.incl(table)
|
client.subscriptions.incl(table)
|
||||||
|
|
||||||
@@ -201,9 +230,11 @@ proc broadcastToTable*(server: WsServer, table: string, msg: string) {.async.} =
|
|||||||
# WebSocket client handler
|
# WebSocket client handler
|
||||||
# ----------------------------------------------------------------------
|
# ----------------------------------------------------------------------
|
||||||
|
|
||||||
proc handleWsClient(server: WsServer, client: AsyncSocket, id: int) {.async.} =
|
proc handleWsClient(server: WsServer, client: AsyncSocket, id: int,
|
||||||
|
username: string = "") {.async.} =
|
||||||
echo "WebSocket client ", id, " connected"
|
echo "WebSocket client ", id, " connected"
|
||||||
var wsClient = WsClient(socket: client, id: id, subscriptions: initHashSet[string]())
|
var wsClient = WsClient(socket: client, id: id, username: username,
|
||||||
|
subscriptions: initHashSet[string]())
|
||||||
server.clients[id] = wsClient
|
server.clients[id] = wsClient
|
||||||
|
|
||||||
var buf = ""
|
var buf = ""
|
||||||
@@ -212,12 +243,22 @@ proc handleWsClient(server: WsServer, client: AsyncSocket, id: int) {.async.} =
|
|||||||
let chunk = await client.recv(4096)
|
let chunk = await client.recv(4096)
|
||||||
if chunk.len == 0:
|
if chunk.len == 0:
|
||||||
break
|
break
|
||||||
|
if buf.len + chunk.len > MaxWsMessageBytes:
|
||||||
|
let closeF = encodeFrame(0x8, "")
|
||||||
|
try: await client.send(closeF) except CatchableError: discard
|
||||||
|
break
|
||||||
buf.add(chunk)
|
buf.add(chunk)
|
||||||
|
|
||||||
while buf.len >= 2:
|
while buf.len >= 2:
|
||||||
let (frame, consumed) = decodeFrame(buf)
|
let (frame, consumed) = decodeFrame(buf)
|
||||||
if consumed == 0:
|
if consumed == 0:
|
||||||
break
|
break
|
||||||
|
if consumed < 0:
|
||||||
|
let closeF = encodeFrame(0x8, "")
|
||||||
|
try: await client.send(closeF) except CatchableError: discard
|
||||||
|
client.close()
|
||||||
|
server.clients.del(id)
|
||||||
|
return
|
||||||
|
|
||||||
case frame.opcode
|
case frame.opcode
|
||||||
of 0x8: # close
|
of 0x8: # close
|
||||||
@@ -231,9 +272,18 @@ proc handleWsClient(server: WsServer, client: AsyncSocket, id: int) {.async.} =
|
|||||||
let msg = frame.payload
|
let msg = frame.payload
|
||||||
if msg.startsWith("SUBSCRIBE "):
|
if msg.startsWith("SUBSCRIBE "):
|
||||||
let table = msg[10..^1].strip()
|
let table = msg[10..^1].strip()
|
||||||
wsClient.subscribe(table)
|
var allowed = validSubscribeTable(table)
|
||||||
let ack = encodeFrame(0x1, "OK subscribed to " & table)
|
if allowed and server.config.authEnabled:
|
||||||
await client.send(ack)
|
if username.len == 0 or server.canSubscribe == nil or
|
||||||
|
not server.canSubscribe(username, table):
|
||||||
|
allowed = false
|
||||||
|
if allowed:
|
||||||
|
wsClient.subscribe(table)
|
||||||
|
let ack = encodeFrame(0x1, "OK subscribed to " & table)
|
||||||
|
await client.send(ack)
|
||||||
|
else:
|
||||||
|
let nack = encodeFrame(0x1, "ERR subscribe denied for " & table)
|
||||||
|
await client.send(nack)
|
||||||
elif msg.startsWith("UNSUBSCRIBE "):
|
elif msg.startsWith("UNSUBSCRIBE "):
|
||||||
let table = msg[12..^1].strip()
|
let table = msg[12..^1].strip()
|
||||||
wsClient.unsubscribe(table)
|
wsClient.unsubscribe(table)
|
||||||
@@ -285,6 +335,7 @@ proc handleConnection(server: WsServer, client: AsyncSocket) {.async.} =
|
|||||||
return
|
return
|
||||||
|
|
||||||
# Auth check
|
# Auth check
|
||||||
|
var username = ""
|
||||||
if server.config.authEnabled:
|
if server.config.authEnabled:
|
||||||
let authHeader = headers.getOrDefault("authorization", "")
|
let authHeader = headers.getOrDefault("authorization", "")
|
||||||
if authHeader.len == 0 or not authHeader.startsWith("Bearer "):
|
if authHeader.len == 0 or not authHeader.startsWith("Bearer "):
|
||||||
@@ -305,6 +356,8 @@ proc handleConnection(server: WsServer, client: AsyncSocket) {.async.} =
|
|||||||
await client.send("HTTP/1.1 401 Unauthorized\r\n\r\n")
|
await client.send("HTTP/1.1 401 Unauthorized\r\n\r\n")
|
||||||
client.close()
|
client.close()
|
||||||
return
|
return
|
||||||
|
if "sub" in token.claims:
|
||||||
|
username = token.claims["sub"].node.str
|
||||||
except CatchableError:
|
except CatchableError:
|
||||||
await client.send("HTTP/1.1 401 Unauthorized\r\n\r\n")
|
await client.send("HTTP/1.1 401 Unauthorized\r\n\r\n")
|
||||||
client.close()
|
client.close()
|
||||||
@@ -321,7 +374,7 @@ proc handleConnection(server: WsServer, client: AsyncSocket) {.async.} =
|
|||||||
await client.send(response)
|
await client.send(response)
|
||||||
|
|
||||||
inc server.nextId
|
inc server.nextId
|
||||||
asyncCheck server.handleWsClient(client, server.nextId)
|
asyncCheck server.handleWsClient(client, server.nextId, username)
|
||||||
|
|
||||||
proc setTcpNoDelay(sock: AsyncSocket) =
|
proc setTcpNoDelay(sock: AsyncSocket) =
|
||||||
## Enable TCP_NODELAY using the correct protocol level (IPPROTO_TCP).
|
## Enable TCP_NODELAY using the correct protocol level (IPPROTO_TCP).
|
||||||
|
|||||||
@@ -221,10 +221,17 @@ proc registerScramUser*(am: AuthManager, username, password: string,
|
|||||||
let cred = createScramCredential(password, iterationCount = iterationCount)
|
let cred = createScramCredential(password, iterationCount = iterationCount)
|
||||||
am.scramUsers[username] = cred
|
am.scramUsers[username] = cred
|
||||||
|
|
||||||
|
proc dummyScramStartWork() =
|
||||||
|
## Match known-user startScram work (urandom nonce + base64) so unknown
|
||||||
|
## users cannot be enumerated by timing.
|
||||||
|
discard generateNonce()
|
||||||
|
discard encode("0123456789abcdef0123456789abcdef")
|
||||||
|
|
||||||
proc startScram*(am: AuthManager, clientFirstMessage: string): string =
|
proc startScram*(am: AuthManager, clientFirstMessage: string): string =
|
||||||
## Start SCRAM authentication. Returns server-first-message.
|
## Start SCRAM authentication. Returns server-first-message.
|
||||||
let (_, username, clientNonce) = parseClientFirst(clientFirstMessage)
|
let (gs2, username, clientNonce) = parseClientFirst(clientFirstMessage)
|
||||||
if username notin am.scramUsers:
|
if username notin am.scramUsers or gs2 notin ["n", "y"]:
|
||||||
|
dummyScramStartWork()
|
||||||
raise newException(ValueError, "Authentication failed")
|
raise newException(ValueError, "Authentication failed")
|
||||||
|
|
||||||
let cred = am.scramUsers[username]
|
let cred = am.scramUsers[username]
|
||||||
@@ -239,6 +246,7 @@ proc startScram*(am: AuthManager, clientFirstMessage: string): string =
|
|||||||
|
|
||||||
var state = ScramServerState(
|
var state = ScramServerState(
|
||||||
username: username,
|
username: username,
|
||||||
|
gs2Flag: gs2,
|
||||||
clientFirstMessageBare: clientFirstMessageBare,
|
clientFirstMessageBare: clientFirstMessageBare,
|
||||||
serverFirstMessage: serverFirst,
|
serverFirstMessage: serverFirst,
|
||||||
authMessage: authMessage,
|
authMessage: authMessage,
|
||||||
@@ -264,6 +272,9 @@ proc finishScram*(am: AuthManager, clientFinalMessage: string): (bool, string) =
|
|||||||
var state = am.scramSessions[nonce]
|
var state = am.scramSessions[nonce]
|
||||||
am.scramSessions.del(nonce)
|
am.scramSessions.del(nonce)
|
||||||
|
|
||||||
|
if stripB64Padding(cbind) != expectedChannelBinding(state.gs2Flag):
|
||||||
|
return (false, "e=channel-bindings-dont-match")
|
||||||
|
|
||||||
# Update authMessage with client-final-message-without-proof
|
# Update authMessage with client-final-message-without-proof
|
||||||
let clientFinalWithoutProof = "c=" & cbind & ",r=" & nonce
|
let clientFinalWithoutProof = "c=" & cbind & ",r=" & nonce
|
||||||
state.authMessage = state.authMessage & "," & clientFinalWithoutProof
|
state.authMessage = state.authMessage & "," & clientFinalWithoutProof
|
||||||
|
|||||||
Binary file not shown.
@@ -19,6 +19,7 @@ type
|
|||||||
|
|
||||||
ScramServerState* = object
|
ScramServerState* = object
|
||||||
username*: string
|
username*: string
|
||||||
|
gs2Flag*: string
|
||||||
clientFirstMessageBare*: string
|
clientFirstMessageBare*: string
|
||||||
serverFirstMessage*: string
|
serverFirstMessage*: string
|
||||||
authMessage*: string
|
authMessage*: string
|
||||||
@@ -189,6 +190,19 @@ proc createScramCredential*(password: string, salt: string = "",
|
|||||||
# SCRAM message parsing / building
|
# SCRAM message parsing / building
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
proc expectedChannelBinding*(gs2Flag: string): string =
|
||||||
|
## RFC 5802 cbind-input with no cbind-data is the gs2-header: flag + "," + authzid + ",".
|
||||||
|
## Authzid is unsupported, so the header is `n,,` or `y,,`.
|
||||||
|
let header = gs2Flag & ",,"
|
||||||
|
result = encode(header)
|
||||||
|
while result.endsWith("="):
|
||||||
|
result.setLen(result.len - 1)
|
||||||
|
|
||||||
|
proc stripB64Padding*(s: string): string =
|
||||||
|
result = s
|
||||||
|
while result.endsWith("="):
|
||||||
|
result.setLen(result.len - 1)
|
||||||
|
|
||||||
proc parseClientFirst*(msg: string): (string, string, string) =
|
proc parseClientFirst*(msg: string): (string, string, string) =
|
||||||
## Parse client-first-message: gs2-header,username,nonce
|
## Parse client-first-message: gs2-header,username,nonce
|
||||||
## Returns: (gs2_header, username, nonce)
|
## Returns: (gs2_header, username, nonce)
|
||||||
|
|||||||
@@ -429,6 +429,8 @@ proc evalExprOld*(expr: IRExpr, row: Table[string, string], ctx: ExecutionContex
|
|||||||
let right = evalExprOld(expr.binRight, row, ctx)
|
let right = evalExprOld(expr.binRight, row, ctx)
|
||||||
case expr.binOp
|
case expr.binOp
|
||||||
of irEq:
|
of irEq:
|
||||||
|
# SQL three-valued logic: any NULL operand → unknown, not true.
|
||||||
|
if isNull(left) or isNull(right): return "\\N"
|
||||||
if left == right: return "true"
|
if left == right: return "true"
|
||||||
# Try numeric comparison
|
# Try numeric comparison
|
||||||
try:
|
try:
|
||||||
@@ -436,32 +438,45 @@ proc evalExprOld*(expr: IRExpr, row: Table[string, string], ctx: ExecutionContex
|
|||||||
except CatchableError: discard
|
except CatchableError: discard
|
||||||
return "false"
|
return "false"
|
||||||
of irNeq:
|
of irNeq:
|
||||||
if left != right: return "true"
|
if isNull(left) or isNull(right): return "\\N"
|
||||||
# Try numeric comparison
|
# Numeric-first so `!=` is the exact complement of `=` (irEq): string
|
||||||
|
# inequality alone would make `5 != 5.0` true while `5 = 5.0` is true.
|
||||||
try:
|
try:
|
||||||
return if parseFloat(left) != parseFloat(right): "true" else: "false"
|
return if parseFloat(left) != parseFloat(right): "true" else: "false"
|
||||||
except CatchableError: return "false"
|
except CatchableError:
|
||||||
|
return if left != right: "true" else: "false"
|
||||||
of irLt:
|
of irLt:
|
||||||
|
if isNull(left) or isNull(right): return "\\N"
|
||||||
try:
|
try:
|
||||||
return if parseFloat(left) < parseFloat(right): "true" else: "false"
|
return if parseFloat(left) < parseFloat(right): "true" else: "false"
|
||||||
except CatchableError: return if left < right: "true" else: "false"
|
except CatchableError: return if left < right: "true" else: "false"
|
||||||
of irLte:
|
of irLte:
|
||||||
|
if isNull(left) or isNull(right): return "\\N"
|
||||||
try:
|
try:
|
||||||
return if parseFloat(left) <= parseFloat(right): "true" else: "false"
|
return if parseFloat(left) <= parseFloat(right): "true" else: "false"
|
||||||
except CatchableError: return if left <= right: "true" else: "false"
|
except CatchableError: return if left <= right: "true" else: "false"
|
||||||
of irGt:
|
of irGt:
|
||||||
|
if isNull(left) or isNull(right): return "\\N"
|
||||||
try:
|
try:
|
||||||
return if parseFloat(left) > parseFloat(right): "true" else: "false"
|
return if parseFloat(left) > parseFloat(right): "true" else: "false"
|
||||||
except CatchableError: return if left > right: "true" else: "false"
|
except CatchableError: return if left > right: "true" else: "false"
|
||||||
of irGte:
|
of irGte:
|
||||||
|
if isNull(left) or isNull(right): return "\\N"
|
||||||
try:
|
try:
|
||||||
return if parseFloat(left) >= parseFloat(right): "true" else: "false"
|
return if parseFloat(left) >= parseFloat(right): "true" else: "false"
|
||||||
except CatchableError: return if left >= right: "true" else: "false"
|
except CatchableError: return if left >= right: "true" else: "false"
|
||||||
of irAnd:
|
of irAnd:
|
||||||
if left == "true" and right == "true": return "true"
|
# false AND x = false; unknown AND true/unknown = unknown; else both true.
|
||||||
return "false"
|
let lNull = isNull(left)
|
||||||
|
let rNull = isNull(right)
|
||||||
|
let lTrue = left == "true"
|
||||||
|
let rTrue = right == "true"
|
||||||
|
if (not lNull and not lTrue) or (not rNull and not rTrue): return "false"
|
||||||
|
if lNull or rNull: return "\\N"
|
||||||
|
return "true"
|
||||||
of irOr:
|
of irOr:
|
||||||
if left == "true" or right == "true": return "true"
|
if left == "true" or right == "true": return "true"
|
||||||
|
if isNull(left) or isNull(right): return "\\N"
|
||||||
return "false"
|
return "false"
|
||||||
of irAdd, irSub, irMul, irDiv, irMod, irPow:
|
of irAdd, irSub, irMul, irDiv, irMod, irPow:
|
||||||
let v = evalExpr(expr, stringTableToValueRow(row), ctx)
|
let v = evalExpr(expr, stringTableToValueRow(row), ctx)
|
||||||
@@ -474,6 +489,7 @@ proc evalExprOld*(expr: IRExpr, row: Table[string, string], ctx: ExecutionContex
|
|||||||
of vkString: return v.strVal
|
of vkString: return v.strVal
|
||||||
else: return "\\N"
|
else: return "\\N"
|
||||||
of irLike:
|
of irLike:
|
||||||
|
if isNull(left) or isNull(right): return "\\N"
|
||||||
proc escapeRe(s: string): string =
|
proc escapeRe(s: string): string =
|
||||||
result = ""
|
result = ""
|
||||||
for ch in s:
|
for ch in s:
|
||||||
@@ -489,6 +505,7 @@ proc evalExprOld*(expr: IRExpr, row: Table[string, string], ctx: ExecutionContex
|
|||||||
except CatchableError: discard
|
except CatchableError: discard
|
||||||
return "false"
|
return "false"
|
||||||
of irILike:
|
of irILike:
|
||||||
|
if isNull(left) or isNull(right): return "\\N"
|
||||||
proc escapeRe(s: string): string =
|
proc escapeRe(s: string): string =
|
||||||
result = ""
|
result = ""
|
||||||
for ch in s:
|
for ch in s:
|
||||||
@@ -504,8 +521,10 @@ proc evalExprOld*(expr: IRExpr, row: Table[string, string], ctx: ExecutionContex
|
|||||||
except CatchableError: discard
|
except CatchableError: discard
|
||||||
return "false"
|
return "false"
|
||||||
of irIn:
|
of irIn:
|
||||||
|
if isNull(left): return "\\N"
|
||||||
if expr.binRight.kind == irekSubquery:
|
if expr.binRight.kind == irekSubquery:
|
||||||
let subRows = requireExecutePlanHook()(ctx, expr.binRight.subqueryPlan)
|
let subRows = requireExecutePlanHook()(ctx, expr.binRight.subqueryPlan)
|
||||||
|
var sawNull = false
|
||||||
for row in subRows:
|
for row in subRows:
|
||||||
# Compare against the first non-internal column only (SQL semantics)
|
# Compare against the first non-internal column only (SQL semantics)
|
||||||
var firstVal = ""
|
var firstVal = ""
|
||||||
@@ -515,8 +534,14 @@ proc evalExprOld*(expr: IRExpr, row: Table[string, string], ctx: ExecutionContex
|
|||||||
firstVal = valueToString(v)
|
firstVal = valueToString(v)
|
||||||
found = true
|
found = true
|
||||||
break
|
break
|
||||||
if found and firstVal == left: return "true"
|
if not found: continue
|
||||||
|
if isNull(firstVal):
|
||||||
|
sawNull = true
|
||||||
|
continue
|
||||||
|
if firstVal == left: return "true"
|
||||||
|
if sawNull: return "\\N"
|
||||||
return "false"
|
return "false"
|
||||||
|
if isNull(right): return "\\N"
|
||||||
try:
|
try:
|
||||||
let lv = parseFloat(left)
|
let lv = parseFloat(left)
|
||||||
let rv = parseFloat(right)
|
let rv = parseFloat(right)
|
||||||
@@ -524,8 +549,10 @@ proc evalExprOld*(expr: IRExpr, row: Table[string, string], ctx: ExecutionContex
|
|||||||
except CatchableError: discard
|
except CatchableError: discard
|
||||||
return if left == right: "true" else: "false"
|
return if left == right: "true" else: "false"
|
||||||
of irNotIn:
|
of irNotIn:
|
||||||
|
if isNull(left): return "\\N"
|
||||||
if expr.binRight.kind == irekSubquery:
|
if expr.binRight.kind == irekSubquery:
|
||||||
let subRows = requireExecutePlanHook()(ctx, expr.binRight.subqueryPlan)
|
let subRows = requireExecutePlanHook()(ctx, expr.binRight.subqueryPlan)
|
||||||
|
var sawNull = false
|
||||||
for row in subRows:
|
for row in subRows:
|
||||||
# Compare against the first non-internal column only (SQL semantics)
|
# Compare against the first non-internal column only (SQL semantics)
|
||||||
var firstVal = ""
|
var firstVal = ""
|
||||||
@@ -535,8 +562,14 @@ proc evalExprOld*(expr: IRExpr, row: Table[string, string], ctx: ExecutionContex
|
|||||||
firstVal = valueToString(v)
|
firstVal = valueToString(v)
|
||||||
found = true
|
found = true
|
||||||
break
|
break
|
||||||
if found and firstVal == left: return "false"
|
if not found: continue
|
||||||
|
if isNull(firstVal):
|
||||||
|
sawNull = true
|
||||||
|
continue
|
||||||
|
if firstVal == left: return "false"
|
||||||
|
if sawNull: return "\\N"
|
||||||
return "true"
|
return "true"
|
||||||
|
if isNull(right): return "\\N"
|
||||||
try:
|
try:
|
||||||
let lv = parseFloat(left)
|
let lv = parseFloat(left)
|
||||||
let rv = parseFloat(right)
|
let rv = parseFloat(right)
|
||||||
@@ -661,6 +694,7 @@ proc evalExprOld*(expr: IRExpr, row: Table[string, string], ctx: ExecutionContex
|
|||||||
case expr.unOp
|
case expr.unOp
|
||||||
of irNot:
|
of irNot:
|
||||||
let v = evalExprOld(expr.unExpr, row, ctx)
|
let v = evalExprOld(expr.unExpr, row, ctx)
|
||||||
|
if isNull(v): return "\\N"
|
||||||
return if v == "true": "false" else: "true"
|
return if v == "true": "false" else: "true"
|
||||||
of irIsNull:
|
of irIsNull:
|
||||||
let v = evalExprOld(expr.unExpr, row, ctx)
|
let v = evalExprOld(expr.unExpr, row, ctx)
|
||||||
|
|||||||
@@ -76,6 +76,8 @@ proc lowerExpr*(node: Node): IRExpr =
|
|||||||
of bkJsonContainedBy: irOp = irJsonContainedBy
|
of bkJsonContainedBy: irOp = irJsonContainedBy
|
||||||
of bkJsonHasAny: irOp = irJsonHasAny
|
of bkJsonHasAny: irOp = irJsonHasAny
|
||||||
of bkJsonHasAll: irOp = irJsonHasAll
|
of bkJsonHasAll: irOp = irJsonHasAll
|
||||||
|
of bkPow: irOp = irPow
|
||||||
|
of bkConcat: irOp = irAdd # irAdd concatenates string operands
|
||||||
else: irOp = irEq
|
else: irOp = irEq
|
||||||
result.binOp = irOp
|
result.binOp = irOp
|
||||||
result.binLeft = lowerExpr(node.binLeft)
|
result.binLeft = lowerExpr(node.binLeft)
|
||||||
@@ -120,6 +122,7 @@ proc lowerExpr*(node: Node): IRExpr =
|
|||||||
else: discard
|
else: discard
|
||||||
result.aggArgs = @[]
|
result.aggArgs = @[]
|
||||||
for arg in node.funcArgs: result.aggArgs.add(lowerExpr(arg))
|
for arg in node.funcArgs: result.aggArgs.add(lowerExpr(arg))
|
||||||
|
result.aggDistinct = node.funcDistinct
|
||||||
if node.funcFilter != nil:
|
if node.funcFilter != nil:
|
||||||
result.aggFilter = lowerExpr(node.funcFilter)
|
result.aggFilter = lowerExpr(node.funcFilter)
|
||||||
else:
|
else:
|
||||||
@@ -408,9 +411,18 @@ proc lowerSelect*(node: Node): IRPlan =
|
|||||||
if node.selLimit != nil or node.selOffset != nil:
|
if node.selLimit != nil or node.selOffset != nil:
|
||||||
let limitPlan = IRPlan(kind: irpkLimit)
|
let limitPlan = IRPlan(kind: irpkLimit)
|
||||||
limitPlan.limitSource = result
|
limitPlan.limitSource = result
|
||||||
limitPlan.limitCount = if node.selLimit != nil and node.selLimit.limitExpr.kind == nkIntLit:
|
# limitCount: -1 = unlimited (OFFSET without LIMIT). LIMIT 0 is empty.
|
||||||
node.selLimit.limitExpr.intVal else: 0
|
# Negative LIMIT/OFFSET are clamped so slicing cannot IndexDefect.
|
||||||
limitPlan.limitOffset = if node.selOffset != nil and node.selOffset.offsetExpr.kind == nkIntLit:
|
if node.selLimit != nil:
|
||||||
node.selOffset.offsetExpr.intVal else: 0
|
if node.selLimit.limitExpr.kind == nkIntLit:
|
||||||
|
limitPlan.limitCount = max(0'i64, node.selLimit.limitExpr.intVal)
|
||||||
|
else:
|
||||||
|
limitPlan.limitCount = 0
|
||||||
|
else:
|
||||||
|
limitPlan.limitCount = -1
|
||||||
|
if node.selOffset != nil and node.selOffset.offsetExpr.kind == nkIntLit:
|
||||||
|
limitPlan.limitOffset = max(0'i64, node.selOffset.offsetExpr.intVal)
|
||||||
|
else:
|
||||||
|
limitPlan.limitOffset = 0
|
||||||
result = limitPlan
|
result = limitPlan
|
||||||
|
|
||||||
|
|||||||
@@ -5,6 +5,7 @@
|
|||||||
## executor split). Pure code motion — no behavior changes.
|
## executor split). Pure code motion — no behavior changes.
|
||||||
import std/strutils
|
import std/strutils
|
||||||
import std/tables
|
import std/tables
|
||||||
|
import std/sets
|
||||||
import std/sequtils
|
import std/sequtils
|
||||||
import std/algorithm
|
import std/algorithm
|
||||||
import ../ir
|
import ../ir
|
||||||
@@ -19,6 +20,19 @@ import eval
|
|||||||
import scan
|
import scan
|
||||||
import window
|
import window
|
||||||
|
|
||||||
|
# ----------------------------------------------------------------------
|
||||||
|
# Aggregate DISTINCT helpers
|
||||||
|
# ----------------------------------------------------------------------
|
||||||
|
|
||||||
|
proc shouldKeepDistinct(seen: var HashSet[string], s: string, doDistinct: bool): bool =
|
||||||
|
## Returns true if `s` should be counted/included (first occurrence when distinct).
|
||||||
|
if not doDistinct:
|
||||||
|
return true
|
||||||
|
if s in seen:
|
||||||
|
return false
|
||||||
|
seen.incl(s)
|
||||||
|
return true
|
||||||
|
|
||||||
# ----------------------------------------------------------------------
|
# ----------------------------------------------------------------------
|
||||||
# IR Plan Execution (with actual filter/sort/projection)
|
# IR Plan Execution (with actual filter/sort/projection)
|
||||||
# ----------------------------------------------------------------------
|
# ----------------------------------------------------------------------
|
||||||
@@ -113,49 +127,71 @@ proc executePlan*(ctx: ExecutionContext, plan: IRPlan): seq[Row] =
|
|||||||
newRow[alias] = $filteredRows.len
|
newRow[alias] = $filteredRows.len
|
||||||
else:
|
else:
|
||||||
var count = 0
|
var count = 0
|
||||||
|
var seen: HashSet[string]
|
||||||
for row in filteredRows:
|
for row in filteredRows:
|
||||||
let v = evalExpr(expr.aggArgs[0], row, ctx)
|
let v = evalExpr(expr.aggArgs[0], row, ctx)
|
||||||
if v.kind != vkNull: count += 1
|
if v.kind != vkNull:
|
||||||
|
let s = valueToString(v)
|
||||||
|
if shouldKeepDistinct(seen, s, expr.aggDistinct):
|
||||||
|
count += 1
|
||||||
newRow[alias] = $count
|
newRow[alias] = $count
|
||||||
of irSum:
|
of irSum:
|
||||||
var sum = 0.0
|
var sum = 0.0
|
||||||
|
var seen: HashSet[string]
|
||||||
for row in filteredRows:
|
for row in filteredRows:
|
||||||
let v = evalExpr(expr.aggArgs[0], row, ctx)
|
let v = evalExpr(expr.aggArgs[0], row, ctx)
|
||||||
try: sum += parseFloat(valueToString(v)) except CatchableError: discard
|
let s = valueToString(v)
|
||||||
|
if shouldKeepDistinct(seen, s, expr.aggDistinct):
|
||||||
|
try: sum += parseFloat(s) except CatchableError: discard
|
||||||
newRow[alias] = $sum
|
newRow[alias] = $sum
|
||||||
of irAvg:
|
of irAvg:
|
||||||
var sum = 0.0
|
var sum = 0.0
|
||||||
var count = 0
|
var count = 0
|
||||||
|
var seen: HashSet[string]
|
||||||
for row in filteredRows:
|
for row in filteredRows:
|
||||||
let v = evalExpr(expr.aggArgs[0], row, ctx)
|
let v = evalExpr(expr.aggArgs[0], row, ctx)
|
||||||
try: sum += parseFloat(valueToString(v)); count += 1 except CatchableError: discard
|
let s = valueToString(v)
|
||||||
|
if shouldKeepDistinct(seen, s, expr.aggDistinct):
|
||||||
|
try: sum += parseFloat(s); count += 1 except CatchableError: discard
|
||||||
newRow[alias] = if count > 0: $(sum / float(count)) else: "0"
|
newRow[alias] = if count > 0: $(sum / float(count)) else: "0"
|
||||||
of irMin:
|
of irMin:
|
||||||
var minVal = ""
|
var minVal = ""
|
||||||
|
var seen: HashSet[string]
|
||||||
for row in filteredRows:
|
for row in filteredRows:
|
||||||
let v = evalExpr(expr.aggArgs[0], row, ctx)
|
let v = evalExpr(expr.aggArgs[0], row, ctx)
|
||||||
if v.kind == vkNull: continue
|
if v.kind == vkNull: continue
|
||||||
if minVal == "" or cmpMin(valueToString(v), minVal): minVal = valueToString(v)
|
let s = valueToString(v)
|
||||||
|
if shouldKeepDistinct(seen, s, expr.aggDistinct):
|
||||||
|
if minVal == "" or cmpMin(s, minVal): minVal = s
|
||||||
newRow[alias] = minVal
|
newRow[alias] = minVal
|
||||||
of irMax:
|
of irMax:
|
||||||
var maxVal = ""
|
var maxVal = ""
|
||||||
|
var seen: HashSet[string]
|
||||||
for row in filteredRows:
|
for row in filteredRows:
|
||||||
let v = evalExpr(expr.aggArgs[0], row, ctx)
|
let v = evalExpr(expr.aggArgs[0], row, ctx)
|
||||||
if v.kind == vkNull: continue
|
if v.kind == vkNull: continue
|
||||||
if maxVal == "" or cmpMax(valueToString(v), maxVal): maxVal = valueToString(v)
|
let s = valueToString(v)
|
||||||
|
if shouldKeepDistinct(seen, s, expr.aggDistinct):
|
||||||
|
if maxVal == "" or cmpMax(s, maxVal): maxVal = s
|
||||||
newRow[alias] = maxVal
|
newRow[alias] = maxVal
|
||||||
of irArrayAgg:
|
of irArrayAgg:
|
||||||
var arr: seq[string]
|
var arr: seq[string]
|
||||||
|
var seen: HashSet[string]
|
||||||
for row in filteredRows:
|
for row in filteredRows:
|
||||||
if expr.aggArgs.len > 0:
|
if expr.aggArgs.len > 0:
|
||||||
arr.add(valueToString(evalExpr(expr.aggArgs[0], row, ctx)))
|
let s = valueToString(evalExpr(expr.aggArgs[0], row, ctx))
|
||||||
|
if shouldKeepDistinct(seen, s, expr.aggDistinct):
|
||||||
|
arr.add(s)
|
||||||
newRow[alias] = "[" & arr.join(", ") & "]"
|
newRow[alias] = "[" & arr.join(", ") & "]"
|
||||||
of irStringAgg:
|
of irStringAgg:
|
||||||
var parts: seq[string]
|
var parts: seq[string]
|
||||||
|
var seen: HashSet[string]
|
||||||
let delim = if expr.aggArgs.len > 1: evalExpr(expr.aggArgs[1], initTable[string, Value](), ctx) else: Value(kind: vkString, strVal: ",")
|
let delim = if expr.aggArgs.len > 1: evalExpr(expr.aggArgs[1], initTable[string, Value](), ctx) else: Value(kind: vkString, strVal: ",")
|
||||||
for row in filteredRows:
|
for row in filteredRows:
|
||||||
if expr.aggArgs.len > 0:
|
if expr.aggArgs.len > 0:
|
||||||
parts.add(valueToString(evalExpr(expr.aggArgs[0], row, ctx)))
|
let s = valueToString(evalExpr(expr.aggArgs[0], row, ctx))
|
||||||
|
if shouldKeepDistinct(seen, s, expr.aggDistinct):
|
||||||
|
parts.add(s)
|
||||||
newRow[alias] = parts.join(valueToString(delim))
|
newRow[alias] = parts.join(valueToString(delim))
|
||||||
else:
|
else:
|
||||||
let val = evalExpr(expr, if sourceRows.len > 0: sourceRows[0] else: initTable[string, Value](), ctx)
|
let val = evalExpr(expr, if sourceRows.len > 0: sourceRows[0] else: initTable[string, Value](), ctx)
|
||||||
@@ -219,12 +255,18 @@ proc executePlan*(ctx: ExecutionContext, plan: IRPlan): seq[Row] =
|
|||||||
of irpkLimit:
|
of irpkLimit:
|
||||||
let sourceRows = executePlan(ctx, plan.limitSource)
|
let sourceRows = executePlan(ctx, plan.limitSource)
|
||||||
var start = int(plan.limitOffset)
|
var start = int(plan.limitOffset)
|
||||||
|
if start < 0: start = 0
|
||||||
if start > sourceRows.len: start = sourceRows.len
|
if start > sourceRows.len: start = sourceRows.len
|
||||||
|
if plan.limitCount < 0:
|
||||||
|
# OFFSET without LIMIT — return the remainder.
|
||||||
|
return sourceRows[start ..< sourceRows.len]
|
||||||
if plan.limitCount == 0:
|
if plan.limitCount == 0:
|
||||||
return @[]
|
return @[]
|
||||||
var endIdx = start + int(plan.limitCount)
|
var endIdx = start + int(plan.limitCount)
|
||||||
if endIdx > sourceRows.len:
|
if endIdx > sourceRows.len:
|
||||||
endIdx = sourceRows.len
|
endIdx = sourceRows.len
|
||||||
|
if endIdx < start:
|
||||||
|
endIdx = start
|
||||||
return sourceRows[start..<endIdx]
|
return sourceRows[start..<endIdx]
|
||||||
|
|
||||||
of irpkGroupBy:
|
of irpkGroupBy:
|
||||||
@@ -292,49 +334,71 @@ proc executePlan*(ctx: ExecutionContext, plan: IRPlan): seq[Row] =
|
|||||||
aggRow[aggKey] = $filteredRows.len
|
aggRow[aggKey] = $filteredRows.len
|
||||||
else:
|
else:
|
||||||
var count = 0
|
var count = 0
|
||||||
|
var seen: HashSet[string]
|
||||||
for row in filteredRows:
|
for row in filteredRows:
|
||||||
let v = evalExpr(aggExpr.aggArgs[0], row, ctx)
|
let v = evalExpr(aggExpr.aggArgs[0], row, ctx)
|
||||||
if v.kind != vkNull: count += 1
|
if v.kind != vkNull:
|
||||||
|
let s = valueToString(v)
|
||||||
|
if shouldKeepDistinct(seen, s, aggExpr.aggDistinct):
|
||||||
|
count += 1
|
||||||
aggRow[aggKey] = $count
|
aggRow[aggKey] = $count
|
||||||
of irSum:
|
of irSum:
|
||||||
var sum = 0.0
|
var sum = 0.0
|
||||||
|
var seen: HashSet[string]
|
||||||
for row in filteredRows:
|
for row in filteredRows:
|
||||||
let v = evalExpr(aggExpr.aggArgs[0], row, ctx)
|
let v = evalExpr(aggExpr.aggArgs[0], row, ctx)
|
||||||
try: sum += parseFloat(valueToString(v)) except CatchableError: discard
|
let s = valueToString(v)
|
||||||
|
if shouldKeepDistinct(seen, s, aggExpr.aggDistinct):
|
||||||
|
try: sum += parseFloat(s) except CatchableError: discard
|
||||||
aggRow[aggKey] = $sum
|
aggRow[aggKey] = $sum
|
||||||
of irAvg:
|
of irAvg:
|
||||||
var sum = 0.0
|
var sum = 0.0
|
||||||
var count = 0
|
var count = 0
|
||||||
|
var seen: HashSet[string]
|
||||||
for row in filteredRows:
|
for row in filteredRows:
|
||||||
let v = evalExpr(aggExpr.aggArgs[0], row, ctx)
|
let v = evalExpr(aggExpr.aggArgs[0], row, ctx)
|
||||||
try: sum += parseFloat(valueToString(v)); count += 1 except CatchableError: discard
|
let s = valueToString(v)
|
||||||
|
if shouldKeepDistinct(seen, s, aggExpr.aggDistinct):
|
||||||
|
try: sum += parseFloat(s); count += 1 except CatchableError: discard
|
||||||
aggRow[aggKey] = if count > 0: $(sum / float(count)) else: "0"
|
aggRow[aggKey] = if count > 0: $(sum / float(count)) else: "0"
|
||||||
of irMin:
|
of irMin:
|
||||||
var minVal = ""
|
var minVal = ""
|
||||||
|
var seen: HashSet[string]
|
||||||
for row in filteredRows:
|
for row in filteredRows:
|
||||||
let v = evalExpr(aggExpr.aggArgs[0], row, ctx)
|
let v = evalExpr(aggExpr.aggArgs[0], row, ctx)
|
||||||
if v.kind == vkNull: continue
|
if v.kind == vkNull: continue
|
||||||
if minVal == "" or cmpMin(valueToString(v), minVal): minVal = valueToString(v)
|
let s = valueToString(v)
|
||||||
|
if shouldKeepDistinct(seen, s, aggExpr.aggDistinct):
|
||||||
|
if minVal == "" or cmpMin(s, minVal): minVal = s
|
||||||
aggRow[aggKey] = minVal
|
aggRow[aggKey] = minVal
|
||||||
of irMax:
|
of irMax:
|
||||||
var maxVal = ""
|
var maxVal = ""
|
||||||
|
var seen: HashSet[string]
|
||||||
for row in filteredRows:
|
for row in filteredRows:
|
||||||
let v = evalExpr(aggExpr.aggArgs[0], row, ctx)
|
let v = evalExpr(aggExpr.aggArgs[0], row, ctx)
|
||||||
if v.kind == vkNull: continue
|
if v.kind == vkNull: continue
|
||||||
if maxVal == "" or cmpMax(valueToString(v), maxVal): maxVal = valueToString(v)
|
let s = valueToString(v)
|
||||||
|
if shouldKeepDistinct(seen, s, aggExpr.aggDistinct):
|
||||||
|
if maxVal == "" or cmpMax(s, maxVal): maxVal = s
|
||||||
aggRow[aggKey] = maxVal
|
aggRow[aggKey] = maxVal
|
||||||
of irArrayAgg:
|
of irArrayAgg:
|
||||||
var arr: seq[string]
|
var arr: seq[string]
|
||||||
|
var seen: HashSet[string]
|
||||||
for row in filteredRows:
|
for row in filteredRows:
|
||||||
if aggExpr.aggArgs.len > 0:
|
if aggExpr.aggArgs.len > 0:
|
||||||
arr.add(valueToString(evalExpr(aggExpr.aggArgs[0], row, ctx)))
|
let s = valueToString(evalExpr(aggExpr.aggArgs[0], row, ctx))
|
||||||
|
if shouldKeepDistinct(seen, s, aggExpr.aggDistinct):
|
||||||
|
arr.add(s)
|
||||||
aggRow[aggKey] = "[" & arr.join(", ") & "]"
|
aggRow[aggKey] = "[" & arr.join(", ") & "]"
|
||||||
of irStringAgg:
|
of irStringAgg:
|
||||||
var parts: seq[string]
|
var parts: seq[string]
|
||||||
|
var seen: HashSet[string]
|
||||||
let delim = if aggExpr.aggArgs.len > 1: evalExpr(aggExpr.aggArgs[1], initTable[string, Value](), ctx) else: Value(kind: vkString, strVal: ",")
|
let delim = if aggExpr.aggArgs.len > 1: evalExpr(aggExpr.aggArgs[1], initTable[string, Value](), ctx) else: Value(kind: vkString, strVal: ",")
|
||||||
for row in filteredRows:
|
for row in filteredRows:
|
||||||
if aggExpr.aggArgs.len > 0:
|
if aggExpr.aggArgs.len > 0:
|
||||||
parts.add(valueToString(evalExpr(aggExpr.aggArgs[0], row, ctx)))
|
let s = valueToString(evalExpr(aggExpr.aggArgs[0], row, ctx))
|
||||||
|
if shouldKeepDistinct(seen, s, aggExpr.aggDistinct):
|
||||||
|
parts.add(s)
|
||||||
aggRow[aggKey] = parts.join(valueToString(delim))
|
aggRow[aggKey] = parts.join(valueToString(delim))
|
||||||
# Apply HAVING filter
|
# Apply HAVING filter
|
||||||
if plan.groupHaving != nil:
|
if plan.groupHaving != nil:
|
||||||
|
|||||||
@@ -11,20 +11,22 @@ import lower
|
|||||||
# Row-Level Security
|
# Row-Level Security
|
||||||
# ----------------------------------------------------------------------
|
# ----------------------------------------------------------------------
|
||||||
|
|
||||||
proc hasPrivilege*(ctx: ExecutionContext, tableName, command: string): bool =
|
proc hasPrivilegeFor*(ctx: ExecutionContext, username, tableName, command: string): bool =
|
||||||
if ctx.currentUser.len == 0: return true
|
## Privilege check for an explicit username (does not mutate ctx.currentUser).
|
||||||
let user = ctx.users.getOrDefault(ctx.currentUser)
|
if username.len == 0: return true
|
||||||
|
let user = ctx.users.getOrDefault(username)
|
||||||
if user.isSuperuser: return true
|
if user.isSuperuser: return true
|
||||||
# Check table-level policies for user or PUBLIC
|
|
||||||
# For now: if no policies exist, allow everything (backward compatible)
|
|
||||||
if tableName notin ctx.policies: return true
|
if tableName notin ctx.policies: return true
|
||||||
let policies = ctx.policies[tableName]
|
let policies = ctx.policies[tableName]
|
||||||
# If RLS is enabled (policies exist), check if user matches any policy
|
|
||||||
for pol in policies:
|
for pol in policies:
|
||||||
if pol.command == "ALL" or pol.command == command:
|
if pol.command == "ALL" or pol.command == command:
|
||||||
return true
|
return true
|
||||||
return false
|
return false
|
||||||
|
|
||||||
|
proc hasPrivilege*(ctx: ExecutionContext, tableName, command: string): bool =
|
||||||
|
if ctx.currentUser.len == 0: return true
|
||||||
|
hasPrivilegeFor(ctx, ctx.currentUser, tableName, command)
|
||||||
|
|
||||||
proc passesPolicy*(ctx: ExecutionContext, tableName, command: string, row: Row): bool =
|
proc passesPolicy*(ctx: ExecutionContext, tableName, command: string, row: Row): bool =
|
||||||
if ctx.currentUser.len == 0: return true
|
if ctx.currentUser.len == 0: return true
|
||||||
let user = ctx.users.getOrDefault(ctx.currentUser)
|
let user = ctx.users.getOrDefault(ctx.currentUser)
|
||||||
|
|||||||
@@ -177,6 +177,64 @@ proc computeWindowValues*(rows: seq[Row], expr: IRExpr, ctx: ExecutionContext =
|
|||||||
for pos, rowIdx in sortedIdxs:
|
for pos, rowIdx in sortedIdxs:
|
||||||
let (_, fEnd) = resolveFrameBounds(pos, sortedIdxs.len, frameStart, frameEnd)
|
let (_, fEnd) = resolveFrameBounds(pos, sortedIdxs.len, frameStart, frameEnd)
|
||||||
result[rowIdx] = valueToString(evalExpr(expr.wfArgs[0], rows[sortedIdxs[fEnd]], ctx))
|
result[rowIdx] = valueToString(evalExpr(expr.wfArgs[0], rows[sortedIdxs[fEnd]], ctx))
|
||||||
|
of "sum", "avg", "count", "min", "max":
|
||||||
|
let countAll = wfName == "count" and
|
||||||
|
(expr.wfArgs.len == 0 or expr.wfArgs[0].kind == irekStar)
|
||||||
|
for pos, rowIdx in sortedIdxs:
|
||||||
|
let (fStart, fEnd) = resolveFrameBounds(pos, sortedIdxs.len, frameStart, frameEnd)
|
||||||
|
if wfName == "count" and countAll:
|
||||||
|
result[rowIdx] = $(fEnd - fStart + 1)
|
||||||
|
continue
|
||||||
|
if wfName == "count":
|
||||||
|
var cnt = 0
|
||||||
|
if expr.wfArgs.len > 0:
|
||||||
|
for i in fStart .. fEnd:
|
||||||
|
let s = valueToString(evalExpr(expr.wfArgs[0], rows[sortedIdxs[i]], ctx))
|
||||||
|
if not isNull(s) and s.len > 0:
|
||||||
|
inc cnt
|
||||||
|
result[rowIdx] = $cnt
|
||||||
|
continue
|
||||||
|
if expr.wfArgs.len == 0:
|
||||||
|
result[rowIdx] = "\\N"
|
||||||
|
continue
|
||||||
|
var sum = 0.0
|
||||||
|
var cnt = 0
|
||||||
|
var minF = 0.0
|
||||||
|
var maxF = 0.0
|
||||||
|
var minS = ""
|
||||||
|
var maxS = ""
|
||||||
|
var allNumeric = true
|
||||||
|
for i in fStart .. fEnd:
|
||||||
|
let s = valueToString(evalExpr(expr.wfArgs[0], rows[sortedIdxs[i]], ctx))
|
||||||
|
if isNull(s) or s.len == 0: continue
|
||||||
|
inc cnt
|
||||||
|
if cnt == 1 or s < minS: minS = s
|
||||||
|
if cnt == 1 or s > maxS: maxS = s
|
||||||
|
try:
|
||||||
|
let f = parseFloat(s)
|
||||||
|
sum += f
|
||||||
|
if cnt == 1:
|
||||||
|
minF = f
|
||||||
|
maxF = f
|
||||||
|
else:
|
||||||
|
if f < minF: minF = f
|
||||||
|
if f > maxF: maxF = f
|
||||||
|
except CatchableError:
|
||||||
|
allNumeric = false
|
||||||
|
if cnt == 0:
|
||||||
|
result[rowIdx] = "\\N"
|
||||||
|
else:
|
||||||
|
case wfName
|
||||||
|
of "sum":
|
||||||
|
result[rowIdx] = if allNumeric: $sum else: "\\N"
|
||||||
|
of "avg":
|
||||||
|
result[rowIdx] = if allNumeric: $(sum / float(cnt)) else: "\\N"
|
||||||
|
of "min":
|
||||||
|
result[rowIdx] = if allNumeric: $minF else: minS
|
||||||
|
of "max":
|
||||||
|
result[rowIdx] = if allNumeric: $maxF else: maxS
|
||||||
|
else:
|
||||||
|
result[rowIdx] = "\\N"
|
||||||
else:
|
else:
|
||||||
# Unknown window function — fill with null
|
# Unknown window function — fill with null
|
||||||
for rowIdx in sortedIdxs:
|
for rowIdx in sortedIdxs:
|
||||||
|
|||||||
@@ -448,6 +448,26 @@ proc executeQueryImpl(ctx: ExecutionContext, astNode: Node, params: seq[WireValu
|
|||||||
if cols.len == 0:
|
if cols.len == 0:
|
||||||
cols = rightRes.columns
|
cols = rightRes.columns
|
||||||
|
|
||||||
|
# Fingerprint a projected row for set-op dedup. Prefer declared columns;
|
||||||
|
# fall back to non-system keys so UNION/INTERSECT/EXCEPT work without `$value`.
|
||||||
|
proc setOpRowKey(row: Row, colNames: seq[string]): string =
|
||||||
|
var parts: seq[string] = @[]
|
||||||
|
if colNames.len > 0:
|
||||||
|
for c in colNames:
|
||||||
|
if c in row:
|
||||||
|
parts.add(valueToString(row[c]))
|
||||||
|
else:
|
||||||
|
parts.add("")
|
||||||
|
else:
|
||||||
|
var keys: seq[string] = @[]
|
||||||
|
for k, _ in row:
|
||||||
|
if not k.startsWith("$"):
|
||||||
|
keys.add(k)
|
||||||
|
keys.sort()
|
||||||
|
for k in keys:
|
||||||
|
parts.add(k & "=" & valueToString(row[k]))
|
||||||
|
return parts.join("\x1f")
|
||||||
|
|
||||||
var rows: seq[Row] = @[]
|
var rows: seq[Row] = @[]
|
||||||
case stmt.setOpKind
|
case stmt.setOpKind
|
||||||
of sdkUnion:
|
of sdkUnion:
|
||||||
@@ -460,28 +480,30 @@ proc executeQueryImpl(ctx: ExecutionContext, astNode: Node, params: seq[WireValu
|
|||||||
# UNION: deduplicate
|
# UNION: deduplicate
|
||||||
var seen: Table[string, bool]
|
var seen: Table[string, bool]
|
||||||
for row in leftRes.rows:
|
for row in leftRes.rows:
|
||||||
seen[valueToString(row["$value"])] = true
|
seen[setOpRowKey(row, cols)] = true
|
||||||
for row in rightRes.rows:
|
for row in rightRes.rows:
|
||||||
if not seen.getOrDefault(valueToString(row["$value"]), false):
|
let k = setOpRowKey(row, cols)
|
||||||
seen[valueToString(row["$value"])] = true
|
if not seen.getOrDefault(k, false):
|
||||||
|
seen[k] = true
|
||||||
rows.add(row)
|
rows.add(row)
|
||||||
|
|
||||||
of sdkIntersect:
|
of sdkIntersect:
|
||||||
var leftSet: Table[string, bool]
|
var leftSet: Table[string, bool]
|
||||||
for row in leftRes.rows:
|
for row in leftRes.rows:
|
||||||
leftSet[valueToString(row["$value"])] = true
|
leftSet[setOpRowKey(row, cols)] = true
|
||||||
for row in rightRes.rows:
|
for row in rightRes.rows:
|
||||||
if leftSet.getOrDefault(valueToString(row["$value"]), false):
|
let k = setOpRowKey(row, cols)
|
||||||
|
if leftSet.getOrDefault(k, false):
|
||||||
rows.add(row)
|
rows.add(row)
|
||||||
if not stmt.setOpAll:
|
if not stmt.setOpAll:
|
||||||
leftSet.del(valueToString(row["$value"])) # remove to prevent duplicates for INTERSECT (not ALL)
|
leftSet.del(k) # remove to prevent duplicates for INTERSECT (not ALL)
|
||||||
|
|
||||||
of sdkExcept:
|
of sdkExcept:
|
||||||
var rightSet: Table[string, bool]
|
var rightSet: Table[string, bool]
|
||||||
for row in rightRes.rows:
|
for row in rightRes.rows:
|
||||||
rightSet[valueToString(row["$value"])] = true
|
rightSet[setOpRowKey(row, cols)] = true
|
||||||
for row in leftRes.rows:
|
for row in leftRes.rows:
|
||||||
if not rightSet.getOrDefault(valueToString(row["$value"]), false):
|
if not rightSet.getOrDefault(setOpRowKey(row, cols), false):
|
||||||
rows.add(row)
|
rows.add(row)
|
||||||
|
|
||||||
return okResult(rows, cols)
|
return okResult(rows, cols)
|
||||||
@@ -759,21 +781,34 @@ proc executeQueryImpl(ctx: ExecutionContext, astNode: Node, params: seq[WireValu
|
|||||||
let onExpr = lowerExpr(stmt.mergeOn)
|
let onExpr = lowerExpr(stmt.mergeOn)
|
||||||
if valueToString(evalExpr(onExpr, rowWithTarget, ctx)) == "true":
|
if valueToString(evalExpr(onExpr, rowWithTarget, ctx)) == "true":
|
||||||
matched = true
|
matched = true
|
||||||
if stmt.mergeMatchedUpdate.len > 0 and "$key" in tgtRow:
|
# Optional AND <condition> after WHEN MATCHED
|
||||||
var updateSets = initTable[string, string]()
|
var applyMatched = true
|
||||||
for s in stmt.mergeMatchedUpdate:
|
if stmt.mergeMatchedCondition != nil:
|
||||||
if s.kind == nkBinOp and s.binOp == bkAssign:
|
let condExpr = lowerExpr(stmt.mergeMatchedCondition)
|
||||||
if s.binLeft.kind == nkIdent:
|
applyMatched = valueToString(evalExpr(condExpr, rowWithTarget, ctx)) == "true"
|
||||||
let valExpr = lowerExpr(s.binRight)
|
if applyMatched and "$key" in tgtRow:
|
||||||
updateSets[s.binLeft.identName] = valueToString(evalExpr(valExpr, rowWithTarget, ctx))
|
if stmt.mergeMatchedDelete:
|
||||||
var newRow = tgtRow
|
fireTriggers(ctx, stmt.mergeTarget, "before", "delete", tgtRow)
|
||||||
for col, val in updateSets:
|
count += execDelete(ctx, stmt.mergeTarget, valueToString(tgtRow["$key"]), kvPairs)
|
||||||
newRow[col] = Value(kind: vkString, strVal: val)
|
fireTriggers(ctx, stmt.mergeTarget, "after", "delete", tgtRow)
|
||||||
fireTriggers(ctx, stmt.mergeTarget, "before", "update", tgtRow)
|
if ctx.onChange != nil:
|
||||||
count += execUpdateRow(ctx, stmt.mergeTarget, valueToString(tgtRow["$key"]), updateSets, kvPairs)
|
ctx.onChange(ChangeEvent(table: stmt.mergeTarget, kind: ckDelete,
|
||||||
fireTriggers(ctx, stmt.mergeTarget, "after", "update", newRow)
|
key: valueToString(tgtRow["$key"]), data: ""))
|
||||||
if ctx.onChange != nil:
|
elif stmt.mergeMatchedUpdate.len > 0:
|
||||||
ctx.onChange(ChangeEvent(table: stmt.mergeTarget, kind: ckUpdate, key: valueToString(tgtRow["$key"]), data: ""))
|
var updateSets = initTable[string, string]()
|
||||||
|
for s in stmt.mergeMatchedUpdate:
|
||||||
|
if s.kind == nkBinOp and s.binOp == bkAssign:
|
||||||
|
if s.binLeft.kind == nkIdent:
|
||||||
|
let valExpr = lowerExpr(s.binRight)
|
||||||
|
updateSets[s.binLeft.identName] = valueToString(evalExpr(valExpr, rowWithTarget, ctx))
|
||||||
|
var newRow = tgtRow
|
||||||
|
for col, val in updateSets:
|
||||||
|
newRow[col] = Value(kind: vkString, strVal: val)
|
||||||
|
fireTriggers(ctx, stmt.mergeTarget, "before", "update", tgtRow)
|
||||||
|
count += execUpdateRow(ctx, stmt.mergeTarget, valueToString(tgtRow["$key"]), updateSets, kvPairs)
|
||||||
|
fireTriggers(ctx, stmt.mergeTarget, "after", "update", newRow)
|
||||||
|
if ctx.onChange != nil:
|
||||||
|
ctx.onChange(ChangeEvent(table: stmt.mergeTarget, kind: ckUpdate, key: valueToString(tgtRow["$key"]), data: ""))
|
||||||
break
|
break
|
||||||
|
|
||||||
if not matched and stmt.mergeNotMatchedInsert.len > 0:
|
if not matched and stmt.mergeNotMatchedInsert.len > 0:
|
||||||
|
|||||||
@@ -194,6 +194,11 @@ proc parsePrimary(p: var Parser): Node =
|
|||||||
discard p.expect(tkWhere)
|
discard p.expect(tkWhere)
|
||||||
node.funcFilter = p.parseExpr()
|
node.funcFilter = p.parseExpr()
|
||||||
discard p.expect(tkRParen)
|
discard p.expect(tkRParen)
|
||||||
|
# Window aggregate: SUM/AVG/COUNT/MIN/MAX(...) OVER (...)
|
||||||
|
if p.peek().kind == tkOver:
|
||||||
|
let overClause = p.parseOverClause()
|
||||||
|
return Node(kind: nkWindowExpr, winFunc: funcName.toLower(), winArgs: args,
|
||||||
|
winOver: overClause, line: tok.line, col: tok.col)
|
||||||
return node
|
return node
|
||||||
of tkCase:
|
of tkCase:
|
||||||
discard p.advance()
|
discard p.advance()
|
||||||
|
|||||||
@@ -175,6 +175,61 @@ proc scan*[K, V](btree: BTreeIndex[K, V], startKey, endKey: K): seq[(K, seq[V])]
|
|||||||
finally:
|
finally:
|
||||||
release(btree.lock)
|
release(btree.lock)
|
||||||
|
|
||||||
|
proc subtreeMinMax[K, V](node: BTreeNode[K, V]): (bool, K, K) =
|
||||||
|
## Inclusive min/max of keys stored in this subtree's leaves.
|
||||||
|
if node == nil:
|
||||||
|
return (false, default(K), default(K))
|
||||||
|
if node.isLeaf:
|
||||||
|
if node.keys.len == 0:
|
||||||
|
return (false, default(K), default(K))
|
||||||
|
return (true, node.keys[0], node.keys[^1])
|
||||||
|
var have = false
|
||||||
|
var mn, mx: K
|
||||||
|
for c in node.children:
|
||||||
|
let (ok, a, b) = subtreeMinMax(c)
|
||||||
|
if ok:
|
||||||
|
if not have:
|
||||||
|
mn = a
|
||||||
|
mx = b
|
||||||
|
have = true
|
||||||
|
else:
|
||||||
|
if a < mn: mn = a
|
||||||
|
if b > mx: mx = b
|
||||||
|
return (have, mn, mx)
|
||||||
|
|
||||||
|
proc collectSeparatorErrors[K, V](node: BTreeNode[K, V], errors: var seq[string]) =
|
||||||
|
## Search uses `key > separator → right child`, so every key in the left
|
||||||
|
## subtree must be <= sep (otherwise it is routed right and missed — no
|
||||||
|
## prev-leaf pointer). Boundary duplicates are allowed: the same key may be
|
||||||
|
## max(left) and min(right); get/remove walk `next` to collect them.
|
||||||
|
if node == nil or node.isLeaf:
|
||||||
|
return
|
||||||
|
for i in 0..<node.keys.len:
|
||||||
|
if i + 1 >= node.children.len:
|
||||||
|
errors.add("internal node has fewer children than keys+1")
|
||||||
|
break
|
||||||
|
let (lok, _, lmax) = subtreeMinMax(node.children[i])
|
||||||
|
let (rok, rmin, _) = subtreeMinMax(node.children[i + 1])
|
||||||
|
let sep = node.keys[i]
|
||||||
|
if lok and lmax > sep:
|
||||||
|
errors.add("separator[" & $i & "]=" & $sep &
|
||||||
|
" < max(left)=" & $lmax & " (search would miss left keys)")
|
||||||
|
if lok and rok and rmin < lmax:
|
||||||
|
errors.add("separator[" & $i & "]=" & $sep &
|
||||||
|
" leaf order inverted: max(left)=" & $lmax & " > min(right)=" & $rmin)
|
||||||
|
collectSeparatorErrors(node.children[i], errors)
|
||||||
|
if node.children.len > 0:
|
||||||
|
collectSeparatorErrors(node.children[^1], errors)
|
||||||
|
|
||||||
|
proc checkInvariants*[K, V](btree: BTreeIndex[K, V]): seq[string] =
|
||||||
|
## Returns a list of separator/search-routing violations (empty = healthy).
|
||||||
|
acquire(btree.lock)
|
||||||
|
try:
|
||||||
|
result = @[]
|
||||||
|
collectSeparatorErrors(btree.root, result)
|
||||||
|
finally:
|
||||||
|
release(btree.lock)
|
||||||
|
|
||||||
proc len*[K, V](btree: BTreeIndex[K, V]): int =
|
proc len*[K, V](btree: BTreeIndex[K, V]): int =
|
||||||
acquire(btree.lock)
|
acquire(btree.lock)
|
||||||
try:
|
try:
|
||||||
@@ -209,7 +264,9 @@ proc borrowFromLeft[K, V](node: BTreeNode[K, V], parent: BTreeNode[K, V], parent
|
|||||||
node.values.insert(borrowVal, 0)
|
node.values.insert(borrowVal, 0)
|
||||||
sibling.keys.setLen(sibling.keys.len - 1)
|
sibling.keys.setLen(sibling.keys.len - 1)
|
||||||
sibling.values.setLen(sibling.values.len - 1)
|
sibling.values.setLen(sibling.values.len - 1)
|
||||||
parent.keys[parentIdx - 1] = node.keys[0]
|
# Search is `key > sep → right` (left-max). After lending, sep is the
|
||||||
|
# left sibling's new max — not the borrowed key now sitting in `node`.
|
||||||
|
parent.keys[parentIdx - 1] = sibling.keys[^1]
|
||||||
else:
|
else:
|
||||||
# Borrow from internal sibling
|
# Borrow from internal sibling
|
||||||
let borrowKey = sibling.keys[^1]
|
let borrowKey = sibling.keys[^1]
|
||||||
@@ -231,7 +288,8 @@ proc borrowFromRight[K, V](node: BTreeNode[K, V], parent: BTreeNode[K, V], paren
|
|||||||
node.values.add(borrowVal)
|
node.values.add(borrowVal)
|
||||||
sibling.keys.delete(0)
|
sibling.keys.delete(0)
|
||||||
sibling.values.delete(0)
|
sibling.values.delete(0)
|
||||||
parent.keys[parentIdx] = sibling.keys[0]
|
# Borrowed key is now left's max; search must keep it on the left.
|
||||||
|
parent.keys[parentIdx] = node.keys[^1]
|
||||||
else:
|
else:
|
||||||
let borrowKey = sibling.keys[0]
|
let borrowKey = sibling.keys[0]
|
||||||
let borrowChild = sibling.children[0]
|
let borrowChild = sibling.children[0]
|
||||||
@@ -325,10 +383,14 @@ proc rebalanceAfterDelete[K, V](node: BTreeNode[K, V], root: var BTreeNode[K, V]
|
|||||||
mergeWithLeft(node, parent, parentIdx)
|
mergeWithLeft(node, parent, parentIdx)
|
||||||
elif hasRight:
|
elif hasRight:
|
||||||
mergeWithRight(node, parent, parentIdx)
|
mergeWithRight(node, parent, parentIdx)
|
||||||
|
else:
|
||||||
|
return
|
||||||
|
|
||||||
# Recursively rebalance parent if it fell below minimum
|
# Merge dropped a separator from parent — rebalance up the tree.
|
||||||
if parent == root and parent.keys.len == 0 and parent.children.len == 1:
|
if parent == root and parent.keys.len == 0 and parent.children.len == 1:
|
||||||
root = parent.children[0]
|
root = parent.children[0]
|
||||||
|
elif parent.keys.len < minKeysForLeaf(parent, order):
|
||||||
|
rebalanceAfterDelete(parent, root, order)
|
||||||
|
|
||||||
proc remove*[K, V](btree: var BTreeIndex[K, V], key: K, value: V) =
|
proc remove*[K, V](btree: var BTreeIndex[K, V], key: K, value: V) =
|
||||||
acquire(btree.lock)
|
acquire(btree.lock)
|
||||||
@@ -368,15 +430,17 @@ proc remove*[K, V](btree: var BTreeIndex[K, V], key: K, value: V) =
|
|||||||
else:
|
else:
|
||||||
# Internal node: recurse into child
|
# Internal node: recurse into child
|
||||||
let child = node.children[i]
|
let child = node.children[i]
|
||||||
let oldFirstKey = if child.keys.len > 0: child.keys[0] else: default(K)
|
|
||||||
let found = removeRec(child, root, order)
|
let found = removeRec(child, root, order)
|
||||||
if found:
|
if found:
|
||||||
# Update separator if child's first key changed.
|
# Rebalance first — merge/borrow rewrite parent separators.
|
||||||
# Separator node.keys[i-1] represents child's first key (for i > 0).
|
|
||||||
if i > 0 and child.keys.len > 0 and child.keys[0] != oldFirstKey:
|
|
||||||
node.keys[i - 1] = child.keys[0]
|
|
||||||
# Rebalance the child if needed
|
|
||||||
rebalanceAfterDelete(child, root, order)
|
rebalanceAfterDelete(child, root, order)
|
||||||
|
# Leaf children only: refresh left-max separator. Skip if merge
|
||||||
|
# already unlinked this child. Internal keys are promoted
|
||||||
|
# separators, not copies of child.keys[^1]; copying those (the
|
||||||
|
# naive H10 rewrite) makes sep < max(left) and misses keys.
|
||||||
|
if i < node.children.len and node.children[i] == child and
|
||||||
|
child.isLeaf and child.keys.len > 0 and i < node.keys.len:
|
||||||
|
node.keys[i] = child.keys[^1]
|
||||||
return found
|
return found
|
||||||
|
|
||||||
if removeRec(btree.root, btree.root, btree.order):
|
if removeRec(btree.root, btree.root, btree.order):
|
||||||
|
|||||||
@@ -125,13 +125,16 @@ proc compact*(cs: CompactionStrategy, level: int): CompactionResult =
|
|||||||
return cmp(b.timestamp, a.timestamp) # newest first
|
return cmp(b.timestamp, a.timestamp) # newest first
|
||||||
)
|
)
|
||||||
|
|
||||||
# Deduplicate: keep only the newest version of each key
|
# Deduplicate: keep only the newest version of each key.
|
||||||
|
# Use a haveLast flag — sentinel lastKey="" would skip the empty-string key.
|
||||||
var merged: seq[Entry] = @[]
|
var merged: seq[Entry] = @[]
|
||||||
var lastKey = ""
|
var lastKey = ""
|
||||||
|
var haveLast = false
|
||||||
for entry in allEntries:
|
for entry in allEntries:
|
||||||
if entry.key != lastKey:
|
if not haveLast or entry.key != lastKey:
|
||||||
merged.add(entry)
|
merged.add(entry)
|
||||||
lastKey = entry.key
|
lastKey = entry.key
|
||||||
|
haveLast = true
|
||||||
|
|
||||||
# Keep tombstones to prevent deleted keys from resurrecting in lower levels
|
# Keep tombstones to prevent deleted keys from resurrecting in lower levels
|
||||||
var final: seq[Entry] = @[]
|
var final: seq[Entry] = @[]
|
||||||
@@ -155,20 +158,15 @@ proc compact*(cs: CompactionStrategy, level: int): CompactionResult =
|
|||||||
createdAt: tables[^1].createdAt,
|
createdAt: tables[^1].createdAt,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Verify output SSTable before deleting sources
|
# Verify output SSTable before mutating the catalog. Input files stay on
|
||||||
|
# disk until the caller loads the output and writes MANIFEST (crash-safe
|
||||||
|
# order: output durable + catalog updated, then unlink inputs).
|
||||||
let (ok, msg) = verifySSTable(outputPath)
|
let (ok, msg) = verifySSTable(outputPath)
|
||||||
if not ok:
|
if not ok:
|
||||||
echo "[ERROR] Compaction output verification failed: ", msg
|
echo "[ERROR] Compaction output verification failed: ", msg
|
||||||
try: removeFile(outputPath) except CatchableError: discard
|
try: removeFile(outputPath) except CatchableError: discard
|
||||||
return CompactionResult()
|
return CompactionResult()
|
||||||
|
|
||||||
# Remove old SSTable files
|
|
||||||
for t in tables:
|
|
||||||
try:
|
|
||||||
removeFile(t.path)
|
|
||||||
except CatchableError as e:
|
|
||||||
echo "[WARN] Failed to remove old SSTable: ", t.path, ": ", e.msg
|
|
||||||
|
|
||||||
# Update level arrays
|
# Update level arrays
|
||||||
var newTables: seq[SSTableMeta] = @[]
|
var newTables: seq[SSTableMeta] = @[]
|
||||||
for t in cs.levels[level]:
|
for t in cs.levels[level]:
|
||||||
|
|||||||
@@ -696,6 +696,10 @@ proc newLSMTree*(
|
|||||||
var version: uint32 = 0
|
var version: uint32 = 0
|
||||||
if stream.readData(addr magic, 4) == 4 and magic == WALMagic:
|
if stream.readData(addr magic, 4) == 4 and magic == WALMagic:
|
||||||
if stream.readData(addr version, 4) == 4:
|
if stream.readData(addr version, 4) == 4:
|
||||||
|
# Cap per-record sizes to avoid multi-GiB alloc on torn/corrupt WAL.
|
||||||
|
# Kind must be a known WalEntryKind value (1..4) — out-of-range casts
|
||||||
|
# raise CaseStmtError (Defect) and crash the process.
|
||||||
|
const MaxWalRecordField = 64 * 1024 * 1024 # 64 MB
|
||||||
while not stream.atEnd():
|
while not stream.atEnd():
|
||||||
var kind: uint8 = 0
|
var kind: uint8 = 0
|
||||||
var timestamp: uint64 = 0
|
var timestamp: uint64 = 0
|
||||||
@@ -704,10 +708,20 @@ proc newLSMTree*(
|
|||||||
if stream.readData(addr kind, 1) != 1: break
|
if stream.readData(addr kind, 1) != 1: break
|
||||||
if stream.readData(addr timestamp, 8) != 8: break
|
if stream.readData(addr timestamp, 8) != 8: break
|
||||||
if stream.readData(addr keyLen, 4) != 4: break
|
if stream.readData(addr keyLen, 4) != 4: break
|
||||||
|
if keyLen.int > MaxWalRecordField:
|
||||||
|
echo "[WARN] WAL recovery: torn/corrupt record (keyLen=", keyLen, ") — stopping replay"
|
||||||
|
break
|
||||||
|
# Validate kind before allocating or branching (avoids CaseStmtError Defect)
|
||||||
|
if kind < uint8(wekPut) or kind > uint8(wekCommit):
|
||||||
|
echo "[WARN] WAL recovery: invalid entry kind ", kind, " — stopping replay"
|
||||||
|
break
|
||||||
var key = newString(keyLen.int)
|
var key = newString(keyLen.int)
|
||||||
if keyLen > 0:
|
if keyLen > 0:
|
||||||
if stream.readData(addr key[0], keyLen.int) != keyLen.int: break
|
if stream.readData(addr key[0], keyLen.int) != keyLen.int: break
|
||||||
if stream.readData(addr valLen, 4) != 4: break
|
if stream.readData(addr valLen, 4) != 4: break
|
||||||
|
if valLen.int > MaxWalRecordField:
|
||||||
|
echo "[WARN] WAL recovery: torn/corrupt record (valLen=", valLen, ") — stopping replay"
|
||||||
|
break
|
||||||
var value = newSeq[byte](valLen.int)
|
var value = newSeq[byte](valLen.int)
|
||||||
if valLen > 0:
|
if valLen > 0:
|
||||||
if stream.readData(addr value[0], valLen.int) != valLen.int: break
|
if stream.readData(addr value[0], valLen.int) != valLen.int: break
|
||||||
@@ -866,26 +880,36 @@ proc flushUnsafe(db: LSMTree) =
|
|||||||
if db.immutableMem.len == 0 and db.memTable.len == 0:
|
if db.immutableMem.len == 0 and db.memTable.len == 0:
|
||||||
return
|
return
|
||||||
|
|
||||||
# Flush immutable memtable if present, otherwise flush current memtable
|
# Flush immutable memtable if present, otherwise flush current memtable.
|
||||||
var toFlush = db.immutableMem
|
# Do NOT clear the source memtable until the SSTable is written — an IOError
|
||||||
if toFlush.len == 0:
|
# mid-write must leave the data still visible to live reads (WAL still has it).
|
||||||
toFlush = db.memTable
|
var flushingImmutable = false
|
||||||
db.memTable = newMemTable(db.memMaxSize)
|
var toFlush: MemTable
|
||||||
|
if db.immutableMem.len > 0:
|
||||||
|
toFlush = db.immutableMem
|
||||||
|
flushingImmutable = true
|
||||||
else:
|
else:
|
||||||
db.immutableMem = newMemTable(0)
|
toFlush = db.memTable
|
||||||
|
|
||||||
if toFlush.len == 0:
|
if toFlush.len == 0:
|
||||||
return
|
return
|
||||||
|
|
||||||
let path = db.dir / "sstables" / ($db.nextSSTableId & ".sst")
|
let path = db.dir / "sstables" / ($db.nextSSTableId & ".sst")
|
||||||
|
let sstId = db.nextSSTableId
|
||||||
inc db.nextSSTableId
|
inc db.nextSSTableId
|
||||||
|
|
||||||
# Sort once at flush time (O(n log n)) — put/get stay O(1)
|
# Sort once at flush time (O(n log n)) — put/get stay O(1)
|
||||||
var sst = writeSSTable(toFlush.sortedEntries(), path, level = 0)
|
var sst = writeSSTable(toFlush.sortedEntries(), path, level = 0)
|
||||||
sst.id = db.nextSSTableId - 1
|
sst.id = sstId
|
||||||
db.sstables.add(sst)
|
db.sstables.add(sst)
|
||||||
# SSTables are kept in insertion order (newest last) so getUnsafe can search newest-first
|
# SSTables are kept in insertion order (newest last) so getUnsafe can search newest-first
|
||||||
|
|
||||||
|
# Only now drop the in-memory copy — SSTable is durable on disk
|
||||||
|
if flushingImmutable:
|
||||||
|
db.immutableMem = newMemTable(0)
|
||||||
|
else:
|
||||||
|
db.memTable = newMemTable(db.memMaxSize)
|
||||||
|
|
||||||
# Update MANIFEST atomically
|
# Update MANIFEST atomically
|
||||||
inc db.manifestSequence
|
inc db.manifestSequence
|
||||||
try:
|
try:
|
||||||
@@ -930,27 +954,29 @@ proc checkpoint*(db: LSMTree) =
|
|||||||
## rotate WAL, and write MANIFEST. This provides a clean boundary
|
## rotate WAL, and write MANIFEST. This provides a clean boundary
|
||||||
## for online backup without stopping the server.
|
## for online backup without stopping the server.
|
||||||
acquireWrite(db.lock)
|
acquireWrite(db.lock)
|
||||||
|
try:
|
||||||
|
# Flush any pending immutable memtable first
|
||||||
|
if db.immutableMem.len > 0:
|
||||||
|
flushUnsafe(db)
|
||||||
|
|
||||||
# Flush any pending immutable memtable first
|
# Freeze current memtable so writes can continue on a new one
|
||||||
if db.immutableMem.len > 0:
|
if db.memTable.len > 0:
|
||||||
flushUnsafe(db)
|
db.immutableMem = db.memTable
|
||||||
|
db.memTable = newMemTable(db.memMaxSize)
|
||||||
|
|
||||||
# Freeze current memtable so writes can continue on a new one
|
# Flush the frozen memtable
|
||||||
if db.memTable.len > 0:
|
if db.immutableMem.len > 0:
|
||||||
db.immutableMem = db.memTable
|
flushUnsafe(db)
|
||||||
db.memTable = newMemTable(db.memMaxSize)
|
|
||||||
|
|
||||||
# Flush the frozen memtable
|
# Rotate WAL for a clean backup boundary
|
||||||
if db.immutableMem.len > 0:
|
acquire(db.walLock)
|
||||||
flushUnsafe(db)
|
try:
|
||||||
|
db.wal.maybeRotate()
|
||||||
# Rotate WAL for a clean backup boundary
|
db.wal.sync()
|
||||||
acquire(db.walLock)
|
finally:
|
||||||
db.wal.maybeRotate()
|
release(db.walLock)
|
||||||
db.wal.sync()
|
finally:
|
||||||
release(db.walLock)
|
releaseWrite(db.lock)
|
||||||
|
|
||||||
releaseWrite(db.lock)
|
|
||||||
|
|
||||||
proc close*(db: LSMTree) =
|
proc close*(db: LSMTree) =
|
||||||
acquireWrite(db.lock)
|
acquireWrite(db.lock)
|
||||||
|
|||||||
@@ -80,7 +80,8 @@ proc readAt*(mf: MmapFile, offset: int, size: int): seq[byte] =
|
|||||||
if mf.regions.len == 0:
|
if mf.regions.len == 0:
|
||||||
return @[]
|
return @[]
|
||||||
let region = mf.regions[0]
|
let region = mf.regions[0]
|
||||||
if offset < 0 or size < 0 or offset + size > region.size:
|
# overflow-safe bound: offset > size - length (not offset + length > size)
|
||||||
|
if offset < 0 or size < 0 or size > region.size or offset > region.size - size:
|
||||||
return @[]
|
return @[]
|
||||||
result = newSeq[byte](size)
|
result = newSeq[byte](size)
|
||||||
copyMem(addr result[0], unsafeAddr region.data[offset], size)
|
copyMem(addr result[0], unsafeAddr region.data[offset], size)
|
||||||
@@ -91,21 +92,24 @@ proc readByte*(mf: MmapFile, offset: int): byte =
|
|||||||
return mf.regions[0].data[offset]
|
return mf.regions[0].data[offset]
|
||||||
|
|
||||||
proc readUint32*(mf: MmapFile, offset: int): uint32 =
|
proc readUint32*(mf: MmapFile, offset: int): uint32 =
|
||||||
if mf.regions.len == 0 or offset < 0 or offset + 4 > mf.regions[0].size:
|
if mf.regions.len == 0 or offset < 0 or 4 > mf.regions[0].size or
|
||||||
|
offset > mf.regions[0].size - 4:
|
||||||
return 0
|
return 0
|
||||||
var val: uint32
|
var val: uint32
|
||||||
copyMem(addr val, unsafeAddr mf.regions[0].data[offset], 4)
|
copyMem(addr val, unsafeAddr mf.regions[0].data[offset], 4)
|
||||||
return val
|
return val
|
||||||
|
|
||||||
proc readUint64*(mf: MmapFile, offset: int): uint64 =
|
proc readUint64*(mf: MmapFile, offset: int): uint64 =
|
||||||
if mf.regions.len == 0 or offset < 0 or offset + 8 > mf.regions[0].size:
|
if mf.regions.len == 0 or offset < 0 or 8 > mf.regions[0].size or
|
||||||
|
offset > mf.regions[0].size - 8:
|
||||||
return 0
|
return 0
|
||||||
var val: uint64
|
var val: uint64
|
||||||
copyMem(addr val, unsafeAddr mf.regions[0].data[offset], 8)
|
copyMem(addr val, unsafeAddr mf.regions[0].data[offset], 8)
|
||||||
return val
|
return val
|
||||||
|
|
||||||
proc readString*(mf: MmapFile, offset: int, size: int): string =
|
proc readString*(mf: MmapFile, offset: int, size: int): string =
|
||||||
if mf.regions.len == 0 or offset < 0 or size < 0 or offset + size > mf.regions[0].size:
|
if mf.regions.len == 0 or offset < 0 or size < 0 or
|
||||||
|
size > mf.regions[0].size or offset > mf.regions[0].size - size:
|
||||||
return ""
|
return ""
|
||||||
result = newString(size)
|
result = newString(size)
|
||||||
copyMem(addr result[0], unsafeAddr mf.regions[0].data[offset], size)
|
copyMem(addr result[0], unsafeAddr mf.regions[0].data[offset], size)
|
||||||
|
|||||||
@@ -68,6 +68,7 @@ proc scanWAL*(rec: CrashRecovery): seq[RecoveredEntry] =
|
|||||||
var txnId: uint64 = 0
|
var txnId: uint64 = 0
|
||||||
var entryCount = 0
|
var entryCount = 0
|
||||||
|
|
||||||
|
const MaxWalRecordField = 64 * 1024 * 1024 # 64 MB
|
||||||
while not stream.atEnd():
|
while not stream.atEnd():
|
||||||
var kind: uint8 = 0
|
var kind: uint8 = 0
|
||||||
var timestamp: uint64 = 0
|
var timestamp: uint64 = 0
|
||||||
@@ -77,12 +78,15 @@ proc scanWAL*(rec: CrashRecovery): seq[RecoveredEntry] =
|
|||||||
if stream.readData(addr kind, 1) != 1: break
|
if stream.readData(addr kind, 1) != 1: break
|
||||||
if stream.readData(addr timestamp, 8) != 8: break
|
if stream.readData(addr timestamp, 8) != 8: break
|
||||||
if stream.readData(addr keyLen, 4) != 4: break
|
if stream.readData(addr keyLen, 4) != 4: break
|
||||||
|
if keyLen.int > MaxWalRecordField: break
|
||||||
|
if kind < uint8(wekPut) or kind > uint8(wekCommit): break
|
||||||
|
|
||||||
var key = newString(keyLen.int)
|
var key = newString(keyLen.int)
|
||||||
if keyLen > 0:
|
if keyLen > 0:
|
||||||
if stream.readData(addr key[0], keyLen.int) != keyLen.int: break
|
if stream.readData(addr key[0], keyLen.int) != keyLen.int: break
|
||||||
|
|
||||||
if stream.readData(addr valLen, 4) != 4: break
|
if stream.readData(addr valLen, 4) != 4: break
|
||||||
|
if valLen.int > MaxWalRecordField: break
|
||||||
var value = newSeq[byte](valLen.int)
|
var value = newSeq[byte](valLen.int)
|
||||||
if valLen > 0:
|
if valLen > 0:
|
||||||
if stream.readData(addr value[0], valLen.int) != valLen.int: break
|
if stream.readData(addr value[0], valLen.int) != valLen.int: break
|
||||||
|
|||||||
@@ -326,8 +326,9 @@ proc rewriteLive*(wal: var WriteAheadLog,
|
|||||||
|
|
||||||
if wal.stream != nil:
|
if wal.stream != nil:
|
||||||
wal.stream.close()
|
wal.stream.close()
|
||||||
if fileExists(wal.path):
|
wal.stream = nil
|
||||||
removeFile(wal.path)
|
# Atomic replace: moveFile overwrites the destination on POSIX rename(2).
|
||||||
|
# Do not removeFile first — a crash between unlink and rename would lose the WAL.
|
||||||
moveFile(tmpPath, wal.path)
|
moveFile(tmpPath, wal.path)
|
||||||
wal.stream = newFileStream(wal.path, fmAppend)
|
wal.stream = newFileStream(wal.path, fmAppend)
|
||||||
if wal.stream == nil:
|
if wal.stream == nil:
|
||||||
@@ -363,6 +364,7 @@ proc readEntries*(walPath: string, untilTimestamp: uint64 = 0): seq[WalEntry] =
|
|||||||
if s.readData(addr magic, 4) != 4: return
|
if s.readData(addr magic, 4) != 4: return
|
||||||
if s.readData(addr version, 4) != 4: return
|
if s.readData(addr version, 4) != 4: return
|
||||||
if magic != WALMagic: return
|
if magic != WALMagic: return
|
||||||
|
const MaxWalRecordField = 64 * 1024 * 1024 # 64 MB
|
||||||
while not s.atEnd:
|
while not s.atEnd:
|
||||||
var kind: uint8
|
var kind: uint8
|
||||||
if s.readData(addr kind, 1) != 1: break
|
if s.readData(addr kind, 1) != 1: break
|
||||||
@@ -372,11 +374,14 @@ proc readEntries*(walPath: string, untilTimestamp: uint64 = 0): seq[WalEntry] =
|
|||||||
break
|
break
|
||||||
var keyLen: uint32
|
var keyLen: uint32
|
||||||
if s.readData(addr keyLen, 4) != 4: break
|
if s.readData(addr keyLen, 4) != 4: break
|
||||||
|
if keyLen.int > MaxWalRecordField: break
|
||||||
|
if kind < uint8(wekPut) or kind > uint8(wekCommit): break
|
||||||
var key = newSeq[byte](keyLen)
|
var key = newSeq[byte](keyLen)
|
||||||
if keyLen > 0:
|
if keyLen > 0:
|
||||||
if s.readData(addr key[0], int(keyLen)) != int(keyLen): break
|
if s.readData(addr key[0], int(keyLen)) != int(keyLen): break
|
||||||
var valLen: uint32
|
var valLen: uint32
|
||||||
if s.readData(addr valLen, 4) != 4: break
|
if s.readData(addr valLen, 4) != 4: break
|
||||||
|
if valLen.int > MaxWalRecordField: break
|
||||||
var value = newSeq[byte](valLen)
|
var value = newSeq[byte](valLen)
|
||||||
if valLen > 0:
|
if valLen > 0:
|
||||||
if s.readData(addr value[0], int(valLen)) != int(valLen): break
|
if s.readData(addr value[0], int(valLen)) != int(valLen): break
|
||||||
|
|||||||
+40
-22
@@ -41,27 +41,18 @@ proc newCompactionManager*(db: LSMTree): CompactionManager =
|
|||||||
result.strategy.rebuildFromLSM(db)
|
result.strategy.rebuildFromLSM(db)
|
||||||
|
|
||||||
proc applyCompactionResult(db: LSMTree, result: compaction.CompactionResult) =
|
proc applyCompactionResult(db: LSMTree, result: compaction.CompactionResult) =
|
||||||
## Apply compaction output under the caller's lock: update sstables + MANIFEST.
|
## Crash-safe apply: load output while inputs still exist, swap catalog,
|
||||||
## On Linux, compact may already have unlinked inputs; we still close our mmaps.
|
## write MANIFEST, then unlink inputs. A crash before MANIFEST leaves the
|
||||||
|
## old set intact (orphan output is ignored); a crash after MANIFEST leaves
|
||||||
|
## at worst unlinked-but-closed input files.
|
||||||
if result.outputTables.len == 0:
|
if result.outputTables.len == 0:
|
||||||
return
|
return
|
||||||
|
|
||||||
var newSSTables: seq[SSTable] = @[]
|
var loaded: seq[SSTable] = @[]
|
||||||
var removedPaths = initTable[string, bool]()
|
|
||||||
for t in result.inputTables:
|
|
||||||
removedPaths[t.path] = true
|
|
||||||
for sst in db.sstables.mitems:
|
|
||||||
if sst.path notin removedPaths:
|
|
||||||
newSSTables.add(sst)
|
|
||||||
else:
|
|
||||||
# Drop mmap after compact unlinked the path (fd remains valid until close)
|
|
||||||
sst.close()
|
|
||||||
|
|
||||||
for meta in result.outputTables:
|
for meta in result.outputTables:
|
||||||
try:
|
try:
|
||||||
var sst = loadSSTable(meta.path)
|
var sst = loadSSTable(meta.path)
|
||||||
let name = splitFile(meta.path).name
|
let name = splitFile(meta.path).name
|
||||||
# Prefer numeric id from filename; otherwise allocate
|
|
||||||
let parsed = try: parseInt(name) except CatchableError: -1
|
let parsed = try: parseInt(name) except CatchableError: -1
|
||||||
if parsed >= 0:
|
if parsed >= 0:
|
||||||
sst.id = parsed
|
sst.id = parsed
|
||||||
@@ -69,11 +60,30 @@ proc applyCompactionResult(db: LSMTree, result: compaction.CompactionResult) =
|
|||||||
sst.id = db.nextSSTableId
|
sst.id = db.nextSSTableId
|
||||||
inc db.nextSSTableId
|
inc db.nextSSTableId
|
||||||
sst.level = meta.level
|
sst.level = meta.level
|
||||||
newSSTables.add(sst)
|
loaded.add(sst)
|
||||||
db.nextSSTableId = max(db.nextSSTableId, sst.id + 1)
|
db.nextSSTableId = max(db.nextSSTableId, sst.id + 1)
|
||||||
except CatchableError as e:
|
except CatchableError as e:
|
||||||
warn("Compaction output SSTable failed to load: " & meta.path & " — " & e.msg)
|
warn("Compaction output SSTable failed to load: " & meta.path & " — " & e.msg)
|
||||||
|
for s in loaded.mitems:
|
||||||
|
s.close()
|
||||||
|
for m in result.outputTables:
|
||||||
|
try: removeFile(m.path) except CatchableError: discard
|
||||||
|
return
|
||||||
|
|
||||||
|
var removedPaths = initTable[string, bool]()
|
||||||
|
for t in result.inputTables:
|
||||||
|
removedPaths[t.path] = true
|
||||||
|
|
||||||
|
var newSSTables: seq[SSTable] = @[]
|
||||||
|
var toDrop: seq[SSTable] = @[]
|
||||||
|
for sst in db.sstables.mitems:
|
||||||
|
if sst.path notin removedPaths:
|
||||||
|
newSSTables.add(sst)
|
||||||
|
else:
|
||||||
|
toDrop.add(sst)
|
||||||
|
|
||||||
|
for sst in loaded:
|
||||||
|
newSSTables.add(sst)
|
||||||
newSSTables.sort(proc(a, b: SSTable): int = cmp(a.id, b.id))
|
newSSTables.sort(proc(a, b: SSTable): int = cmp(a.id, b.id))
|
||||||
db.sstables = newSSTables
|
db.sstables = newSSTables
|
||||||
db.needsCompaction = db.countL0() >= L0CompactionTrigger
|
db.needsCompaction = db.countL0() >= L0CompactionTrigger
|
||||||
@@ -84,6 +94,13 @@ proc applyCompactionResult(db: LSMTree, result: compaction.CompactionResult) =
|
|||||||
except CatchableError as e:
|
except CatchableError as e:
|
||||||
warn("Failed to write MANIFEST after compaction: " & e.msg)
|
warn("Failed to write MANIFEST after compaction: " & e.msg)
|
||||||
|
|
||||||
|
for sst in toDrop.mitems:
|
||||||
|
try:
|
||||||
|
removeFile(sst.path)
|
||||||
|
except CatchableError as e:
|
||||||
|
warn("Failed to remove compacted SSTable: " & sst.path & " — " & e.msg)
|
||||||
|
sst.close()
|
||||||
|
|
||||||
proc compact*(cm: CompactionManager) =
|
proc compact*(cm: CompactionManager) =
|
||||||
# Gate first (cross-thread), then per-DB write lock
|
# Gate first (cross-thread), then per-DB write lock
|
||||||
withStorageGate:
|
withStorageGate:
|
||||||
@@ -439,19 +456,20 @@ proc main() =
|
|||||||
echo "[raft] Snapshot restore failed: ", e.msg
|
echo "[raft] Snapshot restore failed: ", e.msg
|
||||||
result = false
|
result = false
|
||||||
|
|
||||||
# Leader InstallSnapshot send: archive the default DB's data directory
|
# Leader InstallSnapshot send: tar the default DB's data directory into the
|
||||||
# into the path raft picks (dataDir/raft/snap_out_<snapId>.tar.gz). Like
|
# path raft picks (dataDir/raft/snap_out_<snapId>.tar). The tar runs here on
|
||||||
# restoreSnapshot this runs on the raft event loop and blocks on disk I/O
|
# the raft event loop under the storage gate; sendSnapshot then gzips it on
|
||||||
# (tar+gzip); snapshot sends are rare, so we accept the stall.
|
# a worker thread off the loop (gzipFileAsync) so heartbeats keep flowing.
|
||||||
raftNode.buildSnapshot = proc(destPath: string): bool {.gcsafe.} =
|
raftNode.buildSnapshot = proc(destPath: string): bool {.gcsafe.} =
|
||||||
echo "[raft] Building snapshot archive ", destPath
|
echo "[raft] Building snapshot tar ", destPath
|
||||||
{.cast(gcsafe).}:
|
{.cast(gcsafe).}:
|
||||||
# Hold the storage gate while tarring the data dir so a concurrent
|
# Hold the storage gate while tarring the data dir so a concurrent
|
||||||
# memtable flush (HTTP /query path) cannot write an SSTable
|
# memtable flush (HTTP /query path) cannot write an SSTable
|
||||||
# mid-archive.
|
# mid-archive. Compression happens later, off the gate (see
|
||||||
|
# sendSnapshot), so it neither stalls the loop nor blocks applies.
|
||||||
withStorageGate:
|
withStorageGate:
|
||||||
try:
|
try:
|
||||||
result = backupDataDir(defaultDbDir, destPath)
|
result = tarDataDir(defaultDbDir, destPath)
|
||||||
except CatchableError as e:
|
except CatchableError as e:
|
||||||
echo "[raft] Snapshot build failed: ", e.msg
|
echo "[raft] Snapshot build failed: ", e.msg
|
||||||
result = false
|
result = false
|
||||||
|
|||||||
@@ -7,7 +7,16 @@ import ../src/barabadb/query/exec/params
|
|||||||
import ../src/barabadb/query/exec/dml
|
import ../src/barabadb/query/exec/dml
|
||||||
import ../src/barabadb/core/types
|
import ../src/barabadb/core/types
|
||||||
import ../src/barabadb/core/config
|
import ../src/barabadb/core/config
|
||||||
|
import ../src/barabadb/core/replication
|
||||||
|
import ../src/barabadb/core/disttxn
|
||||||
|
import ../src/barabadb/core/websocket
|
||||||
|
import ../src/barabadb/protocol/auth
|
||||||
|
import ../src/barabadb/protocol/scram
|
||||||
import ../src/barabadb/storage/lsm
|
import ../src/barabadb/storage/lsm
|
||||||
|
import ../src/barabadb/storage/compaction
|
||||||
|
import ../src/barabadb/storage/btree
|
||||||
|
import std/random
|
||||||
|
import std/sets
|
||||||
|
|
||||||
const testDir = "/tmp/baradb_bugfix_test"
|
const testDir = "/tmp/baradb_bugfix_test"
|
||||||
|
|
||||||
@@ -552,3 +561,388 @@ suite "Raft TLS config":
|
|||||||
check cfg.raftTlsKeyFile == "/tmp/raft.key"
|
check cfg.raftTlsKeyFile == "/tmp/raft.key"
|
||||||
check cfg.raftTlsCaFile == "/tmp/raft-ca.crt"
|
check cfg.raftTlsCaFile == "/tmp/raft-ca.crt"
|
||||||
check cfg.raftTlsVerifyPeer == true
|
check cfg.raftTlsVerifyPeer == true
|
||||||
|
|
||||||
|
|
||||||
|
suite "Legacy REP payload encoding — empty value is not a delete":
|
||||||
|
|
||||||
|
test "PK-only put (empty value) round-trips as a put, not a delete":
|
||||||
|
## Regression: the legacy REP receiver used to infer a delete from an empty
|
||||||
|
## value, so PK-only rows (empty LSM value) vanished on the replica.
|
||||||
|
let decoded = decodeRepPayload(encodeRepPayload(false, "pkonly.id=3", @[]))
|
||||||
|
check decoded.op == ropPut
|
||||||
|
check decoded.key == "pkonly.id=3"
|
||||||
|
check decoded.value.len == 0
|
||||||
|
|
||||||
|
test "delete round-trips as a delete":
|
||||||
|
let decoded = decodeRepPayload(encodeRepPayload(true, "users.id=1", @[]))
|
||||||
|
check decoded.op == ropDelete
|
||||||
|
check decoded.key == "users.id=1"
|
||||||
|
check decoded.value.len == 0
|
||||||
|
|
||||||
|
test "put with a non-empty value preserves the value bytes":
|
||||||
|
let decoded = decodeRepPayload(
|
||||||
|
encodeRepPayload(false, "users.id=1", cast[seq[byte]]("bob")))
|
||||||
|
check decoded.op == ropPut
|
||||||
|
check decoded.key == "users.id=1"
|
||||||
|
check cast[string](decoded.value) == "bob"
|
||||||
|
|
||||||
|
test "value containing a null byte survives the round-trip":
|
||||||
|
## Decode splits on the FIRST null (the key/value separator) only.
|
||||||
|
let value = @[byte('a'), byte(0), byte('b')]
|
||||||
|
let decoded = decodeRepPayload(encodeRepPayload(false, "k", value))
|
||||||
|
check decoded.op == ropPut
|
||||||
|
check decoded.key == "k"
|
||||||
|
check decoded.value == value
|
||||||
|
|
||||||
|
test "empty or untagged payloads decode as invalid, not delete":
|
||||||
|
check decodeRepPayload(@[]).op == ropInvalid
|
||||||
|
check decodeRepPayload(cast[seq[byte]]("Xfoo")).op == ropInvalid
|
||||||
|
|
||||||
|
|
||||||
|
suite "Query operator correctness — audit batch 1":
|
||||||
|
|
||||||
|
test "power operator ** evaluates, not lowered to equality":
|
||||||
|
## Regression: bkPow used to fall through to `else: irOp = irEq`, so
|
||||||
|
## `2 ** 3` evaluated as `2 = 3` (false) instead of 8.
|
||||||
|
var ctx = setupCtx()
|
||||||
|
defer: teardown(ctx)
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (1, 'a')"))
|
||||||
|
let r = executeQuery(ctx, parse("SELECT 2 ** 3 AS x FROM users"))
|
||||||
|
check r.success
|
||||||
|
check r.rows.len == 1
|
||||||
|
check parseFloat(valueToString(r.rows[0]["x"])) == 8.0
|
||||||
|
|
||||||
|
test "concat operator ++ concatenates strings":
|
||||||
|
## Regression: bkConcat also fell through to irEq, so `'a' ++ 'b'`
|
||||||
|
## evaluated as `'a' = 'b'` (false) instead of "ab".
|
||||||
|
var ctx = setupCtx()
|
||||||
|
defer: teardown(ctx)
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (1, 'a')"))
|
||||||
|
let r = executeQuery(ctx, parse("SELECT 'a' ++ 'b' AS x FROM users"))
|
||||||
|
check r.success
|
||||||
|
check r.rows.len == 1
|
||||||
|
check valueToString(r.rows[0]["x"]) == "ab"
|
||||||
|
|
||||||
|
test "!= is the complement of = for numerically equal values":
|
||||||
|
## Regression: irNeq short-circuited on string inequality, so `1 != 1.0`
|
||||||
|
## was true while `1 = 1.0` was also true (not complements).
|
||||||
|
var ctx = setupCtx()
|
||||||
|
defer: teardown(ctx)
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (1, 'alice')"))
|
||||||
|
let eq = executeQuery(ctx, parse("SELECT * FROM users WHERE id = 1.0"))
|
||||||
|
let neq = executeQuery(ctx, parse("SELECT * FROM users WHERE id != 1.0"))
|
||||||
|
check eq.rows.len == 1 # 1 = 1.0 -> true
|
||||||
|
check neq.rows.len == 0 # 1 != 1.0 -> false (old bug returned the row)
|
||||||
|
|
||||||
|
|
||||||
|
suite "Query correctness — audit batch 2":
|
||||||
|
|
||||||
|
test "COUNT(DISTINCT) deduplicates values":
|
||||||
|
## Regression: funcDistinct was parsed but never copied to aggDistinct /
|
||||||
|
## never consulted during aggregation.
|
||||||
|
var ctx = setupCtx()
|
||||||
|
defer: teardown(ctx)
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (1, 'alice')"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (2, 'bob')"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (3, 'alice')"))
|
||||||
|
let r = executeQuery(ctx, parse("SELECT COUNT(DISTINCT name) AS c FROM users"))
|
||||||
|
check r.success
|
||||||
|
check r.rows.len == 1
|
||||||
|
check valueToString(r.rows[0]["c"]) == "2"
|
||||||
|
|
||||||
|
test "SUM(DISTINCT) sums unique values only":
|
||||||
|
var ctx = setupCtx()
|
||||||
|
defer: teardown(ctx)
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (1, 'a')"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (2, 'b')"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (5, 'c')"))
|
||||||
|
# ids 1, 2, 5 — insert another row with id-like values via a number col
|
||||||
|
discard executeQuery(ctx, parse("CREATE TABLE nums (id INTEGER PRIMARY KEY, n INTEGER)"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO nums (id, n) VALUES (1, 10)"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO nums (id, n) VALUES (2, 10)"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO nums (id, n) VALUES (3, 20)"))
|
||||||
|
let r = executeQuery(ctx, parse("SELECT SUM(DISTINCT n) AS s FROM nums"))
|
||||||
|
check r.success
|
||||||
|
check r.rows.len == 1
|
||||||
|
check parseFloat(valueToString(r.rows[0]["s"])) == 30.0
|
||||||
|
|
||||||
|
test "UNION deduplicates without KeyError":
|
||||||
|
## Regression: set-op dedup used row["$value"] which projected rows lack.
|
||||||
|
var ctx = setupCtx()
|
||||||
|
defer: teardown(ctx)
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (1, 'alice')"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (2, 'bob')"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (3, 'alice')"))
|
||||||
|
let r = executeQuery(ctx, parse(
|
||||||
|
"SELECT name FROM users WHERE id = 1 UNION SELECT name FROM users WHERE id = 3"))
|
||||||
|
check r.success
|
||||||
|
check r.rows.len == 1
|
||||||
|
check valueToString(r.rows[0]["name"]) == "alice"
|
||||||
|
|
||||||
|
test "INTERSECT returns common rows":
|
||||||
|
var ctx = setupCtx()
|
||||||
|
defer: teardown(ctx)
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (1, 'alice')"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (2, 'bob')"))
|
||||||
|
let r = executeQuery(ctx, parse(
|
||||||
|
"SELECT name FROM users WHERE id <= 2 INTERSECT SELECT name FROM users WHERE id = 1"))
|
||||||
|
check r.success
|
||||||
|
check r.rows.len == 1
|
||||||
|
check valueToString(r.rows[0]["name"]) == "alice"
|
||||||
|
|
||||||
|
test "EXCEPT removes right-side rows":
|
||||||
|
var ctx = setupCtx()
|
||||||
|
defer: teardown(ctx)
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (1, 'alice')"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (2, 'bob')"))
|
||||||
|
let r = executeQuery(ctx, parse(
|
||||||
|
"SELECT name FROM users EXCEPT SELECT name FROM users WHERE id = 1"))
|
||||||
|
check r.success
|
||||||
|
check r.rows.len == 1
|
||||||
|
check valueToString(r.rows[0]["name"]) == "bob"
|
||||||
|
|
||||||
|
test "MERGE WHEN MATCHED THEN DELETE removes the row":
|
||||||
|
## Regression: mergeMatchedDelete was parsed but never executed.
|
||||||
|
var ctx = setupCtx()
|
||||||
|
defer: teardown(ctx)
|
||||||
|
discard executeQuery(ctx, parse("CREATE TABLE inv (id INTEGER PRIMARY KEY, qty INTEGER)"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO inv (id, qty) VALUES (1, 10)"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO inv (id, qty) VALUES (2, 20)"))
|
||||||
|
discard executeQuery(ctx, parse("CREATE TABLE deltas (id INTEGER PRIMARY KEY, qty INTEGER)"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO deltas (id, qty) VALUES (1, 0)"))
|
||||||
|
let r = executeQuery(ctx, parse("""
|
||||||
|
MERGE INTO inv AS t
|
||||||
|
USING deltas AS s
|
||||||
|
ON t.id = s.id
|
||||||
|
WHEN MATCHED THEN DELETE
|
||||||
|
"""))
|
||||||
|
check r.success
|
||||||
|
check r.affectedRows >= 1
|
||||||
|
let left = executeQuery(ctx, parse("SELECT id FROM inv ORDER BY id"))
|
||||||
|
check left.success
|
||||||
|
check left.rows.len == 1
|
||||||
|
check valueToString(left.rows[0]["id"]) == "2"
|
||||||
|
|
||||||
|
test "semi-sync writeLsn returns 0 when replicas do not ack":
|
||||||
|
var rm = newReplicationManager(rmSemiSync, syncCount = 1)
|
||||||
|
rm.addReplica(newReplica("r1", "10.0.0.1", 9472))
|
||||||
|
rm.connectReplica("r1")
|
||||||
|
let lsn = rm.writeLsn(@[1'u8, 2, 3])
|
||||||
|
check lsn == 0
|
||||||
|
|
||||||
|
|
||||||
|
suite "Audit batch 3 — remaining 2026-08 findings":
|
||||||
|
|
||||||
|
test "TLS CA file auto-enables peer verify":
|
||||||
|
putEnv("BARADB_TLS_CA_FILE", "/tmp/ca.crt")
|
||||||
|
defer: delEnv("BARADB_TLS_CA_FILE")
|
||||||
|
var cfg = defaultConfig()
|
||||||
|
loadConfigFromEnv(cfg)
|
||||||
|
check cfg.tlsCaFile == "/tmp/ca.crt"
|
||||||
|
check cfg.tlsVerifyPeer == true
|
||||||
|
|
||||||
|
test "explicit BARADB_TLS_VERIFY_PEER=false wins over CA auto-enable":
|
||||||
|
putEnv("BARADB_TLS_CA_FILE", "/tmp/ca.crt")
|
||||||
|
putEnv("BARADB_TLS_VERIFY_PEER", "false")
|
||||||
|
defer:
|
||||||
|
delEnv("BARADB_TLS_CA_FILE")
|
||||||
|
delEnv("BARADB_TLS_VERIFY_PEER")
|
||||||
|
var cfg = defaultConfig()
|
||||||
|
loadConfigFromEnv(cfg)
|
||||||
|
check cfg.tlsVerifyPeer == false
|
||||||
|
|
||||||
|
test "production TLS without verify is rejected":
|
||||||
|
putEnv("BARADB_ENV", "production")
|
||||||
|
defer: delEnv("BARADB_ENV")
|
||||||
|
var cfg = defaultConfig()
|
||||||
|
cfg.authEnabled = true
|
||||||
|
cfg.jwtSecret = "a".repeat(32)
|
||||||
|
cfg.tlsEnabled = true
|
||||||
|
var msg = ""
|
||||||
|
try:
|
||||||
|
validateProductionConfig(cfg)
|
||||||
|
except ValueError as e:
|
||||||
|
msg = e.msg
|
||||||
|
check "TLS" in msg or "verify" in msg.toLower()
|
||||||
|
|
||||||
|
test "OFFSET without LIMIT returns remaining rows":
|
||||||
|
var ctx = setupCtx()
|
||||||
|
defer: teardown(ctx)
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (1, 'a')"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (2, 'b')"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (3, 'c')"))
|
||||||
|
let r = executeQuery(ctx, parse("SELECT id FROM users ORDER BY id OFFSET 1"))
|
||||||
|
check r.success
|
||||||
|
check r.rows.len == 2
|
||||||
|
check valueToString(r.rows[0]["id"]) == "2"
|
||||||
|
check valueToString(r.rows[1]["id"]) == "3"
|
||||||
|
|
||||||
|
test "LIMIT 0 returns no rows":
|
||||||
|
var ctx = setupCtx()
|
||||||
|
defer: teardown(ctx)
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (1, 'a')"))
|
||||||
|
let r = executeQuery(ctx, parse("SELECT id FROM users LIMIT 0"))
|
||||||
|
check r.success
|
||||||
|
check r.rows.len == 0
|
||||||
|
|
||||||
|
test "negative LIMIT is clamped to empty, not IndexDefect":
|
||||||
|
var ctx = setupCtx()
|
||||||
|
defer: teardown(ctx)
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (1, 'a')"))
|
||||||
|
let r = executeQuery(ctx, parse("SELECT id FROM users LIMIT -5"))
|
||||||
|
check r.success
|
||||||
|
check r.rows.len == 0
|
||||||
|
|
||||||
|
test "WebSocket decode rejects unmasked client frames":
|
||||||
|
let raw = encodeFrame(0x1, "SUBSCRIBE t", masked = false)
|
||||||
|
let (_, n) = decodeFrame(raw)
|
||||||
|
check n < 0
|
||||||
|
|
||||||
|
test "WebSocket decode accepts masked client frames":
|
||||||
|
let raw = encodeFrame(0x1, "SUBSCRIBE t", masked = true,
|
||||||
|
maskKey = [1'u8, 2, 3, 4])
|
||||||
|
let (frame, n) = decodeFrame(raw)
|
||||||
|
check n > 0
|
||||||
|
check frame.masked
|
||||||
|
check frame.payload == "SUBSCRIBE t"
|
||||||
|
|
||||||
|
test "WebSocket decode rejects oversized frame length":
|
||||||
|
var raw = newString(10)
|
||||||
|
raw[0] = char(0x81) # FIN + text
|
||||||
|
raw[1] = char(0xFF) # MASK + 127
|
||||||
|
# 8-byte length = 2 GiB
|
||||||
|
raw[2] = char(0)
|
||||||
|
raw[3] = char(0)
|
||||||
|
raw[4] = char(0)
|
||||||
|
raw[5] = char(0)
|
||||||
|
raw[6] = char(0x80)
|
||||||
|
raw[7] = char(0)
|
||||||
|
raw[8] = char(0)
|
||||||
|
raw[9] = char(0)
|
||||||
|
let (_, n) = decodeFrame(raw)
|
||||||
|
check n < 0
|
||||||
|
|
||||||
|
test "disttxn prepare against refused port fails closed":
|
||||||
|
var tm = newDistTxnManager()
|
||||||
|
let txn = tm.beginTransaction("coord")
|
||||||
|
txn.addParticipant("n1", "127.0.0.1", 1)
|
||||||
|
check txn.prepare() == false
|
||||||
|
check txn.isAborted
|
||||||
|
|
||||||
|
test "compact leaves input files on disk for catalog apply":
|
||||||
|
let testDir = "/tmp/baradb_bugfix_compact_order"
|
||||||
|
removeDir(testDir)
|
||||||
|
var db = newLSMTree(testDir, 128)
|
||||||
|
defer:
|
||||||
|
db.close()
|
||||||
|
removeDir(testDir)
|
||||||
|
for round in 0 ..< L0CompactionTrigger:
|
||||||
|
db.put("r" & $round, cast[seq[byte]]("v" & $round))
|
||||||
|
db.flush()
|
||||||
|
var cs = newCompactionStrategy(testDir)
|
||||||
|
cs.rebuildFromLSM(db)
|
||||||
|
let cr = cs.compact(0)
|
||||||
|
check cr.outputTables.len == 1
|
||||||
|
check fileExists(cr.outputTables[0].path)
|
||||||
|
for t in cr.inputTables:
|
||||||
|
check fileExists(t.path)
|
||||||
|
|
||||||
|
test "SCRAM unknown user fails without leaking existence":
|
||||||
|
var am = newAuthManager()
|
||||||
|
am.registerScramUser("alice", "wonderland")
|
||||||
|
var msg = ""
|
||||||
|
try:
|
||||||
|
discard am.startScram("n,,n=eve,r=abcnonceabcnonceabcn")
|
||||||
|
except ValueError as e:
|
||||||
|
msg = e.msg
|
||||||
|
check msg == "Authentication failed"
|
||||||
|
|
||||||
|
test "SCRAM rejects mismatched channel binding":
|
||||||
|
var am = newAuthManager()
|
||||||
|
am.registerScramUser("alice", "wonderland")
|
||||||
|
let clientNonce = generateNonce()
|
||||||
|
let clientFirst = "n,,n=alice,r=" & clientNonce
|
||||||
|
let serverFirst = am.startScram(clientFirst)
|
||||||
|
var combinedNonce = ""
|
||||||
|
for part in serverFirst.split(","):
|
||||||
|
if part.startsWith("r="): combinedNonce = part[2..^1]
|
||||||
|
let (ok, err) = am.finishScram("c=AAAA,r=" & combinedNonce & ",p=AA")
|
||||||
|
check ok == false
|
||||||
|
check err == "e=channel-bindings-dont-match"
|
||||||
|
|
||||||
|
test "SCRAM expected cbind for gs2 n is biws":
|
||||||
|
check expectedChannelBinding("n") == "biws"
|
||||||
|
|
||||||
|
|
||||||
|
suite "Audit batch 4 — B-tree separator (H10) and NULL equality (L4)":
|
||||||
|
|
||||||
|
test "B-tree separators stay valid after interleaved insert/remove":
|
||||||
|
## Search routes with `key > sep → right`. After delete, removeRec used to
|
||||||
|
## copy the right child's first key into the separator (right-min), which
|
||||||
|
## breaks max(left) <= sep < min(right).
|
||||||
|
var rng = initRand(20260828)
|
||||||
|
var btree = newBTreeIndex[int, string](order = 5)
|
||||||
|
var tracker = initTable[int, seq[string]]()
|
||||||
|
for i in 0..<400:
|
||||||
|
let k = rng.rand(0..80)
|
||||||
|
if rng.rand(0..2) < 2:
|
||||||
|
let v = "v" & $i
|
||||||
|
btree.insert(k, v)
|
||||||
|
if k notin tracker: tracker[k] = @[]
|
||||||
|
tracker[k].add(v)
|
||||||
|
else:
|
||||||
|
if k in tracker and tracker[k].len > 0:
|
||||||
|
let v = tracker[k][0]
|
||||||
|
btree.remove(k, v)
|
||||||
|
tracker[k].del(0)
|
||||||
|
if tracker[k].len == 0: tracker.del(k)
|
||||||
|
if i mod 50 == 49:
|
||||||
|
check btree.checkInvariants().len == 0
|
||||||
|
for k, vals in tracker:
|
||||||
|
check btree.get(k).toHashSet == vals.toHashSet
|
||||||
|
|
||||||
|
test "B-tree sequential fill + prefix delete keeps separators":
|
||||||
|
var btree = newBTreeIndex[int, string](order = 4)
|
||||||
|
for i in 0..<60:
|
||||||
|
btree.insert(i, "v" & $i)
|
||||||
|
for i in 0..<30:
|
||||||
|
btree.remove(i, "v" & $i)
|
||||||
|
check btree.checkInvariants().len == 0
|
||||||
|
for i in 30..<60:
|
||||||
|
check btree.get(i) == @["v" & $i]
|
||||||
|
for i in 0..<30:
|
||||||
|
check btree.get(i).len == 0
|
||||||
|
|
||||||
|
test "NULL = NULL is unknown, not true":
|
||||||
|
var ctx = setupCtx()
|
||||||
|
defer: teardown(ctx)
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (1, NULL)"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (2, 'bob')"))
|
||||||
|
let eqNull = executeQuery(ctx, parse("SELECT id FROM users WHERE name = NULL"))
|
||||||
|
check eqNull.success
|
||||||
|
check eqNull.rows.len == 0
|
||||||
|
let isNull = executeQuery(ctx, parse("SELECT id FROM users WHERE name IS NULL"))
|
||||||
|
check isNull.success
|
||||||
|
check isNull.rows.len == 1
|
||||||
|
check valueToString(isNull.rows[0]["id"]) == "1"
|
||||||
|
|
||||||
|
test "NULL != value is unknown so WHERE excludes the row":
|
||||||
|
var ctx = setupCtx()
|
||||||
|
defer: teardown(ctx)
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (1, NULL)"))
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (2, 'bob')"))
|
||||||
|
let r = executeQuery(ctx, parse("SELECT id FROM users WHERE name != 'bob'"))
|
||||||
|
check r.success
|
||||||
|
check r.rows.len == 0
|
||||||
|
|
||||||
|
test "SELECT NULL = NULL yields NULL, IS NULL is true":
|
||||||
|
var ctx = setupCtx()
|
||||||
|
defer: teardown(ctx)
|
||||||
|
discard executeQuery(ctx, parse("INSERT INTO users (id, name) VALUES (1, 'a')"))
|
||||||
|
let eq = executeQuery(ctx, parse("SELECT (NULL = NULL) AS x FROM users"))
|
||||||
|
check eq.success
|
||||||
|
check eq.rows.len == 1
|
||||||
|
check valueToString(eq.rows[0]["x"]) == "\\N"
|
||||||
|
let isn = executeQuery(ctx, parse("SELECT (NULL IS NULL) AS x FROM users"))
|
||||||
|
check isn.success
|
||||||
|
check valueToString(isn.rows[0]["x"]) == "true"
|
||||||
|
|||||||
+105
-7
@@ -1659,14 +1659,28 @@ suite "Replication":
|
|||||||
rm.connectReplica("r2")
|
rm.connectReplica("r2")
|
||||||
rm.connectReplica("r3")
|
rm.connectReplica("r3")
|
||||||
|
|
||||||
|
# Unreachable replicas cannot ack — semi-sync must fail closed (return 0)
|
||||||
let lsn = rm.writeLsn(@[1'u8])
|
let lsn = rm.writeLsn(@[1'u8])
|
||||||
check not rm.isFullyAcked(lsn) # needs 2 acks
|
check lsn == 0
|
||||||
|
|
||||||
rm.ackLsn("r1", lsn)
|
# No connected replicas → nothing to wait for; write succeeds
|
||||||
check not rm.isFullyAcked(lsn) # still needs 1 more
|
var rm2 = newReplicationManager(rmSemiSync, syncCount = 2)
|
||||||
|
rm2.addReplica(newReplica("r1", "10.0.0.1", 9472))
|
||||||
|
# not connected
|
||||||
|
let lsn2 = rm2.writeLsn(@[1'u8])
|
||||||
|
check lsn2 > 0
|
||||||
|
check rm2.isFullyAcked(lsn2)
|
||||||
|
|
||||||
rm.ackLsn("r2", lsn)
|
# ackLsn bookkeeping still clears pendingAcks at the required quorum
|
||||||
check rm.isFullyAcked(lsn) # 2 acks received
|
var rm3 = newReplicationManager(rmSemiSync, syncCount = 2)
|
||||||
|
rm3.pendingAcks[1'u64] = initHashSet[string]()
|
||||||
|
rm3.pendingAcks[1'u64].incl("r1")
|
||||||
|
rm3.pendingAcks[1'u64].incl("r2")
|
||||||
|
check not rm3.isFullyAcked(1)
|
||||||
|
rm3.ackLsn("r1", 1)
|
||||||
|
check not rm3.isFullyAcked(1)
|
||||||
|
rm3.ackLsn("r2", 1)
|
||||||
|
check rm3.isFullyAcked(1)
|
||||||
|
|
||||||
test "Replica status":
|
test "Replica status":
|
||||||
var rm = newReplicationManager(rmAsync)
|
var rm = newReplicationManager(rmAsync)
|
||||||
@@ -2895,6 +2909,29 @@ suite "Raft InstallSnapshot Send":
|
|||||||
check node.matchIndex["peer-1"] == 101
|
check node.matchIndex["peer-1"] == 101
|
||||||
check node.nextIndex["peer-1"] == 102
|
check node.nextIndex["peer-1"] == 102
|
||||||
|
|
||||||
|
test "commit requires strict majority for even-sized clusters":
|
||||||
|
## Regression: the commit quorum used (N+1) div 2, which for a 4-node
|
||||||
|
## cluster commits at 2/4 (a minority). Strict majority is N div 2 + 1.
|
||||||
|
var node = newRaftNode("leader", @["p1", "p2", "p3"])
|
||||||
|
node.currentTerm = 5
|
||||||
|
node.state = rsLeader
|
||||||
|
let e = node.appendLog("put", cast[seq[byte]]("k\x00v"))
|
||||||
|
check e.index == 1
|
||||||
|
check e.term == 5
|
||||||
|
node.nextIndex["p1"] = 2
|
||||||
|
node.nextIndex["p2"] = 2
|
||||||
|
node.nextIndex["p3"] = 2
|
||||||
|
# Leader + 1 peer (count=2) is NOT a majority of 4.
|
||||||
|
node.handleAppendReply("p1", RaftMessage(
|
||||||
|
kind: rmkAppendEntriesReply, term: 5, senderId: "p1",
|
||||||
|
success: true, matchIdx: 1))
|
||||||
|
check node.commitIndex == 0
|
||||||
|
# Leader + 2 peers (count=3) IS a strict majority of 4 -> commits.
|
||||||
|
node.handleAppendReply("p2", RaftMessage(
|
||||||
|
kind: rmkAppendEntriesReply, term: 5, senderId: "p2",
|
||||||
|
success: true, matchIdx: 1))
|
||||||
|
check node.commitIndex == 1
|
||||||
|
|
||||||
test "InstallSnapshotReply success advances match/next index and clears streak":
|
test "InstallSnapshotReply success advances match/next index and clears streak":
|
||||||
var node = newRaftNode("leader", @["peer-1"])
|
var node = newRaftNode("leader", @["peer-1"])
|
||||||
node.currentTerm = 5
|
node.currentTerm = 5
|
||||||
@@ -3021,7 +3058,14 @@ suite "Raft InstallSnapshot Send":
|
|||||||
bt: uint64): bool {.gcsafe.} =
|
bt: uint64): bool {.gcsafe.} =
|
||||||
gotBaseIndex = bi
|
gotBaseIndex = bi
|
||||||
gotBaseTerm = bt
|
gotBaseTerm = bt
|
||||||
result = readFile(p) == payload
|
# sendSnapshot now gzips the tar off the event loop, so the assembled
|
||||||
|
# archive is gzip-compressed; decompress before comparing the bytes.
|
||||||
|
let raw = p & ".raw"
|
||||||
|
defer:
|
||||||
|
if fileExists(raw): removeFile(raw)
|
||||||
|
if not gunzipFile(p, raw):
|
||||||
|
return false
|
||||||
|
result = readFile(raw) == payload
|
||||||
|
|
||||||
let netL = newRaftNetwork(leader)
|
let netL = newRaftNetwork(leader)
|
||||||
let netF = newRaftNetwork(follower)
|
let netF = newRaftNetwork(follower)
|
||||||
@@ -3049,8 +3093,10 @@ suite "Raft InstallSnapshot Send":
|
|||||||
check follower.lastSnapshotTerm == 4
|
check follower.lastSnapshotTerm == 4
|
||||||
check gotBaseIndex == 100
|
check gotBaseIndex == 100
|
||||||
check gotBaseTerm == 4
|
check gotBaseTerm == 4
|
||||||
# Temp archive cleaned up after the transfer
|
# Temp archives (uncompressed tar + compressed .tar.gz) cleaned up after
|
||||||
|
# the transfer
|
||||||
check not fileExists(tmp / "raft-l" / "snap_out_100.tar.gz")
|
check not fileExists(tmp / "raft-l" / "snap_out_100.tar.gz")
|
||||||
|
check not fileExists(tmp / "raft-l" / "snap_out_100.tar")
|
||||||
scenario()
|
scenario()
|
||||||
|
|
||||||
test "sendSnapshot single-flight guard skips a concurrent send":
|
test "sendSnapshot single-flight guard skips a concurrent send":
|
||||||
@@ -3441,6 +3487,43 @@ suite "Raft SQL Write Path":
|
|||||||
let (found, _) = db.get("t.id=1")
|
let (found, _) = db.get("t.id=1")
|
||||||
check not found
|
check not found
|
||||||
|
|
||||||
|
test "REP receiver chain (encode -> decode -> apply) maintains indexes":
|
||||||
|
## Mirrors server.nim's legacy REP handler: decodeRepPayload decides the op,
|
||||||
|
## then applyReplicatedPut/Delete keep secondary indexes consistent. Guards
|
||||||
|
## the wiring the receiver relies on — an indexed put must populate the
|
||||||
|
## B-tree and a PK-only put (empty value) must apply as a put, not vanish.
|
||||||
|
var testDir = getTempDir() / "baradb_rep_recv_idx_" & $getCurrentProcessId() & "_" & $getMonoTime().ticks
|
||||||
|
createDir(testDir)
|
||||||
|
defer: removeDir(testDir)
|
||||||
|
var db = newLSMTree(testDir)
|
||||||
|
var ctx = qexec.newExecutionContext(db)
|
||||||
|
discard qexec.executeQuery(ctx, parse(
|
||||||
|
"CREATE TABLE t (id INT PRIMARY KEY, name STRING)"))
|
||||||
|
discard qexec.executeQuery(ctx, parse(
|
||||||
|
"CREATE INDEX idx_name ON t (name)"))
|
||||||
|
# Leader ships an indexed put; the receiver decodes and applies it.
|
||||||
|
let put = decodeRepPayload(
|
||||||
|
encodeRepPayload(false, "t.id=1", cast[seq[byte]]("name=alice")))
|
||||||
|
check put.op == ropPut
|
||||||
|
if put.op == ropPut:
|
||||||
|
applyReplicatedPut(ctx, put.key, put.value)
|
||||||
|
check ctx.btrees["t.name"].get("alice").len >= 1
|
||||||
|
# A PK-only put (empty value) must apply as a put, not a delete.
|
||||||
|
let pk = decodeRepPayload(encodeRepPayload(false, "t.id=2", @[]))
|
||||||
|
check pk.op == ropPut
|
||||||
|
if pk.op == ropPut:
|
||||||
|
applyReplicatedPut(ctx, pk.key, pk.value)
|
||||||
|
let (foundPk, _) = db.get("t.id=2")
|
||||||
|
check foundPk
|
||||||
|
# Leader ships a delete; the receiver drops the row and the index entry.
|
||||||
|
let del = decodeRepPayload(encodeRepPayload(true, "t.id=1", @[]))
|
||||||
|
check del.op == ropDelete
|
||||||
|
if del.op == ropDelete:
|
||||||
|
applyReplicatedDelete(ctx, del.key)
|
||||||
|
check ctx.btrees["t.name"].get("alice").len == 0
|
||||||
|
let (foundDel, _) = db.get("t.id=1")
|
||||||
|
check not foundDel
|
||||||
|
|
||||||
test "applyReplicatedPut updates in-memory graphs":
|
test "applyReplicatedPut updates in-memory graphs":
|
||||||
var testDir = getTempDir() / "baradb_raft_apply_g_" & $getCurrentProcessId() & "_" & $getMonoTime().ticks
|
var testDir = getTempDir() / "baradb_raft_apply_g_" & $getCurrentProcessId() & "_" & $getMonoTime().ticks
|
||||||
createDir(testDir)
|
createDir(testDir)
|
||||||
@@ -4215,6 +4298,21 @@ suite "Window Functions":
|
|||||||
if row["name"] == "Bob":
|
if row["name"] == "Bob":
|
||||||
check row["last_sal"] == "90000"
|
check row["last_sal"] == "90000"
|
||||||
|
|
||||||
|
test "SUM/AVG/COUNT window aggregates over a running frame":
|
||||||
|
## Default frame is UNBOUNDED PRECEDING .. CURRENT ROW (lower.nim).
|
||||||
|
let r = qexec.executeQuery(ctx, parse(
|
||||||
|
"SELECT name, salary, SUM(salary) OVER (ORDER BY salary) AS running, COUNT(*) OVER (ORDER BY salary) AS cnt FROM employees"))
|
||||||
|
check r.success
|
||||||
|
check r.rows.len == 5
|
||||||
|
for row in r.rows:
|
||||||
|
if row["name"] == "Charlie":
|
||||||
|
check parseFloat($row["running"]) == 70000.0
|
||||||
|
check $row["cnt"] == "1"
|
||||||
|
if row["name"] == "Eve":
|
||||||
|
# 70000+75000+80000+90000+95000
|
||||||
|
check parseFloat($row["running"]) == 410000.0
|
||||||
|
check $row["cnt"] == "5"
|
||||||
|
|
||||||
suite "GROUP BY Aggregates":
|
suite "GROUP BY Aggregates":
|
||||||
var db: LSMTree
|
var db: LSMTree
|
||||||
var ctx: qexec.ExecutionContext
|
var ctx: qexec.ExecutionContext
|
||||||
|
|||||||
Reference in New Issue
Block a user