Compare commits
2
Commits
v2.3.0
..
167671fd79
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
167671fd79 | ||
|
|
339a5bd06a |
@@ -22,19 +22,5 @@ jobs:
|
|||||||
run: rustup component add rustfmt clippy
|
run: rustup component add rustfmt clippy
|
||||||
- name: Install thumbv7em-none-eabihf target
|
- name: Install thumbv7em-none-eabihf target
|
||||||
run: rustup target add thumbv7em-none-eabihf
|
run: rustup target add thumbv7em-none-eabihf
|
||||||
- name: Install Python interop dependencies
|
|
||||||
# The interop suites used to skip silently when python3/h5py were
|
|
||||||
# missing, so they never ran in CI. Install them and make a missing
|
|
||||||
# dependency a failure (CLAWHDF5_REQUIRE_INTEROP below).
|
|
||||||
run: |
|
|
||||||
apt-get update
|
|
||||||
apt-get install -y --no-install-recommends python3 python3-venv
|
|
||||||
python3 -m venv /opt/interop
|
|
||||||
/opt/interop/bin/pip install --no-cache-dir h5py numpy netCDF4 xarray
|
|
||||||
echo "/opt/interop/bin" >> "$GITHUB_PATH"
|
|
||||||
- name: Show interop library versions
|
|
||||||
run: python3 -c "import h5py, netCDF4; print('h5py', h5py.__version__, 'HDF5', h5py.version.hdf5_version, 'netCDF4', netCDF4.__version__)"
|
|
||||||
- name: Run CI script
|
- name: Run CI script
|
||||||
env:
|
|
||||||
CLAWHDF5_REQUIRE_INTEROP: "1"
|
|
||||||
run: bash scripts/ci-test.sh
|
run: bash scripts/ci-test.sh
|
||||||
|
|||||||
+1
-161
@@ -1,156 +1,6 @@
|
|||||||
# Changelog
|
# Changelog
|
||||||
|
|
||||||
## v2.3.0 (2026-09-19)
|
## Unreleased
|
||||||
|
|
||||||
### Upgrade Notes
|
|
||||||
- **A memory store now has a single writer.** `HDF5Memory::create`/`open` take
|
|
||||||
an exclusive lock (`<store>.h5.lock`); a second open of the same store — in
|
|
||||||
the same or another process — returns `MemoryError::Locked`. Code that opened
|
|
||||||
a second handle just to read should use `HDF5Memory::open_read_only`.
|
|
||||||
- **Unsigned array attributes arrive as `AttrValue::U64Array`**, not
|
|
||||||
`I64Array`, and `attrs()` may now return `AttrValue::Raw`. Exhaustive matches
|
|
||||||
on `AttrValue` need the two new arms.
|
|
||||||
- **WAL header version 3 → 4.** v3 files are read and upgraded in place, but a
|
|
||||||
store written by 2.3.0 with a pending WAL cannot be opened by 2.2.0 or
|
|
||||||
earlier (it is refused, not corrupted). Checkpoint first
|
|
||||||
(`flush_wal`) if you need to downgrade.
|
|
||||||
- `MemoryConfig::compression` now uses deflate unless the agent's new `zstd`
|
|
||||||
feature is enabled; it previously failed outright in a default build.
|
|
||||||
- `MemoryError` gained `Locked`; `FormatError` gained `UnresolvedSharedMessage`,
|
|
||||||
`ExternalDataFilesUnsupported` and `ExternalLinkUnsupported`; `MessageType`
|
|
||||||
gained `ExternalDataFiles`.
|
|
||||||
|
|
||||||
### Bug Fixes
|
|
||||||
- `clawhdf5-format`: compound datatypes written with **default libver bounds**
|
|
||||||
(datatype message version 1 — what plain `h5py.File(path, 'w')` produces)
|
|
||||||
were mis-parsed. The v1 member layout carries 28 bytes of legacy array
|
|
||||||
fields after the byte offset (the parser skipped 24), and v2 pads member
|
|
||||||
names to 8 bytes and has no array fields at all (the parser did neither), so
|
|
||||||
every member after the first byte offset was read from the wrong position —
|
|
||||||
typically surfacing as `Overflow("compound member ...")` on read. Found by
|
|
||||||
adding a default-libver axis to the h5py interop tests; byte-level regression
|
|
||||||
tests for v1 and v2 added.
|
|
||||||
- `clawhdf5-gpu`: `gpu_tests` could hang forever under the default parallel
|
|
||||||
test runner — every test created its own wgpu instance and device at once.
|
|
||||||
Tests now serialise GPU access, and GPU→CPU readback waits are bounded
|
|
||||||
(30 s) so a wedged driver returns `GpuError::BufferMap` instead of blocking.
|
|
||||||
- `clawhdf5-agent`: `benches/bench.rs` and `benches/memory_bench.rs` no longer
|
|
||||||
compiled against the current `strategy`/`consolidation` APIs.
|
|
||||||
|
|
||||||
### HDF5 Compatibility
|
|
||||||
- `clawhdf5-format`/`clawhdf5`: datasets and attributes that use a **committed
|
|
||||||
(named) datatype** now read correctly. They store a shared-message reference;
|
|
||||||
the facade parsed the reference bytes as the datatype (`Time { size: 0 }`,
|
|
||||||
unreadable data) and silently dropped such attributes. The shared-reference
|
|
||||||
parser itself was wrong for real files: version 2 has no reserved bytes, and
|
|
||||||
the version 3 types were inverted (1 = SOHM heap, 2 = committed).
|
|
||||||
- **Fill values are applied on read.** There was no Fill Value message parser:
|
|
||||||
the holes of a sparse chunked dataset read as zeros even when the fill value
|
|
||||||
was not zero (silently wrong data), and a dataset that was created but never
|
|
||||||
written failed with `NoDataAllocated` where h5py returns a filled array.
|
|
||||||
Messages v1–v3 and the old 0x0004 form are parsed; the fill value is written
|
|
||||||
into exactly the chunk-grid cells missing from the chunk index.
|
|
||||||
- **Soft links are followed** during path resolution, in old- and new-style
|
|
||||||
groups (absolute/relative targets, links to groups, links through links),
|
|
||||||
with a depth limit so a link cycle is an error rather than a hang. A dangling
|
|
||||||
link reports the target it could not find.
|
|
||||||
- Things the reader does not follow are now explicit errors instead of wrong
|
|
||||||
answers: an external link is `ExternalLinkUnsupported { filename,
|
|
||||||
object_path }` (was `PathNotFound`), and a dataset whose raw data lives in
|
|
||||||
external files (message 0x0007, now a known `MessageType`) is
|
|
||||||
`ExternalDataFilesUnsupported` (it would otherwise read as fill values).
|
|
||||||
- **`attrs()` no longer drops attributes.** Any attribute whose datatype had
|
|
||||||
no `AttrValue` variant was omitted with no error — including every Python
|
|
||||||
`bool` (h5py stores `attrs["flag"] = True` as an enum), complex numbers,
|
|
||||||
compound values and object references. Now:
|
|
||||||
- numpy/h5py-style booleans (an enum of exactly `FALSE`=0 / `TRUE`=1) decode
|
|
||||||
as `I64` / `I64Array` of 0/1;
|
|
||||||
- new `AttrValue::U64Array` keeps unsigned arrays unsigned (they were cast to
|
|
||||||
`I64Array`, so values above `i64::MAX` came back negative). **Behaviour
|
|
||||||
change:** code matching `I64Array` for an unsigned attribute must also
|
|
||||||
match `U64Array` (the netCDF-4 CF helpers and Python bindings do);
|
|
||||||
- new `AttrValue::Raw { datatype, shape, data }` carries everything else
|
|
||||||
verbatim, decodable with `clawhdf5_format::data_read` against `datatype`.
|
|
||||||
Both new variants are writable, so an attribute can be copied between files
|
|
||||||
unchanged. Python receives `Raw` as `{"dtype", "shape", "data"}`.
|
|
||||||
- All of the above are covered by h5py interop tests under both default and
|
|
||||||
`libver='latest'` bounds, compared against h5py's own readback.
|
|
||||||
|
|
||||||
### Security
|
|
||||||
- `clawhdf5`: virtual-dataset source file names are untrusted input but were
|
|
||||||
joined straight onto the opened file's directory, so a crafted file could
|
|
||||||
make the reader open any path the process can reach (absolute path, or `..`
|
|
||||||
components). Only plain relative paths inside that directory are accepted.
|
|
||||||
|
|
||||||
### Durability & Integrity
|
|
||||||
- `clawhdf5-agent`: a crash between writing a checkpoint and truncating the WAL
|
|
||||||
no longer **duplicates every pending entry** on the next open. Each
|
|
||||||
checkpoint records a `WalMark` (byte length + chained CRC of the WAL prefix it
|
|
||||||
folded in) in `/meta`; `open()` skips exactly that prefix when it is still
|
|
||||||
present. No WAL format change for this; older files behave as before.
|
|
||||||
- `clawhdf5-agent`: checkpoints and snapshots are durable as a unit — the temp
|
|
||||||
file is synced before the rename and the directory after it. Individual WAL
|
|
||||||
appends remain unsynced by design (documented in `CLAUDE.md`).
|
|
||||||
- `clawhdf5-agent`: `save_or_update` hits are logged as a new `Update` WAL
|
|
||||||
record, so replay updates in place instead of appending a duplicate. WAL
|
|
||||||
header version 3 → 4 (so older builds refuse the file rather than truncating
|
|
||||||
a record they can't parse); v3 files are read and upgraded in place.
|
|
||||||
- `clawhdf5-agent`: loading validates every per-record dataset length (a
|
|
||||||
truncated store is now `MemoryError::Schema`, not a later panic), fixes the
|
|
||||||
`n.len() == n.len()` tautology that trusted a norms dataset of any length,
|
|
||||||
and rejects `embedding_dim == 0` with records present.
|
|
||||||
- `clawhdf5-agent`: eight behavioural `MemoryConfig` fields are now persisted in
|
|
||||||
`/meta`. Previously they reset to defaults on every open — a compressed store
|
|
||||||
was rewritten uncompressed, `wal_enabled = false` flipped back to `true`.
|
|
||||||
- `clawhdf5-agent`: `compression = true` never worked in a default build (it
|
|
||||||
requested Zstd without enabling the feature, so every checkpoint failed with
|
|
||||||
`unsupported filter: 32015`). Default builds now use deflate; Zstd is the new
|
|
||||||
opt-in `zstd` feature.
|
|
||||||
- `clawhdf5-agent`: **single-writer lock** (`<store>.h5.lock`,
|
|
||||||
`MemoryError::Locked`) — two handles on one store used to silently destroy
|
|
||||||
each other's data. New `HDF5Memory::open_read_only` gives a lock-free,
|
|
||||||
never-writing view; the CLI's read-only subcommands use it.
|
|
||||||
- `clawhdf5-agent`: an unreadable WAL (torn header / bad magic) is quarantined
|
|
||||||
(`HDF5Memory::quarantined_wal()`) instead of blocking `open()` of a healthy
|
|
||||||
store. A WAL from an unknown newer version still fails and is left intact.
|
|
||||||
- `clawhdf5-agent`: provenance records are renumbered on compaction (they
|
|
||||||
weren't, so every later `save_or_update` raised a false High integrity
|
|
||||||
alert); pending anomaly alerts and tracked sessions are bounded;
|
|
||||||
`snapshot()` includes entries still in the WAL.
|
|
||||||
- `clawhdf5-agent`: hybrid ranking is deterministic (index tie-breaks instead
|
|
||||||
of `HashMap` order); a set of identical positive scores — including a single
|
|
||||||
candidate — normalises to 1.0 rather than 0.0; the Hebbian boost no longer
|
|
||||||
reinforces zero-score filler results.
|
|
||||||
- `clawhdf5-format`: chunked/VDS/hyperslab reads size their buffers with
|
|
||||||
overflow-checked arithmetic and fallible allocation, so crafted dimensions
|
|
||||||
are `FormatError::Overflow` instead of a wrapped size or a process abort;
|
|
||||||
`parallel_read` bounds checks use `checked_add`.
|
|
||||||
- `clawhdf5`: a malformed filter-pipeline message is an error instead of being
|
|
||||||
treated as "no filters" (which returned compressed bytes as data);
|
|
||||||
`FileBuilder::write` is atomic and synced instead of truncating the
|
|
||||||
destination first.
|
|
||||||
|
|
||||||
### CI / Testing
|
|
||||||
- CI now lints every target (`cargo clippy --all-targets`) plus
|
|
||||||
`clawhdf5-format`'s optional features, compiles all benches, and tests the
|
|
||||||
format feature matrix. Previously test/bench code and feature-gated modules
|
|
||||||
were never linted; the accumulated clippy backlog is fixed.
|
|
||||||
- CI installs python3 + h5py/numpy/netCDF4/xarray and sets
|
|
||||||
`CLAWHDF5_REQUIRE_INTEROP=1`, which turns a missing interop dependency into a
|
|
||||||
test **failure**. Until now every h5py/netCDF4 interop test silently skipped
|
|
||||||
in CI, which is how the HDF5 2.0 compound bug fixed in v2.2.0 reached a user.
|
|
||||||
The `#[ignore]`d `writer_h5py_tests` suite is run explicitly.
|
|
||||||
- h5py-generated-file tests now cover default libver bounds as well as
|
|
||||||
`libver='latest'` (HDF5 2.0 raised the default low bound to 1.8).
|
|
||||||
- `clawhdf5-agent`: WAL property tests (round trip; after any corruption the
|
|
||||||
entries read back are an exact prefix of what was written — 1500 seeded
|
|
||||||
cases), a crash-recovery matrix (an on-disk image after every operation, the
|
|
||||||
checkpoint window, and the WAL torn at every byte length, each reopened and
|
|
||||||
checked against a model), and a WAL fuzz target.
|
|
||||||
- Optional fuzz smoke run (`CLAWHDF5_FUZZ_SECONDS=N scripts/ci-test.sh`); new
|
|
||||||
datatype corpus seeds for v1 compound and native complex messages.
|
|
||||||
|
|
||||||
## v2.2.0 (2026-09-18)
|
|
||||||
|
|
||||||
### Security
|
### Security
|
||||||
- `clawhdf5-format`: bounded decompression output (`MAX_DECOMPRESS_SIZE`) for
|
- `clawhdf5-format`: bounded decompression output (`MAX_DECOMPRESS_SIZE`) for
|
||||||
@@ -395,16 +245,6 @@
|
|||||||
reading compound types and — critically — every chunked/compressed dataset
|
reading compound types and — critically — every chunked/compressed dataset
|
||||||
written by HDF5 2.0. Found by running the h5py interop tests against
|
written by HDF5 2.0. Found by running the h5py interop tests against
|
||||||
h5py 3.16 / HDF5 2.0.
|
h5py 3.16 / HDF5 2.0.
|
||||||
Independently reported (with a patch) against the v2.1.0 tag by
|
|
||||||
M. Scot Breitenfeld (The HDF Group) — v2.1.0 predates this fix.
|
|
||||||
- `clawhdf5-format`: parse HDF5 2.0 native complex datatypes (class 11,
|
|
||||||
datatype version 5, e.g. `H5T_COMPLEX_IEEE_F64LE`). The properties are a
|
|
||||||
single base floating-point datatype, not a compound-style member list; the
|
|
||||||
old parser read the base type's bytes as member names, producing a garbage
|
|
||||||
datatype, and failed with `UnexpectedEof` when a complex type was nested in
|
|
||||||
a compound. It is now surfaced as the equivalent `{r, i}` compound (the
|
|
||||||
shape h5py writes for numpy complex dtypes), with a size check against the
|
|
||||||
base type. Validated end-to-end against an HDF5 2.0-written file.
|
|
||||||
|
|
||||||
### Performance
|
### Performance
|
||||||
- `clawhdf5-format`: chunked writes now compress all chunks up front via
|
- `clawhdf5-format`: chunked writes now compress all chunks up front via
|
||||||
|
|||||||
@@ -33,48 +33,7 @@ Cargo workspace with 16 crates under `crates/` (plus `libaec-sys`, an internal F
|
|||||||
the approximate `clawhdf5-ann` index for the vector stage (the index mirrors
|
the approximate `clawhdf5-ann` index for the vector stage (the index mirrors
|
||||||
the cache and self-heals on drift). Build the agent with
|
the cache and self-heals on drift). Build the agent with
|
||||||
`--no-default-features --features float16` to force the exact linear cosine scan.
|
`--no-default-features --features float16` to force the exact linear cosine scan.
|
||||||
- WAL (write-ahead log) for crash-safe persistence, with a chained CRC32
|
- WAL (write-ahead log) for crash-safe persistence, with a CRC32 trailer per entry so a corrupted entry stops replay cleanly instead of loading bad data
|
||||||
trailer per entry (each entry's CRC folds in the previous entry's CRC) so a
|
|
||||||
corrupted, reordered, duplicated, or spliced entry stops replay cleanly
|
|
||||||
instead of loading bad or tampered data. The pre-chaining per-entry-CRC
|
|
||||||
format (v2) is still fully readable; the oldest no-CRC format (v1) is only
|
|
||||||
reachable through the one-time migration path in `HDF5Memory::open`, not
|
|
||||||
through the public `WalFile::read_entries`.
|
|
||||||
**What the WAL guarantees:** integrity, ordering, and recovery from a
|
|
||||||
*process* crash at any point — including between a checkpoint and the WAL
|
|
||||||
truncate (each checkpoint records a `WalMark` in `/meta`, and `open()` skips
|
|
||||||
the WAL prefix the `.h5` already contains, so entries are never applied
|
|
||||||
twice). Checkpoints and snapshots are made durable as a unit (temp file
|
|
||||||
synced, renamed, directory synced). **What it does not guarantee:**
|
|
||||||
individual WAL appends are *not* fsynced (a deliberate latency trade-off), so
|
|
||||||
saves made since the last checkpoint can be lost on power failure or kernel
|
|
||||||
panic. Current header version is 4 (adds the `Update` record used by
|
|
||||||
`save_or_update`); v3 files are read and upgraded in place.
|
|
||||||
- A store has a **single writer**: `HDF5Memory::create`/`open` hold an exclusive
|
|
||||||
advisory lock on `<store>.h5.lock` and a second opener gets
|
|
||||||
`MemoryError::Locked`. Use `HDF5Memory::open_read_only` for a lock-free,
|
|
||||||
never-writing point-in-time view (the CLI's `recall`/`stats`/`agents-md`/
|
|
||||||
`export` do). An unreadable WAL (torn header, bad magic) is quarantined to
|
|
||||||
`<store>.h5.wal.corrupt-<ts>` rather than blocking `open()`; a WAL with an
|
|
||||||
unknown *newer* version still fails and is left untouched.
|
|
||||||
- `MemoryConfig::compression` uses deflate by default; enable the agent's
|
|
||||||
`zstd` feature to compress embeddings with Zstd instead (links libzstd).
|
|
||||||
- `Dataset::verify_provenance()` (clawhdf5 facade, `provenance` feature, on by
|
|
||||||
default) recomputes a dataset's SHA-256 and compares it against the
|
|
||||||
`_provenance_sha256` attribute written automatically on save when
|
|
||||||
`DatasetBuilder::with_provenance` is used. It's opt-in per call, not run
|
|
||||||
automatically on open — it decodes and hashes the whole dataset. The hash
|
|
||||||
is unkeyed (tamper-*evident*, not tamper-*proof*): it detects accidental
|
|
||||||
corruption, not a deliberate actor able to modify both the data and the
|
|
||||||
stored hash.
|
|
||||||
- `clawhdf5-agent`'s `HDF5Memory::save`/`save_batch`/`save_or_update` run every
|
|
||||||
write through an in-memory (session-scoped, not persisted to disk)
|
|
||||||
provenance ledger and write-anomaly detector: a content hash per record
|
|
||||||
(`provenance.rs`) for detecting accidental mid-session corruption, plus
|
|
||||||
rate-limit/injection-pattern/source-distribution checks (`anomaly.rs`).
|
|
||||||
Alerts never block a save — drain them with `HDF5Memory::take_anomaly_alerts`.
|
|
||||||
`MemorySource` for this bookkeeping is inferred from the caller-supplied
|
|
||||||
`source_channel` string (a heuristic, not an authenticated trust boundary).
|
|
||||||
- GPU-accelerated batch I/O for large dataset processing
|
- GPU-accelerated batch I/O for large dataset processing
|
||||||
- Python and Node.js bindings for cross-language use
|
- Python and Node.js bindings for cross-language use
|
||||||
- NetCDF-4 compatibility for scientific data interop
|
- NetCDF-4 compatibility for scientific data interop
|
||||||
|
|||||||
+2
-2
@@ -21,10 +21,10 @@ members = [
|
|||||||
resolver = "2"
|
resolver = "2"
|
||||||
|
|
||||||
[workspace.package]
|
[workspace.package]
|
||||||
version = "2.3.0"
|
version = "2.1.0"
|
||||||
edition = "2024"
|
edition = "2024"
|
||||||
license = "MIT"
|
license = "MIT"
|
||||||
repository = "https://git.redclaw.dev/quantumclaw/clawhdf5"
|
repository = "https://github.com/redclawsystems/clawhdf5"
|
||||||
|
|
||||||
[workspace.dependencies]
|
[workspace.dependencies]
|
||||||
tempfile = "3"
|
tempfile = "3"
|
||||||
|
|||||||
@@ -1,10 +1,10 @@
|
|||||||
[package]
|
[package]
|
||||||
name = "clawhdf5-accel"
|
name = "clawhdf5-accel"
|
||||||
version = "2.3.0"
|
version = "2.1.0"
|
||||||
edition = "2024"
|
edition = "2024"
|
||||||
description = "SIMD-accelerated operations for rustyhdf5"
|
description = "SIMD-accelerated operations for rustyhdf5"
|
||||||
license = "MIT"
|
license = "MIT"
|
||||||
repository = "https://git.redclaw.dev/quantumclaw/clawhdf5"
|
repository = "https://github.com/redclawsystems/clawhdf5"
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
keywords = ["hdf5", "simd", "acceleration", "performance"]
|
keywords = ["hdf5", "simd", "acceleration", "performance"]
|
||||||
categories = ["science", "algorithms"]
|
categories = ["science", "algorithms"]
|
||||||
|
|||||||
@@ -111,11 +111,7 @@ pub unsafe fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
|
|||||||
}
|
}
|
||||||
|
|
||||||
let denom = (norm_a * norm_b).sqrt();
|
let denom = (norm_a * norm_b).sqrt();
|
||||||
if denom < f32::EPSILON {
|
if denom == 0.0 { 0.0 } else { dot / denom }
|
||||||
0.0
|
|
||||||
} else {
|
|
||||||
dot / denom
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -89,11 +89,7 @@ pub unsafe fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
|
|||||||
}
|
}
|
||||||
|
|
||||||
let denom = (norm_a * norm_b).sqrt();
|
let denom = (norm_a * norm_b).sqrt();
|
||||||
if denom < f32::EPSILON {
|
if denom == 0.0 { 0.0 } else { dot / denom }
|
||||||
0.0
|
|
||||||
} else {
|
|
||||||
dot / denom
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -361,18 +361,6 @@ mod tests {
|
|||||||
assert!(approx_eq(cosine_similarity(&a, &b), 0.0, EPSILON));
|
assert!(approx_eq(cosine_similarity(&a, &b), 0.0, EPSILON));
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_cosine_near_zero_norm_clamped() {
|
|
||||||
// denom = 1e-4 * 1e-4 = 1e-8, comfortably below f32::EPSILON
|
|
||||||
// (~1.19e-7) but not exactly 0.0 — must still clamp to 0.0 so
|
|
||||||
// callers computing `1.0 - cosine_similarity(...)` treat these
|
|
||||||
// as maximally dissimilar, matching the pre-SIMD scalar guard.
|
|
||||||
let a = [1e-4f32];
|
|
||||||
let b = [1e-4f32];
|
|
||||||
assert_eq!(cosine_similarity(&a, &b), 0.0);
|
|
||||||
assert_eq!(scalar::cosine_similarity(&a, &b), 0.0);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_cosine_scalar_vs_dispatch() {
|
fn test_cosine_scalar_vs_dispatch() {
|
||||||
let a: Vec<f32> = (0..384).map(|i| (i as f32).sin()).collect();
|
let a: Vec<f32> = (0..384).map(|i| (i as f32).sin()).collect();
|
||||||
|
|||||||
@@ -94,11 +94,7 @@ pub unsafe fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
|
|||||||
}
|
}
|
||||||
|
|
||||||
let denom = (norm_a * norm_b).sqrt();
|
let denom = (norm_a * norm_b).sqrt();
|
||||||
if denom < f32::EPSILON {
|
if denom == 0.0 { 0.0 } else { dot / denom }
|
||||||
0.0
|
|
||||||
} else {
|
|
||||||
dot / denom
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// NEON L2 distance.
|
/// NEON L2 distance.
|
||||||
|
|||||||
@@ -21,11 +21,7 @@ pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
|
|||||||
norm_b += y * y;
|
norm_b += y * y;
|
||||||
}
|
}
|
||||||
let denom = (norm_a * norm_b).sqrt();
|
let denom = (norm_a * norm_b).sqrt();
|
||||||
if denom < f32::EPSILON {
|
if denom == 0.0 { 0.0 } else { dot / denom }
|
||||||
0.0
|
|
||||||
} else {
|
|
||||||
dot / denom
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn batch_cosine(query: &[f32], vectors: &[&[f32]], results: &mut [(usize, f32)]) {
|
pub fn batch_cosine(query: &[f32], vectors: &[&[f32]], results: &mut [(usize, f32)]) {
|
||||||
|
|||||||
@@ -1,21 +1,21 @@
|
|||||||
[package]
|
[package]
|
||||||
name = "clawhdf5-agent"
|
name = "clawhdf5-agent"
|
||||||
version = "2.3.0"
|
version = "2.1.0"
|
||||||
edition = "2024"
|
edition = "2024"
|
||||||
description = "HDF5-backed persistent memory store for on-device AI agents"
|
description = "HDF5-backed persistent memory store for on-device AI agents"
|
||||||
license = "MIT"
|
license = "MIT"
|
||||||
repository = "https://git.redclaw.dev/quantumclaw/clawhdf5"
|
repository = "https://github.com/redclawsystems/clawhdf5"
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
keywords = ["agent", "memory", "hdf5", "vector-search", "embedding"]
|
keywords = ["agent", "memory", "hdf5", "vector-search", "embedding"]
|
||||||
categories = ["database", "science", "algorithms"]
|
categories = ["database", "science", "algorithms"]
|
||||||
|
|
||||||
[dependencies]
|
[dependencies]
|
||||||
clawhdf5-format = { path = "../clawhdf5-format", version = "2.3.0", features = ["parallel", "fast-checksum"] }
|
clawhdf5-format = { path = "../clawhdf5-format", version = "2.1.0", features = ["parallel", "fast-checksum"] }
|
||||||
clawhdf5 = { path = "../clawhdf5", version = "2.3.0" }
|
clawhdf5 = { path = "../clawhdf5", version = "2.1.0" }
|
||||||
clawhdf5-io = { path = "../clawhdf5-io", version = "2.3.0", features = ["mmap"] }
|
clawhdf5-io = { path = "../clawhdf5-io", version = "2.1.0", features = ["mmap"] }
|
||||||
clawhdf5-accel = { path = "../clawhdf5-accel", version = "2.3.0" }
|
clawhdf5-accel = { path = "../clawhdf5-accel", version = "2.1.0" }
|
||||||
clawhdf5-ann = { path = "../clawhdf5-ann", version = "2.3.0", optional = true }
|
clawhdf5-ann = { path = "../clawhdf5-ann", version = "2.1.0", optional = true }
|
||||||
clawhdf5-gpu = { path = "../clawhdf5-gpu", version = "2.3.0", optional = true, default-features = false }
|
clawhdf5-gpu = { path = "../clawhdf5-gpu", version = "2.1.0", optional = true, default-features = false }
|
||||||
serde = { workspace = true }
|
serde = { workspace = true }
|
||||||
byteorder = "1"
|
byteorder = "1"
|
||||||
half = { workspace = true, optional = true }
|
half = { workspace = true, optional = true }
|
||||||
@@ -48,9 +48,6 @@ harness = false
|
|||||||
default = ["float16", "hnsw"]
|
default = ["float16", "hnsw"]
|
||||||
float16 = ["half"]
|
float16 = ["half"]
|
||||||
parallel = ["rayon"]
|
parallel = ["rayon"]
|
||||||
# Compress embeddings with Zstd instead of deflate when
|
|
||||||
# `MemoryConfig::compression` is on. Off by default: it links libzstd (C).
|
|
||||||
zstd = ["clawhdf5/zstd"]
|
|
||||||
# HNSW approximate-nearest-neighbour acceleration for the vector stage of
|
# HNSW approximate-nearest-neighbour acceleration for the vector stage of
|
||||||
# hybrid_search. On by default; the index is rebuilt from the cache on demand
|
# hybrid_search. On by default; the index is rebuilt from the cache on demand
|
||||||
# and stays self-consistent with the persisted memory store. Disable with
|
# and stays self-consistent with the persisted memory store. Disable with
|
||||||
|
|||||||
@@ -483,7 +483,7 @@ fn rayon_benches(c: &mut Criterion) {
|
|||||||
use rayon::prelude::*;
|
use rayon::prelude::*;
|
||||||
let query_norm = vector_search::compute_norm(&query);
|
let query_norm = vector_search::compute_norm(&query);
|
||||||
let num_cores = rayon::current_num_threads().max(1);
|
let num_cores = rayon::current_num_threads().max(1);
|
||||||
let chunk_size = n.div_ceil(num_cores);
|
let chunk_size = (n + num_cores - 1) / num_cores;
|
||||||
let mut results: Vec<(usize, f32)> = vectors
|
let mut results: Vec<(usize, f32)> = vectors
|
||||||
.par_chunks(chunk_size)
|
.par_chunks(chunk_size)
|
||||||
.enumerate()
|
.enumerate()
|
||||||
@@ -537,7 +537,7 @@ fn rayon_benches(c: &mut Criterion) {
|
|||||||
use rayon::prelude::*;
|
use rayon::prelude::*;
|
||||||
let query_norm = vector_search::compute_norm(&query);
|
let query_norm = vector_search::compute_norm(&query);
|
||||||
let num_cores = rayon::current_num_threads().max(1);
|
let num_cores = rayon::current_num_threads().max(1);
|
||||||
let chunk_size = n.div_ceil(num_cores);
|
let chunk_size = (n + num_cores - 1) / num_cores;
|
||||||
let mut results: Vec<(usize, f32)> = vectors
|
let mut results: Vec<(usize, f32)> = vectors
|
||||||
.par_chunks(chunk_size)
|
.par_chunks(chunk_size)
|
||||||
.enumerate()
|
.enumerate()
|
||||||
@@ -766,22 +766,12 @@ fn adaptive_benches(c: &mut Criterion) {
|
|||||||
.map(|v| vector_search::compute_norm(v))
|
.map(|v| vector_search::compute_norm(v))
|
||||||
.collect();
|
.collect();
|
||||||
let tombstones = vec![0u8; n];
|
let tombstones = vec![0u8; n];
|
||||||
let flat: Vec<f32> = vectors.iter().flatten().copied().collect();
|
|
||||||
|
|
||||||
c.bench_function("adaptive_search_10k", |b| {
|
c.bench_function("adaptive_search_10k", |b| {
|
||||||
let hw = HardwareCapabilities::detect();
|
let hw = HardwareCapabilities::detect();
|
||||||
let strat = strategy::auto_select_strategy(n, &hw);
|
let strat = strategy::auto_select_strategy(n, &hw);
|
||||||
b.iter(|| {
|
b.iter(|| {
|
||||||
strategy::search_with_metrics(
|
strategy::search_with_metrics(&query, &vectors, &norms, &tombstones, 10, strat, None)
|
||||||
&query,
|
|
||||||
&vectors,
|
|
||||||
&flat,
|
|
||||||
&norms,
|
|
||||||
&tombstones,
|
|
||||||
10,
|
|
||||||
strat,
|
|
||||||
None,
|
|
||||||
)
|
|
||||||
});
|
});
|
||||||
});
|
});
|
||||||
|
|
||||||
@@ -791,7 +781,6 @@ fn adaptive_benches(c: &mut Criterion) {
|
|||||||
strategy::search_with_metrics(
|
strategy::search_with_metrics(
|
||||||
&query,
|
&query,
|
||||||
&vectors,
|
&vectors,
|
||||||
&flat,
|
|
||||||
&norms,
|
&norms,
|
||||||
&tombstones,
|
&tombstones,
|
||||||
10,
|
10,
|
||||||
@@ -806,7 +795,6 @@ fn adaptive_benches(c: &mut Criterion) {
|
|||||||
strategy::search_with_metrics(
|
strategy::search_with_metrics(
|
||||||
&query,
|
&query,
|
||||||
&vectors,
|
&vectors,
|
||||||
&flat,
|
|
||||||
&norms,
|
&norms,
|
||||||
&tombstones,
|
&tombstones,
|
||||||
10,
|
10,
|
||||||
@@ -821,7 +809,6 @@ fn adaptive_benches(c: &mut Criterion) {
|
|||||||
strategy::search_with_metrics(
|
strategy::search_with_metrics(
|
||||||
&query,
|
&query,
|
||||||
&vectors,
|
&vectors,
|
||||||
&flat,
|
|
||||||
&norms,
|
&norms,
|
||||||
&tombstones,
|
&tombstones,
|
||||||
10,
|
10,
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
use clawhdf5_agent::bm25::BM25Index;
|
use clawhdf5_agent::bm25::BM25Index;
|
||||||
use clawhdf5_agent::consolidation::{
|
use clawhdf5_agent::consolidation::{
|
||||||
ConsolidationConfig, ConsolidationEngine, ImportanceScorer, ImportanceWeights, MemorySource,
|
ConsolidationConfig, ConsolidationEngine, ImportanceScorer, ImportanceWeights, MemorySource,
|
||||||
UntrustedSource,
|
|
||||||
};
|
};
|
||||||
use clawhdf5_agent::hybrid::{hybrid_search, rrf_hybrid_search};
|
use clawhdf5_agent::hybrid::{hybrid_search, rrf_hybrid_search};
|
||||||
use clawhdf5_agent::knowledge::KnowledgeCache;
|
use clawhdf5_agent::knowledge::KnowledgeCache;
|
||||||
@@ -286,12 +285,7 @@ fn consolidation_benches(c: &mut Criterion) {
|
|||||||
for i in 0..n {
|
for i in 0..n {
|
||||||
let embedding = make_vec(&mut rng, DIM);
|
let embedding = make_vec(&mut rng, DIM);
|
||||||
let chunk = format!("memory record {i} with some content");
|
let chunk = format!("memory record {i} with some content");
|
||||||
engine.add_memory(
|
engine.add_memory(chunk, embedding, MemorySource::User, now + i as f64);
|
||||||
chunk,
|
|
||||||
embedding,
|
|
||||||
UntrustedSource::User,
|
|
||||||
now + i as f64,
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
engine
|
engine
|
||||||
},
|
},
|
||||||
@@ -313,10 +307,9 @@ fn consolidation_benches(c: &mut Criterion) {
|
|||||||
for i in 0..50usize {
|
for i in 0..50usize {
|
||||||
let embedding = make_vec(&mut rng, DIM);
|
let embedding = make_vec(&mut rng, DIM);
|
||||||
let chunk = format!("existing record {i}");
|
let chunk = format!("existing record {i}");
|
||||||
engine.add_memory(chunk, embedding, UntrustedSource::User, now + i as f64);
|
engine.add_memory(chunk, embedding, MemorySource::User, now + i as f64);
|
||||||
}
|
}
|
||||||
let records = engine.records().to_vec();
|
let records = engine.records().to_vec();
|
||||||
let record_refs: Vec<&_> = records.iter().collect();
|
|
||||||
let weights = ImportanceWeights::default();
|
let weights = ImportanceWeights::default();
|
||||||
let query_embedding = make_vec(&mut rng, DIM);
|
let query_embedding = make_vec(&mut rng, DIM);
|
||||||
let sample_text =
|
let sample_text =
|
||||||
@@ -324,7 +317,7 @@ fn consolidation_benches(c: &mut Criterion) {
|
|||||||
|
|
||||||
group.bench_function("bench_importance_scoring", |b| {
|
group.bench_function("bench_importance_scoring", |b| {
|
||||||
b.iter(|| {
|
b.iter(|| {
|
||||||
let surprise = ImportanceScorer::score_surprise(&query_embedding, &record_refs);
|
let surprise = ImportanceScorer::score_surprise(&query_embedding, &records);
|
||||||
let correction = ImportanceScorer::score_correction(&MemorySource::Correction);
|
let correction = ImportanceScorer::score_correction(&MemorySource::Correction);
|
||||||
let length = ImportanceScorer::score_length(sample_text);
|
let length = ImportanceScorer::score_length(sample_text);
|
||||||
ImportanceScorer::score_combined(surprise, correction, length, &weights)
|
ImportanceScorer::score_combined(surprise, correction, length, &weights)
|
||||||
@@ -361,7 +354,7 @@ fn temporal_benches(c: &mut Criterion) {
|
|||||||
// Insert benchmark: measure time to insert 10k timestamps one by one
|
// Insert benchmark: measure time to insert 10k timestamps one by one
|
||||||
group.bench_function("bench_temporal_insert_10k", |b| {
|
group.bench_function("bench_temporal_insert_10k", |b| {
|
||||||
b.iter_batched(
|
b.iter_batched(
|
||||||
TemporalIndex::new,
|
|| TemporalIndex::new(),
|
||||||
|mut idx| {
|
|mut idx| {
|
||||||
for i in 0..N {
|
for i in 0..N {
|
||||||
// Shuffle insertion order slightly using a simple offset pattern
|
// Shuffle insertion order slightly using a simple offset pattern
|
||||||
@@ -449,8 +442,7 @@ fn large_consolidation_benches(c: &mut Criterion) {
|
|||||||
let mut group = c.benchmark_group("consolidation_large");
|
let mut group = c.benchmark_group("consolidation_large");
|
||||||
group.sample_size(10);
|
group.sample_size(10);
|
||||||
|
|
||||||
{
|
for (label, n) in [("10k", 10_000usize)] {
|
||||||
let (label, n) = ("10k", 10_000usize);
|
|
||||||
group.bench_with_input(
|
group.bench_with_input(
|
||||||
BenchmarkId::new("bench_consolidation_cycle", label),
|
BenchmarkId::new("bench_consolidation_cycle", label),
|
||||||
&n,
|
&n,
|
||||||
@@ -467,12 +459,7 @@ fn large_consolidation_benches(c: &mut Criterion) {
|
|||||||
for i in 0..n {
|
for i in 0..n {
|
||||||
let embedding = make_vec(&mut rng, DIM);
|
let embedding = make_vec(&mut rng, DIM);
|
||||||
let chunk = format!("memory record {i} with content");
|
let chunk = format!("memory record {i} with content");
|
||||||
engine.add_memory(
|
engine.add_memory(chunk, embedding, MemorySource::User, now + i as f64);
|
||||||
chunk,
|
|
||||||
embedding,
|
|
||||||
UntrustedSource::User,
|
|
||||||
now + i as f64,
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
engine
|
engine
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -1,3 +0,0 @@
|
|||||||
target/
|
|
||||||
artifacts/
|
|
||||||
coverage/
|
|
||||||
@@ -1,23 +0,0 @@
|
|||||||
[package]
|
|
||||||
name = "clawhdf5-agent-fuzz"
|
|
||||||
version = "0.0.0"
|
|
||||||
publish = false
|
|
||||||
edition = "2024"
|
|
||||||
|
|
||||||
[package.metadata]
|
|
||||||
cargo-fuzz = true
|
|
||||||
|
|
||||||
[dependencies]
|
|
||||||
libfuzzer-sys = "0.4"
|
|
||||||
tempfile = "3"
|
|
||||||
|
|
||||||
[dependencies.clawhdf5-agent]
|
|
||||||
path = ".."
|
|
||||||
|
|
||||||
[workspace]
|
|
||||||
members = ["."]
|
|
||||||
|
|
||||||
[[bin]]
|
|
||||||
name = "fuzz_wal_replay"
|
|
||||||
path = "fuzz_targets/fuzz_wal_replay.rs"
|
|
||||||
doc = false
|
|
||||||
@@ -1,36 +0,0 @@
|
|||||||
#![no_main]
|
|
||||||
//! Arbitrary bytes as a WAL file. Reading, and opening for append (which scans
|
|
||||||
//! the chain and truncates an unverifiable tail), must never panic, hang, or
|
|
||||||
//! allocate without bound — and after `open` repairs the file, everything
|
|
||||||
//! `read_entries` returned before must still be returned.
|
|
||||||
//!
|
|
||||||
//! The deterministic counterpart that runs in ordinary CI is
|
|
||||||
//! `tests/wal_properties.rs`; this target explores inputs it cannot reach.
|
|
||||||
|
|
||||||
use std::io::Write as _;
|
|
||||||
|
|
||||||
use clawhdf5_agent::wal::WalFile;
|
|
||||||
use libfuzzer_sys::fuzz_target;
|
|
||||||
|
|
||||||
fuzz_target!(|data: &[u8]| {
|
|
||||||
let Ok(mut tmp) = tempfile::NamedTempFile::new() else {
|
|
||||||
return;
|
|
||||||
};
|
|
||||||
if tmp.write_all(data).and_then(|()| tmp.flush()).is_err() {
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
let before = WalFile::read_entries(tmp.path()).map(|e| e.len());
|
|
||||||
// Only the chained formats (header versions 3 and 4) are repaired in
|
|
||||||
// place. `open` deliberately recreates a legacy-format file from scratch:
|
|
||||||
// `HDF5Memory::open` has already replayed its entries by then.
|
|
||||||
let chained = matches!(data.get(4), Some(3 | 4));
|
|
||||||
let opened = WalFile::open(tmp.path());
|
|
||||||
if !chained {
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
if let (Ok(before), Ok(wal)) = (before, opened) {
|
|
||||||
drop(wal);
|
|
||||||
let after = WalFile::read_entries(tmp.path()).map(|e| e.len());
|
|
||||||
assert_eq!(after.ok(), Some(before), "open() changed what is replayable");
|
|
||||||
}
|
|
||||||
});
|
|
||||||
@@ -82,68 +82,6 @@ impl Default for AnomalyConfig {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
// Pattern-match normalization
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
/// `true` for characters used to invisibly break up text without being
|
|
||||||
/// rendered (zero-width joiners/spacers, bidi control marks, the BOM/ZWNBSP,
|
|
||||||
/// soft hyphen, and the invisible math operators) — a common trick for
|
|
||||||
/// splitting a flagged word so a literal-substring check misses it while the
|
|
||||||
/// text still displays normally.
|
|
||||||
fn is_invisible_format_char(ch: char) -> bool {
|
|
||||||
matches!(
|
|
||||||
ch,
|
|
||||||
'\u{00AD}' // soft hyphen
|
|
||||||
| '\u{200B}' // zero width space
|
|
||||||
| '\u{200C}' // zero width non-joiner
|
|
||||||
| '\u{200D}' // zero width joiner
|
|
||||||
| '\u{200E}' // left-to-right mark
|
|
||||||
| '\u{200F}' // right-to-left mark
|
|
||||||
| '\u{2060}' // word joiner
|
|
||||||
| '\u{2061}'..='\u{2064}' // invisible times/plus/separator/function application
|
|
||||||
| '\u{202A}'..='\u{202E}' // bidi embedding/override controls
|
|
||||||
| '\u{FEFF}' // BOM / zero width no-break space
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Normalize text before suspicious-pattern matching so the cheapest evasion
|
|
||||||
/// tricks — extra whitespace, zero-width characters, or punctuation spliced
|
|
||||||
/// between letters (e.g. `"s.y.s.t.e.m"`) — don't defeat a literal-substring
|
|
||||||
/// check. Lowercases, drops invisible-format and control characters, drops
|
|
||||||
/// punctuation entirely (not just collapses it, so split words rejoin), and
|
|
||||||
/// collapses whitespace runs to a single space.
|
|
||||||
///
|
|
||||||
/// Does not perform Unicode NFKC normalization or confusable/homoglyph
|
|
||||||
/// folding (see [`WriteAnomalyDetector::check_pattern_anomaly`]).
|
|
||||||
fn normalize_for_pattern_match(text: &str) -> String {
|
|
||||||
let mut out = String::with_capacity(text.len());
|
|
||||||
let mut last_was_space = true; // trims leading whitespace for free
|
|
||||||
for ch in text.chars() {
|
|
||||||
if ch.is_control() || is_invisible_format_char(ch) {
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
if ch.is_whitespace() {
|
|
||||||
if !last_was_space {
|
|
||||||
out.push(' ');
|
|
||||||
last_was_space = true;
|
|
||||||
}
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
if ch.is_ascii_punctuation() {
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
for lower in ch.to_lowercase() {
|
|
||||||
out.push(lower);
|
|
||||||
}
|
|
||||||
last_was_space = false;
|
|
||||||
}
|
|
||||||
while out.ends_with(' ') {
|
|
||||||
out.pop();
|
|
||||||
}
|
|
||||||
out
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
// WriteEvent
|
// WriteEvent
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
@@ -161,9 +99,6 @@ pub struct WriteEvent {
|
|||||||
// WriteAnomalyDetector
|
// WriteAnomalyDetector
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
/// Upper bound on distinct session ids the detector tracks at once.
|
|
||||||
const MAX_TRACKED_SESSIONS: usize = 4096;
|
|
||||||
|
|
||||||
/// Tracks write events and raises alerts for suspicious behaviour.
|
/// Tracks write events and raises alerts for suspicious behaviour.
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
pub struct WriteAnomalyDetector {
|
pub struct WriteAnomalyDetector {
|
||||||
@@ -192,23 +127,6 @@ impl WriteAnomalyDetector {
|
|||||||
if event.timestamp > self.last_timestamp {
|
if event.timestamp > self.last_timestamp {
|
||||||
self.last_timestamp = event.timestamp;
|
self.last_timestamp = event.timestamp;
|
||||||
}
|
}
|
||||||
// Bound the per-session map: a long-lived process sees an unbounded
|
|
||||||
// number of distinct session ids. When it overflows, forget the
|
|
||||||
// sessions with the fewest writes (they are furthest from the limit
|
|
||||||
// this map exists to enforce); the current one is re-added below.
|
|
||||||
if self.session_counts.len() >= MAX_TRACKED_SESSIONS
|
|
||||||
&& !self.session_counts.contains_key(&event.session_id)
|
|
||||||
{
|
|
||||||
let mut counts: Vec<u32> = self.session_counts.values().copied().collect();
|
|
||||||
let keep_from = counts.len() / 2;
|
|
||||||
counts.select_nth_unstable(keep_from);
|
|
||||||
let threshold = counts[keep_from];
|
|
||||||
self.session_counts.retain(|_, c| *c >= threshold);
|
|
||||||
if self.session_counts.len() >= MAX_TRACKED_SESSIONS {
|
|
||||||
// Every session had the same count: drop them all.
|
|
||||||
self.session_counts.clear();
|
|
||||||
}
|
|
||||||
}
|
|
||||||
*self
|
*self
|
||||||
.session_counts
|
.session_counts
|
||||||
.entry(event.session_id.clone())
|
.entry(event.session_id.clone())
|
||||||
@@ -228,13 +146,6 @@ impl WriteAnomalyDetector {
|
|||||||
/// Returns an alert if the number of writes in the last 60 seconds exceeds
|
/// Returns an alert if the number of writes in the last 60 seconds exceeds
|
||||||
/// `config.max_writes_per_minute`, or if any session has exceeded
|
/// `config.max_writes_per_minute`, or if any session has exceeded
|
||||||
/// `config.max_writes_per_session`.
|
/// `config.max_writes_per_session`.
|
||||||
///
|
|
||||||
/// The 60-second window is a single shared window across all
|
|
||||||
/// sessions/sources, so when it trips the alert additionally names the
|
|
||||||
/// top-contributing session and source within that window — a session
|
|
||||||
/// can never account for more of the window than the aggregate count, so
|
|
||||||
/// this attributes the same trip to its actual offender rather than
|
|
||||||
/// reporting only the anonymous aggregate total.
|
|
||||||
pub fn check_rate_anomaly(&self) -> Option<AnomalyAlert> {
|
pub fn check_rate_anomaly(&self) -> Option<AnomalyAlert> {
|
||||||
let recent = self.window.len() as u32;
|
let recent = self.window.len() as u32;
|
||||||
if recent > self.config.max_writes_per_minute {
|
if recent > self.config.max_writes_per_minute {
|
||||||
@@ -245,31 +156,11 @@ impl WriteAnomalyDetector {
|
|||||||
} else {
|
} else {
|
||||||
Severity::Medium
|
Severity::Medium
|
||||||
};
|
};
|
||||||
|
|
||||||
let mut per_session: std::collections::HashMap<&str, u32> =
|
|
||||||
std::collections::HashMap::new();
|
|
||||||
// MemorySource isn't Eq/Hash, so key by its Display string instead.
|
|
||||||
let mut per_source: std::collections::HashMap<String, u32> =
|
|
||||||
std::collections::HashMap::new();
|
|
||||||
for e in &self.window {
|
|
||||||
*per_session.entry(e.session_id.as_str()).or_insert(0) += 1;
|
|
||||||
*per_source.entry(e.source.to_string()).or_insert(0) += 1;
|
|
||||||
}
|
|
||||||
let top_session = per_session.iter().max_by_key(|&(_, &c)| c);
|
|
||||||
let top_source = per_source.iter().max_by_key(|&(_, &c)| c);
|
|
||||||
|
|
||||||
let attribution = match (top_session, top_source) {
|
|
||||||
(Some((session, s_count)), Some((source, r_count))) => format!(
|
|
||||||
"; top contributor: session '{session}' with {s_count} writes, \
|
|
||||||
source {source} with {r_count} writes"
|
|
||||||
),
|
|
||||||
_ => String::new(),
|
|
||||||
};
|
|
||||||
return Some(AnomalyAlert {
|
return Some(AnomalyAlert {
|
||||||
severity,
|
severity,
|
||||||
message: format!(
|
message: format!(
|
||||||
"Rate limit exceeded: {} writes in last 60s (max {}){}",
|
"Rate limit exceeded: {} writes in last 60s (max {})",
|
||||||
recent, self.config.max_writes_per_minute, attribution
|
recent, self.config.max_writes_per_minute
|
||||||
),
|
),
|
||||||
timestamp: self.last_timestamp,
|
timestamp: self.last_timestamp,
|
||||||
});
|
});
|
||||||
@@ -297,24 +188,11 @@ impl WriteAnomalyDetector {
|
|||||||
// -----------------------------------------------------------------------
|
// -----------------------------------------------------------------------
|
||||||
|
|
||||||
/// Returns an alert if `chunk` contains any of the configured suspicious
|
/// Returns an alert if `chunk` contains any of the configured suspicious
|
||||||
/// patterns, after normalizing both sides to defeat the cheapest evasion
|
/// patterns (case-insensitive).
|
||||||
/// tricks (case, extra whitespace, punctuation between letters,
|
|
||||||
/// zero-width/invisible-formatting characters).
|
|
||||||
///
|
|
||||||
/// This does not perform Unicode NFKC normalization or confusable/
|
|
||||||
/// homoglyph folding (e.g. Cyrillic 'а' standing in for Latin 'a') —
|
|
||||||
/// that needs a per-codepoint confusable table (Unicode's
|
|
||||||
/// `confusables.txt`) beyond what's practical to hand-roll correctly,
|
|
||||||
/// and no such crate is a dependency of this crate today. A determined
|
|
||||||
/// attacker using homoglyphs can still evade these patterns.
|
|
||||||
pub fn check_pattern_anomaly(&self, chunk: &str) -> Option<AnomalyAlert> {
|
pub fn check_pattern_anomaly(&self, chunk: &str) -> Option<AnomalyAlert> {
|
||||||
let normalized = normalize_for_pattern_match(chunk);
|
let lower = chunk.to_lowercase();
|
||||||
for pattern in &self.config.suspicious_patterns {
|
for pattern in &self.config.suspicious_patterns {
|
||||||
let normalized_pattern = normalize_for_pattern_match(pattern);
|
if lower.contains(pattern.as_str()) {
|
||||||
if normalized_pattern.is_empty() {
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
if normalized.contains(&normalized_pattern) {
|
|
||||||
let severity = if pattern.contains("ignore") || pattern.contains("override") {
|
let severity = if pattern.contains("ignore") || pattern.contains("override") {
|
||||||
Severity::Critical
|
Severity::Critical
|
||||||
} else if pattern.contains("system") || pattern.contains("jailbreak") {
|
} else if pattern.contains("system") || pattern.contains("jailbreak") {
|
||||||
@@ -449,57 +327,6 @@ mod tests {
|
|||||||
assert!(alert.unwrap().severity >= Severity::Medium);
|
assert!(alert.unwrap().severity >= Severity::Medium);
|
||||||
}
|
}
|
||||||
|
|
||||||
/// A single session dominating the shared 60s window must be named in
|
|
||||||
/// the alert, not just the anonymous aggregate count — this is the case
|
|
||||||
/// the separate cumulative max_writes_per_session check doesn't cover
|
|
||||||
/// (the window can trip before the session's lifetime total does).
|
|
||||||
#[test]
|
|
||||||
fn rate_anomaly_names_offending_session() {
|
|
||||||
let mut det = WriteAnomalyDetector::new(cfg());
|
|
||||||
for i in 0..11 {
|
|
||||||
det.record_write(event(
|
|
||||||
1.0 + i as f64 * 0.1,
|
|
||||||
"flood-session",
|
|
||||||
MemorySource::User,
|
|
||||||
));
|
|
||||||
}
|
|
||||||
let alert = det.check_rate_anomaly().unwrap();
|
|
||||||
assert!(
|
|
||||||
alert.message.contains("flood-session"),
|
|
||||||
"expected the offending session to be named, got: {}",
|
|
||||||
alert.message
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
/// When many distinct sessions jointly trip the shared window, the top
|
|
||||||
/// contributor named must actually be the one with the most writes.
|
|
||||||
#[test]
|
|
||||||
fn rate_anomaly_attributes_top_contributor_among_many_sessions() {
|
|
||||||
let mut det = WriteAnomalyDetector::new(cfg());
|
|
||||||
// 5 sessions with 1 write each (below any per-session limit)...
|
|
||||||
for i in 0..5 {
|
|
||||||
det.record_write(event(
|
|
||||||
1.0 + i as f64 * 0.1,
|
|
||||||
"minor-session",
|
|
||||||
MemorySource::User,
|
|
||||||
));
|
|
||||||
}
|
|
||||||
// ...plus one session responsible for the majority of the flood.
|
|
||||||
for i in 0..8 {
|
|
||||||
det.record_write(event(
|
|
||||||
2.0 + i as f64 * 0.1,
|
|
||||||
"major-session",
|
|
||||||
MemorySource::User,
|
|
||||||
));
|
|
||||||
}
|
|
||||||
let alert = det.check_rate_anomaly().unwrap();
|
|
||||||
assert!(
|
|
||||||
alert.message.contains("major-session"),
|
|
||||||
"expected the top contributor to be named, got: {}",
|
|
||||||
alert.message
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn rate_anomaly_critical_3x() {
|
fn rate_anomaly_critical_3x() {
|
||||||
let mut det = WriteAnomalyDetector::new(cfg());
|
let mut det = WriteAnomalyDetector::new(cfg());
|
||||||
@@ -568,71 +395,6 @@ mod tests {
|
|||||||
assert!(alert.is_some());
|
assert!(alert.is_some());
|
||||||
}
|
}
|
||||||
|
|
||||||
// --- Pattern-match evasion hardening ---
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn pattern_defeats_extra_whitespace() {
|
|
||||||
let det = WriteAnomalyDetector::new(cfg());
|
|
||||||
let alert = det.check_pattern_anomaly("please ignore previous instructions");
|
|
||||||
assert!(alert.is_some(), "extra whitespace must not defeat matching");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn pattern_defeats_punctuation_splicing() {
|
|
||||||
let det = WriteAnomalyDetector::new(cfg());
|
|
||||||
let alert = det.check_pattern_anomaly("i.g.n.o.r.e p-r-e-v-i-o-u-s instructions");
|
|
||||||
assert!(
|
|
||||||
alert.is_some(),
|
|
||||||
"punctuation spliced between letters must not defeat matching"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn pattern_defeats_zero_width_space() {
|
|
||||||
let det = WriteAnomalyDetector::new(cfg());
|
|
||||||
// Zero-width space (U+200B) inserted mid-word.
|
|
||||||
let chunk = "ign\u{200B}ore previ\u{200B}ous instructions";
|
|
||||||
let alert = det.check_pattern_anomaly(chunk);
|
|
||||||
assert!(
|
|
||||||
alert.is_some(),
|
|
||||||
"zero-width space injection must not defeat matching"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn pattern_defeats_zero_width_joiner_and_bom() {
|
|
||||||
let det = WriteAnomalyDetector::new(cfg());
|
|
||||||
let chunk = "jail\u{200D}break\u{FEFF} attempt";
|
|
||||||
let alert = det.check_pattern_anomaly(chunk);
|
|
||||||
assert!(
|
|
||||||
alert.is_some(),
|
|
||||||
"ZWJ/BOM injection must not defeat matching"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn pattern_still_clean_after_normalization() {
|
|
||||||
let det = WriteAnomalyDetector::new(cfg());
|
|
||||||
// Normalization must not introduce false positives on ordinary text
|
|
||||||
// that merely contains punctuation and extra whitespace.
|
|
||||||
let alert =
|
|
||||||
det.check_pattern_anomaly("Well, I think... the weather is nice today, right?");
|
|
||||||
assert!(alert.is_none());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn normalize_for_pattern_match_examples() {
|
|
||||||
assert_eq!(
|
|
||||||
normalize_for_pattern_match("i.g.n.o.r.e p-r-e-v-i-o-u-s"),
|
|
||||||
"ignore previous"
|
|
||||||
);
|
|
||||||
assert_eq!(
|
|
||||||
normalize_for_pattern_match("ign\u{200B}ore previous"),
|
|
||||||
"ignore previous"
|
|
||||||
);
|
|
||||||
assert_eq!(normalize_for_pattern_match("SYSTEM:"), "system");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn pattern_jailbreak() {
|
fn pattern_jailbreak() {
|
||||||
let det = WriteAnomalyDetector::new(cfg());
|
let det = WriteAnomalyDetector::new(cfg());
|
||||||
|
|||||||
@@ -408,10 +408,6 @@ impl AsyncHDF5Memory {
|
|||||||
let (tx, rx) = oneshot::channel();
|
let (tx, rx) = oneshot::channel();
|
||||||
let _ = self.write_tx.send(WriteCmd::Shutdown(tx)).await;
|
let _ = self.write_tx.send(WriteCmd::Shutdown(tx)).await;
|
||||||
let _ = rx.await;
|
let _ = rx.await;
|
||||||
// The writer task has stopped, so nothing can write through this
|
|
||||||
// handle any more: release the single-writer lock now rather than at
|
|
||||||
// drop, so the store can be reopened while `self` is still in scope.
|
|
||||||
self.inner.lock().await.release_store_lock();
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -8,28 +8,7 @@
|
|||||||
//! - Sorted posting lists by doc_id for cache-friendly access
|
//! - Sorted posting lists by doc_id for cache-friendly access
|
||||||
//! - Block-Max WAND early termination
|
//! - Block-Max WAND early termination
|
||||||
|
|
||||||
use std::cmp::Reverse;
|
use std::collections::HashMap;
|
||||||
use std::collections::{BinaryHeap, HashMap};
|
|
||||||
|
|
||||||
/// `f32` wrapper providing a total order (via `total_cmp`) so BM25 scores can
|
|
||||||
/// be kept in a `BinaryHeap`. Scores are always finite in practice (no NaN
|
|
||||||
/// inputs reach this path), so `total_cmp`'s NaN ordering is never exercised.
|
|
||||||
#[derive(Debug, Clone, Copy, PartialEq)]
|
|
||||||
struct HeapScore(f32);
|
|
||||||
|
|
||||||
impl Eq for HeapScore {}
|
|
||||||
|
|
||||||
impl PartialOrd for HeapScore {
|
|
||||||
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
|
|
||||||
Some(self.cmp(other))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl Ord for HeapScore {
|
|
||||||
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
|
|
||||||
self.0.total_cmp(&other.0)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Default BM25 term-frequency saturation parameter.
|
/// Default BM25 term-frequency saturation parameter.
|
||||||
const DEFAULT_K1: f32 = 1.2;
|
const DEFAULT_K1: f32 = 1.2;
|
||||||
@@ -118,11 +97,9 @@ impl BM25Index {
|
|||||||
|
|
||||||
let total_max_contribution: f32 = max_tf_score.iter().sum();
|
let total_max_contribution: f32 = max_tf_score.iter().sum();
|
||||||
|
|
||||||
// Threshold for WAND early termination. `top_k_heap` is a min-heap of
|
// Threshold for WAND early termination
|
||||||
// size k (worst-of-the-top-k at the head) so it can be maintained in
|
|
||||||
// O(log k) per update instead of re-sorting the whole buffer.
|
|
||||||
let mut threshold = 0.0f32;
|
let mut threshold = 0.0f32;
|
||||||
let mut top_k_heap: BinaryHeap<Reverse<HeapScore>> = BinaryHeap::with_capacity(k);
|
let mut top_k_scores: Vec<f32> = Vec::with_capacity(k);
|
||||||
|
|
||||||
for (term_idx, (_, idf, postings)) in query_terms.iter().enumerate() {
|
for (term_idx, (_, idf, postings)) in query_terms.iter().enumerate() {
|
||||||
for &(doc_id, freq) in *postings {
|
for &(doc_id, freq) in *postings {
|
||||||
@@ -141,17 +118,24 @@ impl BM25Index {
|
|||||||
if term_idx == query_terms.len() - 1 {
|
if term_idx == query_terms.len() - 1 {
|
||||||
// Last term: check if this doc beats threshold
|
// Last term: check if this doc beats threshold
|
||||||
let final_score = *entry;
|
let final_score = *entry;
|
||||||
if top_k_heap.len() >= k {
|
if final_score > threshold && top_k_scores.len() >= k {
|
||||||
if final_score > threshold {
|
// Update threshold
|
||||||
// Replace the current worst-of-top-k.
|
top_k_scores
|
||||||
top_k_heap.pop();
|
.sort_by(|a, b| b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal));
|
||||||
top_k_heap.push(Reverse(HeapScore(final_score)));
|
if final_score > top_k_scores[k - 1] {
|
||||||
threshold = top_k_heap.peek().map(|Reverse(s)| s.0).unwrap_or(0.0);
|
top_k_scores[k - 1] = final_score;
|
||||||
|
top_k_scores.sort_by(|a, b| {
|
||||||
|
b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
|
||||||
|
});
|
||||||
|
threshold = top_k_scores[k - 1];
|
||||||
}
|
}
|
||||||
} else {
|
} else if top_k_scores.len() < k {
|
||||||
top_k_heap.push(Reverse(HeapScore(final_score)));
|
top_k_scores.push(final_score);
|
||||||
if top_k_heap.len() == k {
|
if top_k_scores.len() == k {
|
||||||
threshold = top_k_heap.peek().map(|Reverse(s)| s.0).unwrap_or(0.0);
|
top_k_scores.sort_by(|a, b| {
|
||||||
|
b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
|
||||||
|
});
|
||||||
|
threshold = top_k_scores[k - 1];
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -7,11 +7,6 @@ use crate::vector_search;
|
|||||||
pub struct MemoryCache {
|
pub struct MemoryCache {
|
||||||
pub chunks: Vec<String>,
|
pub chunks: Vec<String>,
|
||||||
pub embeddings: Vec<Vec<f32>>,
|
pub embeddings: Vec<Vec<f32>>,
|
||||||
/// `embeddings` flattened into one contiguous `[N × embedding_dim]`
|
|
||||||
/// buffer, maintained incrementally alongside `embeddings` (push/update/
|
|
||||||
/// compact) so BLAS/Accelerate batch search can read it directly instead
|
|
||||||
/// of re-flattening the whole corpus on every query.
|
|
||||||
pub embeddings_flat: Vec<f32>,
|
|
||||||
pub source_channels: Vec<String>,
|
pub source_channels: Vec<String>,
|
||||||
pub timestamps: Vec<f64>,
|
pub timestamps: Vec<f64>,
|
||||||
pub session_ids: Vec<String>,
|
pub session_ids: Vec<String>,
|
||||||
@@ -29,7 +24,6 @@ impl MemoryCache {
|
|||||||
Self {
|
Self {
|
||||||
chunks: Vec::new(),
|
chunks: Vec::new(),
|
||||||
embeddings: Vec::new(),
|
embeddings: Vec::new(),
|
||||||
embeddings_flat: Vec::new(),
|
|
||||||
source_channels: Vec::new(),
|
source_channels: Vec::new(),
|
||||||
timestamps: Vec::new(),
|
timestamps: Vec::new(),
|
||||||
session_ids: Vec::new(),
|
session_ids: Vec::new(),
|
||||||
@@ -41,17 +35,6 @@ impl MemoryCache {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Rebuild `embeddings_flat` from `embeddings` from scratch. Callers that
|
|
||||||
/// populate `embeddings` directly (bulk loads) must call this afterward.
|
|
||||||
pub fn rebuild_flat(&mut self) {
|
|
||||||
self.embeddings_flat.clear();
|
|
||||||
self.embeddings_flat
|
|
||||||
.reserve(self.embeddings.len() * self.embedding_dim);
|
|
||||||
for emb in &self.embeddings {
|
|
||||||
self.embeddings_flat.extend_from_slice(emb);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Total number of entries (including tombstoned).
|
/// Total number of entries (including tombstoned).
|
||||||
pub fn len(&self) -> usize {
|
pub fn len(&self) -> usize {
|
||||||
self.chunks.len()
|
self.chunks.len()
|
||||||
@@ -79,7 +62,6 @@ impl MemoryCache {
|
|||||||
let idx = self.chunks.len();
|
let idx = self.chunks.len();
|
||||||
let norm = vector_search::compute_norm(&embedding);
|
let norm = vector_search::compute_norm(&embedding);
|
||||||
self.chunks.push(chunk);
|
self.chunks.push(chunk);
|
||||||
self.embeddings_flat.extend_from_slice(&embedding);
|
|
||||||
self.embeddings.push(embedding);
|
self.embeddings.push(embedding);
|
||||||
self.source_channels.push(source_channel);
|
self.source_channels.push(source_channel);
|
||||||
self.timestamps.push(timestamp);
|
self.timestamps.push(timestamp);
|
||||||
@@ -118,20 +100,7 @@ impl MemoryCache {
|
|||||||
if idx < self.chunks.len() {
|
if idx < self.chunks.len() {
|
||||||
let norm = vector_search::compute_norm(&embedding);
|
let norm = vector_search::compute_norm(&embedding);
|
||||||
self.chunks[idx] = chunk;
|
self.chunks[idx] = chunk;
|
||||||
let dim = self.embedding_dim;
|
|
||||||
let flat_start = idx * dim;
|
|
||||||
let matches_dim =
|
|
||||||
embedding.len() == dim && flat_start + dim <= self.embeddings_flat.len();
|
|
||||||
self.embeddings[idx] = embedding;
|
self.embeddings[idx] = embedding;
|
||||||
if matches_dim {
|
|
||||||
self.embeddings_flat[flat_start..flat_start + dim]
|
|
||||||
.copy_from_slice(&self.embeddings[idx]);
|
|
||||||
} else {
|
|
||||||
// Embedding length doesn't match embedding_dim (shouldn't
|
|
||||||
// happen in practice) — fall back to a full rebuild rather
|
|
||||||
// than leave embeddings_flat misaligned with embeddings.
|
|
||||||
self.rebuild_flat();
|
|
||||||
}
|
|
||||||
self.source_channels[idx] = source_channel;
|
self.source_channels[idx] = source_channel;
|
||||||
self.timestamps[idx] = timestamp;
|
self.timestamps[idx] = timestamp;
|
||||||
self.session_ids[idx] = session_id;
|
self.session_ids[idx] = session_id;
|
||||||
@@ -204,125 +173,16 @@ impl MemoryCache {
|
|||||||
self.tombstones = new_tombstones;
|
self.tombstones = new_tombstones;
|
||||||
self.norms = new_norms;
|
self.norms = new_norms;
|
||||||
self.activation_weights = new_activation_weights;
|
self.activation_weights = new_activation_weights;
|
||||||
self.rebuild_flat();
|
|
||||||
|
|
||||||
(removed, index_map)
|
(removed, index_map)
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Flatten all embeddings into a single Vec<f32> for HDF5 storage.
|
/// Flatten all embeddings into a single Vec<f32> for HDF5 storage.
|
||||||
/// `embeddings_flat` is already maintained incrementally, so this just
|
|
||||||
/// clones it — kept as a method for callers that want an owned copy.
|
|
||||||
pub fn flat_embeddings(&self) -> Vec<f32> {
|
pub fn flat_embeddings(&self) -> Vec<f32> {
|
||||||
self.embeddings_flat.clone()
|
let mut flat = Vec::with_capacity(self.embeddings.len() * self.embedding_dim);
|
||||||
|
for emb in &self.embeddings {
|
||||||
|
flat.extend_from_slice(emb);
|
||||||
}
|
}
|
||||||
}
|
flat
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
/// `embeddings_flat` must always equal a from-scratch flatten of `embeddings`.
|
|
||||||
fn assert_flat_in_sync(cache: &MemoryCache) {
|
|
||||||
let expected: Vec<f32> = cache.embeddings.iter().flatten().copied().collect();
|
|
||||||
assert_eq!(cache.embeddings_flat, expected);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn push_keeps_flat_buffer_in_sync() {
|
|
||||||
let mut cache = MemoryCache::new(3);
|
|
||||||
cache.push(
|
|
||||||
"a".into(),
|
|
||||||
vec![1.0, 2.0, 3.0],
|
|
||||||
"chan".into(),
|
|
||||||
0.0,
|
|
||||||
"s1".into(),
|
|
||||||
String::new(),
|
|
||||||
);
|
|
||||||
cache.push(
|
|
||||||
"b".into(),
|
|
||||||
vec![4.0, 5.0, 6.0],
|
|
||||||
"chan".into(),
|
|
||||||
1.0,
|
|
||||||
"s1".into(),
|
|
||||||
String::new(),
|
|
||||||
);
|
|
||||||
assert_flat_in_sync(&cache);
|
|
||||||
assert_eq!(cache.embeddings_flat, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn update_keeps_flat_buffer_in_sync() {
|
|
||||||
let mut cache = MemoryCache::new(3);
|
|
||||||
cache.push(
|
|
||||||
"a".into(),
|
|
||||||
vec![1.0, 2.0, 3.0],
|
|
||||||
"chan".into(),
|
|
||||||
0.0,
|
|
||||||
"s1".into(),
|
|
||||||
String::new(),
|
|
||||||
);
|
|
||||||
cache.push(
|
|
||||||
"b".into(),
|
|
||||||
vec![4.0, 5.0, 6.0],
|
|
||||||
"chan".into(),
|
|
||||||
1.0,
|
|
||||||
"s1".into(),
|
|
||||||
String::new(),
|
|
||||||
);
|
|
||||||
cache.update(
|
|
||||||
0,
|
|
||||||
"a2".into(),
|
|
||||||
vec![7.0, 8.0, 9.0],
|
|
||||||
"chan".into(),
|
|
||||||
2.0,
|
|
||||||
"s1".into(),
|
|
||||||
);
|
|
||||||
assert_flat_in_sync(&cache);
|
|
||||||
assert_eq!(
|
|
||||||
cache.embeddings_flat,
|
|
||||||
vec![7.0, 8.0, 9.0, 4.0, 5.0, 6.0],
|
|
||||||
"update must overwrite the correct flat slice, not just append"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn compact_keeps_flat_buffer_in_sync() {
|
|
||||||
let mut cache = MemoryCache::new(2);
|
|
||||||
cache.push(
|
|
||||||
"a".into(),
|
|
||||||
vec![1.0, 1.0],
|
|
||||||
"chan".into(),
|
|
||||||
0.0,
|
|
||||||
"s1".into(),
|
|
||||||
String::new(),
|
|
||||||
);
|
|
||||||
cache.push(
|
|
||||||
"b".into(),
|
|
||||||
vec![2.0, 2.0],
|
|
||||||
"chan".into(),
|
|
||||||
1.0,
|
|
||||||
"s1".into(),
|
|
||||||
String::new(),
|
|
||||||
);
|
|
||||||
cache.push(
|
|
||||||
"c".into(),
|
|
||||||
vec![3.0, 3.0],
|
|
||||||
"chan".into(),
|
|
||||||
2.0,
|
|
||||||
"s1".into(),
|
|
||||||
String::new(),
|
|
||||||
);
|
|
||||||
cache.mark_deleted(1);
|
|
||||||
cache.compact();
|
|
||||||
assert_flat_in_sync(&cache);
|
|
||||||
assert_eq!(cache.embeddings_flat, vec![1.0, 1.0, 3.0, 3.0]);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn rebuild_flat_matches_manual_flatten() {
|
|
||||||
let mut cache = MemoryCache::new(2);
|
|
||||||
cache.embeddings = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
|
|
||||||
cache.rebuild_flat();
|
|
||||||
assert_eq!(cache.embeddings_flat, vec![1.0, 2.0, 3.0, 4.0]);
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -16,55 +16,6 @@ pub enum MemorySource {
|
|||||||
Correction,
|
Correction,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Source classification for content whose true origin is *not*
|
|
||||||
/// independently verified by the caller of [`ConsolidationEngine::add_memory`]
|
|
||||||
/// — arbitrary text forwarded from a user, a tool's output, or a retrieval
|
|
||||||
/// pipeline. This is the only source set `add_memory` accepts; it cannot
|
|
||||||
/// claim the `System`/`Correction` importance boost (see [`TrustedSource`]
|
|
||||||
/// and [`ConsolidationEngine::add_trusted_memory`]) — a caller passing
|
|
||||||
/// through untrusted content has no way to self-report an elevated trust
|
|
||||||
/// level through this entry point.
|
|
||||||
#[derive(Clone, Debug, PartialEq)]
|
|
||||||
pub enum UntrustedSource {
|
|
||||||
User,
|
|
||||||
Tool,
|
|
||||||
Retrieval,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl From<UntrustedSource> for MemorySource {
|
|
||||||
fn from(s: UntrustedSource) -> Self {
|
|
||||||
match s {
|
|
||||||
UntrustedSource::User => MemorySource::User,
|
|
||||||
UntrustedSource::Tool => MemorySource::Tool,
|
|
||||||
UntrustedSource::Retrieval => MemorySource::Retrieval,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Source classification for content whose elevated trust level has been
|
|
||||||
/// independently verified by the caller — e.g. the library's own
|
|
||||||
/// system-generated text, or a caller that ran its own correction-cue
|
|
||||||
/// detection (as `memory_strategy::SaveOnUserCorrection` does) rather than
|
|
||||||
/// forwarding a caller-supplied label verbatim. `MemorySource::System`/
|
|
||||||
/// `Correction` get elevated importance weighting in
|
|
||||||
/// [`ImportanceScorer::score_correction`]; only reachable through
|
|
||||||
/// [`ConsolidationEngine::add_trusted_memory`], a distinct entry point from
|
|
||||||
/// the one untrusted content is passed through.
|
|
||||||
#[derive(Clone, Debug, PartialEq)]
|
|
||||||
pub enum TrustedSource {
|
|
||||||
System,
|
|
||||||
Correction,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl From<TrustedSource> for MemorySource {
|
|
||||||
fn from(s: TrustedSource) -> Self {
|
|
||||||
match s {
|
|
||||||
TrustedSource::System => MemorySource::System,
|
|
||||||
TrustedSource::Correction => MemorySource::Correction,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Clone, Debug, PartialEq)]
|
#[derive(Clone, Debug, PartialEq)]
|
||||||
pub enum MemoryTier {
|
pub enum MemoryTier {
|
||||||
Working,
|
Working,
|
||||||
@@ -167,7 +118,7 @@ impl ImportanceScorer {
|
|||||||
|
|
||||||
/// Novelty score: 1.0 − max cosine similarity against all existing records.
|
/// Novelty score: 1.0 − max cosine similarity against all existing records.
|
||||||
/// Returns 1.0 when there are no existing memories.
|
/// Returns 1.0 when there are no existing memories.
|
||||||
pub fn score_surprise(embedding: &[f32], existing_memories: &[&MemoryRecord]) -> f32 {
|
pub fn score_surprise(embedding: &[f32], existing_memories: &[MemoryRecord]) -> f32 {
|
||||||
if existing_memories.is_empty() {
|
if existing_memories.is_empty() {
|
||||||
return 1.0;
|
return 1.0;
|
||||||
}
|
}
|
||||||
@@ -248,51 +199,21 @@ impl ConsolidationEngine {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Add a new memory to the Working tier from an untrusted/ordinary origin
|
/// Add a new memory to the Working tier.
|
||||||
/// (User, Tool, or Retrieval). This is the entry point for arbitrary
|
|
||||||
/// caller-supplied content — it cannot claim the elevated System/
|
|
||||||
/// Correction importance boost. Use [`Self::add_trusted_memory`] for
|
|
||||||
/// content whose elevated trust level the caller has independently
|
|
||||||
/// verified.
|
|
||||||
///
|
///
|
||||||
/// Importance is scored against existing Working-tier records only.
|
/// Importance is scored against existing Working-tier records only.
|
||||||
pub fn add_memory(
|
pub fn add_memory(
|
||||||
&mut self,
|
|
||||||
chunk: String,
|
|
||||||
embedding: Vec<f32>,
|
|
||||||
source: UntrustedSource,
|
|
||||||
now: f64,
|
|
||||||
) -> u64 {
|
|
||||||
self.add_memory_with_source(chunk, embedding, source.into(), now)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Add a new memory tagged System or Correction, which get elevated
|
|
||||||
/// importance weighting in [`ImportanceScorer::score_correction`]. Only
|
|
||||||
/// call this from code that has independently verified the origin (the
|
|
||||||
/// library's own system-generated text, or a caller that ran its own
|
|
||||||
/// correction-cue detection) — never from a path that forwards a
|
|
||||||
/// caller-supplied trust label verbatim.
|
|
||||||
pub fn add_trusted_memory(
|
|
||||||
&mut self,
|
|
||||||
chunk: String,
|
|
||||||
embedding: Vec<f32>,
|
|
||||||
source: TrustedSource,
|
|
||||||
now: f64,
|
|
||||||
) -> u64 {
|
|
||||||
self.add_memory_with_source(chunk, embedding, source.into(), now)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn add_memory_with_source(
|
|
||||||
&mut self,
|
&mut self,
|
||||||
chunk: String,
|
chunk: String,
|
||||||
embedding: Vec<f32>,
|
embedding: Vec<f32>,
|
||||||
source: MemorySource,
|
source: MemorySource,
|
||||||
now: f64,
|
now: f64,
|
||||||
) -> u64 {
|
) -> u64 {
|
||||||
let working: Vec<&MemoryRecord> = self
|
let working: Vec<MemoryRecord> = self
|
||||||
.records
|
.records
|
||||||
.iter()
|
.iter()
|
||||||
.filter(|r| r.tier == MemoryTier::Working)
|
.filter(|r| r.tier == MemoryTier::Working)
|
||||||
|
.cloned()
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
let surprise = ImportanceScorer::score_surprise(&embedding, &working);
|
let surprise = ImportanceScorer::score_surprise(&embedding, &working);
|
||||||
@@ -360,7 +281,7 @@ impl ConsolidationEngine {
|
|||||||
if working_count > capacity {
|
if working_count > capacity {
|
||||||
let evict_n = working_count - capacity;
|
let evict_n = working_count - capacity;
|
||||||
// Collect the ids of the records to evict (lowest decay = first in sorted list).
|
// Collect the ids of the records to evict (lowest decay = first in sorted list).
|
||||||
let evict_ids: std::collections::HashSet<u64> = working_indices[..evict_n]
|
let evict_ids: Vec<u64> = working_indices[..evict_n]
|
||||||
.iter()
|
.iter()
|
||||||
.map(|&i| self.records[i].id)
|
.map(|&i| self.records[i].id)
|
||||||
.collect();
|
.collect();
|
||||||
@@ -421,7 +342,7 @@ impl ConsolidationEngine {
|
|||||||
});
|
});
|
||||||
|
|
||||||
let evict_n = episodic_count - episodic_capacity;
|
let evict_n = episodic_count - episodic_capacity;
|
||||||
let evict_ids: std::collections::HashSet<u64> = episodic_indices[..evict_n]
|
let evict_ids: Vec<u64> = episodic_indices[..evict_n]
|
||||||
.iter()
|
.iter()
|
||||||
.map(|&i| self.records[i].id)
|
.map(|&i| self.records[i].id)
|
||||||
.collect();
|
.collect();
|
||||||
@@ -498,44 +419,13 @@ mod tests {
|
|||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
// 2. Add memory — basic
|
// 2. Add memory — basic
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
/// add_trusted_memory(TrustedSource::Correction) must actually produce a
|
|
||||||
/// MemorySource::Correction record — the only way to reach that elevated
|
|
||||||
/// classification, since add_memory's UntrustedSource has no such variant.
|
|
||||||
#[test]
|
|
||||||
fn test_add_trusted_memory_sets_correction_source() {
|
|
||||||
let mut engine = ConsolidationEngine::new(ConsolidationConfig::default());
|
|
||||||
let id = engine.add_trusted_memory(
|
|
||||||
"verified correction".to_string(),
|
|
||||||
unit_vec(4, 0),
|
|
||||||
TrustedSource::Correction,
|
|
||||||
0.0,
|
|
||||||
);
|
|
||||||
let rec = engine.get_by_id(id).unwrap();
|
|
||||||
assert_eq!(rec.source, MemorySource::Correction);
|
|
||||||
}
|
|
||||||
|
|
||||||
/// add_trusted_memory(TrustedSource::System) must produce a
|
|
||||||
/// MemorySource::System record.
|
|
||||||
#[test]
|
|
||||||
fn test_add_trusted_memory_sets_system_source() {
|
|
||||||
let mut engine = ConsolidationEngine::new(ConsolidationConfig::default());
|
|
||||||
let id = engine.add_trusted_memory(
|
|
||||||
"bootstrap text".to_string(),
|
|
||||||
unit_vec(4, 0),
|
|
||||||
TrustedSource::System,
|
|
||||||
0.0,
|
|
||||||
);
|
|
||||||
let rec = engine.get_by_id(id).unwrap();
|
|
||||||
assert_eq!(rec.source, MemorySource::System);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_add_memory_basic() {
|
fn test_add_memory_basic() {
|
||||||
let mut engine = ConsolidationEngine::new(ConsolidationConfig::default());
|
let mut engine = ConsolidationEngine::new(ConsolidationConfig::default());
|
||||||
let id = engine.add_memory(
|
let id = engine.add_memory(
|
||||||
"Hello world".to_string(),
|
"Hello world".to_string(),
|
||||||
unit_vec(4, 0),
|
unit_vec(4, 0),
|
||||||
UntrustedSource::User,
|
MemorySource::User,
|
||||||
1_000_000.0,
|
1_000_000.0,
|
||||||
);
|
);
|
||||||
assert_eq!(id, 0);
|
assert_eq!(id, 0);
|
||||||
@@ -563,7 +453,7 @@ mod tests {
|
|||||||
#[test]
|
#[test]
|
||||||
fn test_importance_scorer_surprise_identical() {
|
fn test_importance_scorer_surprise_identical() {
|
||||||
let emb = unit_vec(4, 0);
|
let emb = unit_vec(4, 0);
|
||||||
let existing = [MemoryRecord {
|
let existing = vec![MemoryRecord {
|
||||||
id: 0,
|
id: 0,
|
||||||
chunk: "existing".to_string(),
|
chunk: "existing".to_string(),
|
||||||
embedding: emb.clone(),
|
embedding: emb.clone(),
|
||||||
@@ -574,8 +464,7 @@ mod tests {
|
|||||||
created_at: 0.0,
|
created_at: 0.0,
|
||||||
source: MemorySource::User,
|
source: MemorySource::User,
|
||||||
}];
|
}];
|
||||||
let existing_refs: Vec<&MemoryRecord> = existing.iter().collect();
|
let score = ImportanceScorer::score_surprise(&emb, &existing);
|
||||||
let score = ImportanceScorer::score_surprise(&emb, &existing_refs);
|
|
||||||
assert!(score < 0.01, "expected ~0.0, got {score}");
|
assert!(score < 0.01, "expected ~0.0, got {score}");
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -603,20 +492,23 @@ mod tests {
|
|||||||
fn test_importance_scorer_length() {
|
fn test_importance_scorer_length() {
|
||||||
assert!((ImportanceScorer::score_length("")).abs() < f32::EPSILON);
|
assert!((ImportanceScorer::score_length("")).abs() < f32::EPSILON);
|
||||||
// 50 words → 0.5
|
// 50 words → 0.5
|
||||||
let fifty_words = std::iter::repeat_n("word", 50)
|
let fifty_words = std::iter::repeat("word")
|
||||||
|
.take(50)
|
||||||
.collect::<Vec<_>>()
|
.collect::<Vec<_>>()
|
||||||
.join(" ");
|
.join(" ");
|
||||||
let s50 = ImportanceScorer::score_length(&fifty_words);
|
let s50 = ImportanceScorer::score_length(&fifty_words);
|
||||||
assert!((s50 - 0.5).abs() < 1e-5, "expected 0.5, got {s50}");
|
assert!((s50 - 0.5).abs() < 1e-5, "expected 0.5, got {s50}");
|
||||||
|
|
||||||
// 100 words → 1.0
|
// 100 words → 1.0
|
||||||
let hundred_words = std::iter::repeat_n("word", 100)
|
let hundred_words = std::iter::repeat("word")
|
||||||
|
.take(100)
|
||||||
.collect::<Vec<_>>()
|
.collect::<Vec<_>>()
|
||||||
.join(" ");
|
.join(" ");
|
||||||
assert_eq!(ImportanceScorer::score_length(&hundred_words), 1.0);
|
assert_eq!(ImportanceScorer::score_length(&hundred_words), 1.0);
|
||||||
|
|
||||||
// 200 words → still 1.0 (clamped)
|
// 200 words → still 1.0 (clamped)
|
||||||
let two_hundred = std::iter::repeat_n("word", 200)
|
let two_hundred = std::iter::repeat("word")
|
||||||
|
.take(200)
|
||||||
.collect::<Vec<_>>()
|
.collect::<Vec<_>>()
|
||||||
.join(" ");
|
.join(" ");
|
||||||
assert_eq!(ImportanceScorer::score_length(&two_hundred), 1.0);
|
assert_eq!(ImportanceScorer::score_length(&two_hundred), 1.0);
|
||||||
@@ -690,11 +582,9 @@ mod tests {
|
|||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
#[test]
|
#[test]
|
||||||
fn test_consolidate_eviction_working() {
|
fn test_consolidate_eviction_working() {
|
||||||
let cfg = ConsolidationConfig {
|
let mut cfg = ConsolidationConfig::default();
|
||||||
working_capacity: 3,
|
cfg.working_capacity = 3;
|
||||||
working_to_episodic_threshold: 2.0, // never promote in this test
|
cfg.working_to_episodic_threshold = 2.0; // never promote in this test
|
||||||
..Default::default()
|
|
||||||
};
|
|
||||||
let mut engine = ConsolidationEngine::new(cfg);
|
let mut engine = ConsolidationEngine::new(cfg);
|
||||||
|
|
||||||
// Add 5 records; all have very low importance so none get promoted.
|
// Add 5 records; all have very low importance so none get promoted.
|
||||||
@@ -702,7 +592,7 @@ mod tests {
|
|||||||
let id = engine.add_memory(
|
let id = engine.add_memory(
|
||||||
"x".to_string(),
|
"x".to_string(),
|
||||||
unit_vec(4, i as usize),
|
unit_vec(4, i as usize),
|
||||||
UntrustedSource::User,
|
MemorySource::User,
|
||||||
i as f64,
|
i as f64,
|
||||||
);
|
);
|
||||||
// Force low importance so promotion threshold is not crossed.
|
// Force low importance so promotion threshold is not crossed.
|
||||||
@@ -735,10 +625,10 @@ mod tests {
|
|||||||
let cfg = ConsolidationConfig::default();
|
let cfg = ConsolidationConfig::default();
|
||||||
let mut engine = ConsolidationEngine::new(cfg);
|
let mut engine = ConsolidationEngine::new(cfg);
|
||||||
|
|
||||||
let id = engine.add_trusted_memory(
|
let id = engine.add_memory(
|
||||||
"important memory".to_string(),
|
"important memory".to_string(),
|
||||||
unit_vec(4, 0),
|
unit_vec(4, 0),
|
||||||
TrustedSource::Correction,
|
MemorySource::Correction,
|
||||||
0.0,
|
0.0,
|
||||||
);
|
);
|
||||||
// Force importance above threshold.
|
// Force importance above threshold.
|
||||||
@@ -771,7 +661,7 @@ mod tests {
|
|||||||
let id = engine.add_memory(
|
let id = engine.add_memory(
|
||||||
"frequently accessed".to_string(),
|
"frequently accessed".to_string(),
|
||||||
unit_vec(4, 0),
|
unit_vec(4, 0),
|
||||||
UntrustedSource::User,
|
MemorySource::User,
|
||||||
0.0,
|
0.0,
|
||||||
);
|
);
|
||||||
|
|
||||||
@@ -799,12 +689,7 @@ mod tests {
|
|||||||
#[test]
|
#[test]
|
||||||
fn test_access_memory_reactivation() {
|
fn test_access_memory_reactivation() {
|
||||||
let mut engine = ConsolidationEngine::new(ConsolidationConfig::default());
|
let mut engine = ConsolidationEngine::new(ConsolidationConfig::default());
|
||||||
let id = engine.add_memory(
|
let id = engine.add_memory("chunk".to_string(), unit_vec(4, 0), MemorySource::User, 0.0);
|
||||||
"chunk".to_string(),
|
|
||||||
unit_vec(4, 0),
|
|
||||||
UntrustedSource::User,
|
|
||||||
0.0,
|
|
||||||
);
|
|
||||||
|
|
||||||
engine.access_memory(id, 5000.0);
|
engine.access_memory(id, 5000.0);
|
||||||
let rec = engine.get_by_id(id).unwrap();
|
let rec = engine.get_by_id(id).unwrap();
|
||||||
@@ -825,11 +710,11 @@ mod tests {
|
|||||||
let mut engine = ConsolidationEngine::new(ConsolidationConfig::default());
|
let mut engine = ConsolidationEngine::new(ConsolidationConfig::default());
|
||||||
|
|
||||||
// 2 Working
|
// 2 Working
|
||||||
engine.add_memory("w1".to_string(), unit_vec(4, 0), UntrustedSource::User, 0.0);
|
engine.add_memory("w1".to_string(), unit_vec(4, 0), MemorySource::User, 0.0);
|
||||||
engine.add_memory("w2".to_string(), unit_vec(4, 1), UntrustedSource::User, 0.0);
|
engine.add_memory("w2".to_string(), unit_vec(4, 1), MemorySource::User, 0.0);
|
||||||
|
|
||||||
// 1 Episodic (manually set)
|
// 1 Episodic (manually set)
|
||||||
let id_e = engine.add_memory("e1".to_string(), unit_vec(4, 2), UntrustedSource::User, 0.0);
|
let id_e = engine.add_memory("e1".to_string(), unit_vec(4, 2), MemorySource::User, 0.0);
|
||||||
engine
|
engine
|
||||||
.records
|
.records
|
||||||
.iter_mut()
|
.iter_mut()
|
||||||
@@ -838,7 +723,7 @@ mod tests {
|
|||||||
.tier = MemoryTier::Episodic;
|
.tier = MemoryTier::Episodic;
|
||||||
|
|
||||||
// 1 Semantic (manually set)
|
// 1 Semantic (manually set)
|
||||||
let id_s = engine.add_memory("s1".to_string(), unit_vec(4, 3), UntrustedSource::User, 0.0);
|
let id_s = engine.add_memory("s1".to_string(), unit_vec(4, 3), MemorySource::User, 0.0);
|
||||||
engine
|
engine
|
||||||
.records
|
.records
|
||||||
.iter_mut()
|
.iter_mut()
|
||||||
@@ -857,11 +742,9 @@ mod tests {
|
|||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
#[test]
|
#[test]
|
||||||
fn test_consolidate_episodic_eviction() {
|
fn test_consolidate_episodic_eviction() {
|
||||||
let cfg = ConsolidationConfig {
|
let mut cfg = ConsolidationConfig::default();
|
||||||
episodic_capacity: 3,
|
cfg.episodic_capacity = 3;
|
||||||
working_to_episodic_threshold: 2.0, // never auto-promote from Working
|
cfg.working_to_episodic_threshold = 2.0; // never auto-promote from Working
|
||||||
..Default::default()
|
|
||||||
};
|
|
||||||
let mut engine = ConsolidationEngine::new(cfg);
|
let mut engine = ConsolidationEngine::new(cfg);
|
||||||
|
|
||||||
// Seed 5 records directly in Episodic.
|
// Seed 5 records directly in Episodic.
|
||||||
@@ -869,7 +752,7 @@ mod tests {
|
|||||||
let id = engine.add_memory(
|
let id = engine.add_memory(
|
||||||
"episodic chunk".to_string(),
|
"episodic chunk".to_string(),
|
||||||
unit_vec(4, i as usize),
|
unit_vec(4, i as usize),
|
||||||
UntrustedSource::User,
|
MemorySource::User,
|
||||||
i as f64,
|
i as f64,
|
||||||
);
|
);
|
||||||
let rec = engine.records.iter_mut().find(|r| r.id == id).unwrap();
|
let rec = engine.records.iter_mut().find(|r| r.id == id).unwrap();
|
||||||
|
|||||||
@@ -777,10 +777,8 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_tech_disabled() {
|
fn test_tech_disabled() {
|
||||||
let config = ExtractorConfig {
|
let mut config = ExtractorConfig::default();
|
||||||
extract_technology: false,
|
config.extract_technology = false;
|
||||||
..Default::default()
|
|
||||||
};
|
|
||||||
let e = EntityExtractor::new(config);
|
let e = EntityExtractor::new(config);
|
||||||
let entities = e.extract("We use Rust and Docker.");
|
let entities = e.extract("We use Rust and Docker.");
|
||||||
assert!(
|
assert!(
|
||||||
@@ -849,10 +847,8 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_date_disabled() {
|
fn test_date_disabled() {
|
||||||
let config = ExtractorConfig {
|
let mut config = ExtractorConfig::default();
|
||||||
extract_dates: false,
|
config.extract_dates = false;
|
||||||
..Default::default()
|
|
||||||
};
|
|
||||||
let e = EntityExtractor::new(config);
|
let e = EntityExtractor::new(config);
|
||||||
let entities = e.extract("Released on 2024-03-19.");
|
let entities = e.extract("Released on 2024-03-19.");
|
||||||
assert!(
|
assert!(
|
||||||
@@ -985,10 +981,8 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_confidence_filter() {
|
fn test_confidence_filter() {
|
||||||
let config = ExtractorConfig {
|
let mut config = ExtractorConfig::default();
|
||||||
min_confidence: 0.95,
|
config.min_confidence = 0.95;
|
||||||
..Default::default()
|
|
||||||
};
|
|
||||||
let e = EntityExtractor::new(config);
|
let e = EntityExtractor::new(config);
|
||||||
// Only dates (0.95) and techs (0.9) should survive; 0.9 < 0.95 filters techs.
|
// Only dates (0.95) and techs (0.9) should survive; 0.9 < 0.95 filters techs.
|
||||||
let entities = e.extract("We use Rust since 2024-01-01.");
|
let entities = e.extract("We use Rust since 2024-01-01.");
|
||||||
@@ -1008,7 +1002,7 @@ mod tests {
|
|||||||
fn test_batch_dedup() {
|
fn test_batch_dedup() {
|
||||||
let e = default_extractor();
|
let e = default_extractor();
|
||||||
let texts = ["We use Rust.", "Rust is fast.", "Also Rust for safety."];
|
let texts = ["We use Rust.", "Rust is fast.", "Also Rust for safety."];
|
||||||
let entities = e.extract_batch(&texts);
|
let entities = e.extract_batch(&texts.iter().map(|s| *s).collect::<Vec<_>>());
|
||||||
let rust_count = entities.iter().filter(|x| x.text == "Rust").count();
|
let rust_count = entities.iter().filter(|x| x.text == "Rust").count();
|
||||||
assert_eq!(rust_count, 1, "Rust should appear exactly once after dedup");
|
assert_eq!(rust_count, 1, "Rust should appear exactly once after dedup");
|
||||||
}
|
}
|
||||||
@@ -1017,7 +1011,7 @@ mod tests {
|
|||||||
fn test_batch_multiple_types() {
|
fn test_batch_multiple_types() {
|
||||||
let e = default_extractor();
|
let e = default_extractor();
|
||||||
let texts = ["Deploy with Docker.", "We merged last week."];
|
let texts = ["Deploy with Docker.", "We merged last week."];
|
||||||
let entities = e.extract_batch(&texts);
|
let entities = e.extract_batch(&texts.iter().map(|s| *s).collect::<Vec<_>>());
|
||||||
assert!(
|
assert!(
|
||||||
entities
|
entities
|
||||||
.iter()
|
.iter()
|
||||||
|
|||||||
@@ -91,22 +91,14 @@ pub fn merge_vector_keyword(
|
|||||||
}
|
}
|
||||||
|
|
||||||
let mut results: Vec<(usize, f32)> = merged.into_iter().collect();
|
let mut results: Vec<(usize, f32)> = merged.into_iter().collect();
|
||||||
// Index tie-break: `merged` is a HashMap, so without it the ties that
|
results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
|
||||||
// survive `truncate` differ from run to run.
|
|
||||||
results.sort_by(|a, b| {
|
|
||||||
b.1.partial_cmp(&a.1)
|
|
||||||
.unwrap_or(std::cmp::Ordering::Equal)
|
|
||||||
.then(a.0.cmp(&b.0))
|
|
||||||
});
|
|
||||||
results.truncate(k);
|
results.truncate(k);
|
||||||
results
|
results
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Normalize a set of scores to the [0, 1] range using min-max normalization.
|
/// Normalize a set of scores to the [0, 1] range using min-max normalization.
|
||||||
///
|
///
|
||||||
/// If all scores are identical there is no spread to normalise: each entry
|
/// If all scores are identical, returns 0.0 for each entry.
|
||||||
/// gets 1.0 when that score is positive (all equally the best match) and 0.0
|
|
||||||
/// otherwise (nothing matched).
|
|
||||||
fn normalize_scores(scores: &[(usize, f32)]) -> Vec<(usize, f32)> {
|
fn normalize_scores(scores: &[(usize, f32)]) -> Vec<(usize, f32)> {
|
||||||
if scores.is_empty() {
|
if scores.is_empty() {
|
||||||
return Vec::new();
|
return Vec::new();
|
||||||
@@ -120,13 +112,7 @@ fn normalize_scores(scores: &[(usize, f32)]) -> Vec<(usize, f32)> {
|
|||||||
|
|
||||||
let range = max - min;
|
let range = max - min;
|
||||||
if range == 0.0 {
|
if range == 0.0 {
|
||||||
// All candidates scored the same (including the single-candidate
|
return scores.iter().map(|(idx, _)| (*idx, 0.0)).collect();
|
||||||
// case), so min-max has no spread to work with. They are all equally
|
|
||||||
// the best match if that score is positive, and all non-matches
|
|
||||||
// otherwise. This used to return 0.0 unconditionally, which erased a
|
|
||||||
// lone perfect match from the fused score.
|
|
||||||
let level = if max > 0.0 { 1.0 } else { 0.0 };
|
|
||||||
return scores.iter().map(|(idx, _)| (*idx, level)).collect();
|
|
||||||
}
|
}
|
||||||
|
|
||||||
scores
|
scores
|
||||||
@@ -338,18 +324,10 @@ mod tests {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn normalize_scores_single() {
|
fn normalize_scores_single() {
|
||||||
// A lone positive score is the best match there is, not a non-match.
|
|
||||||
let result = normalize_scores(&[(0, 5.0)]);
|
let result = normalize_scores(&[(0, 5.0)]);
|
||||||
assert_eq!(result.len(), 1);
|
assert_eq!(result.len(), 1);
|
||||||
assert_eq!(result[0].1, 1.0);
|
// Single score normalizes to 0.0 (range is 0)
|
||||||
}
|
assert_eq!(result[0].1, 0.0);
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn normalize_scores_all_equal() {
|
|
||||||
let matched = normalize_scores(&[(0, 0.4), (1, 0.4)]);
|
|
||||||
assert!(matched.iter().all(|(_, s)| *s == 1.0));
|
|
||||||
let unmatched = normalize_scores(&[(0, 0.0), (1, 0.0)]);
|
|
||||||
assert!(unmatched.iter().all(|(_, s)| *s == 0.0));
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
|
|||||||
@@ -50,9 +50,6 @@ impl RelationType {
|
|||||||
pub struct Entity {
|
pub struct Entity {
|
||||||
pub id: u64,
|
pub id: u64,
|
||||||
pub name: String,
|
pub name: String,
|
||||||
/// Lowercased `name`, cached at construction time to avoid re-allocating
|
|
||||||
/// and re-lowercasing on every entity-resolution scan.
|
|
||||||
pub name_lower: String,
|
|
||||||
pub entity_type: String,
|
pub entity_type: String,
|
||||||
/// Index into the memory embeddings array, or -1 if none.
|
/// Index into the memory embeddings array, or -1 if none.
|
||||||
pub embedding_idx: i64,
|
pub embedding_idx: i64,
|
||||||
@@ -72,7 +69,6 @@ impl Default for Entity {
|
|||||||
Self {
|
Self {
|
||||||
id: 0,
|
id: 0,
|
||||||
name: String::new(),
|
name: String::new(),
|
||||||
name_lower: String::new(),
|
|
||||||
entity_type: String::new(),
|
entity_type: String::new(),
|
||||||
embedding_idx: -1,
|
embedding_idx: -1,
|
||||||
properties: HashMap::new(),
|
properties: HashMap::new(),
|
||||||
@@ -155,55 +151,6 @@ fn levenshtein(a: &str, b: &str) -> usize {
|
|||||||
prev[nb]
|
prev[nb]
|
||||||
}
|
}
|
||||||
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
// AdjacencyIndex
|
|
||||||
// ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
/// Adjacency index over a snapshot of `entities`/`relations`: an entity-id ->
|
|
||||||
/// entities-slice-index map, and an entity-id -> relation-indices map (edges
|
|
||||||
/// touching that entity as either source or target).
|
|
||||||
///
|
|
||||||
/// Built fresh per traversal call rather than cached on `KnowledgeCache`:
|
|
||||||
/// entities/relations are plain `pub` `Vec`s that get pushed to directly
|
|
||||||
/// (e.g. `schema.rs`'s load path bypasses `add_entity`/`add_relation`), so a
|
|
||||||
/// persistent index would need extra bookkeeping to avoid drifting stale. A
|
|
||||||
/// one-off O(V+E) build per call is still a large win over the O(V·E) (BFS)
|
|
||||||
/// / O(steps·active·E) (spreading activation) scans it replaces.
|
|
||||||
struct AdjacencyIndex {
|
|
||||||
entity_index: HashMap<u64, usize>,
|
|
||||||
by_entity: HashMap<u64, Vec<usize>>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl AdjacencyIndex {
|
|
||||||
fn build(entities: &[Entity], relations: &[Relation]) -> Self {
|
|
||||||
let mut entity_index = HashMap::with_capacity(entities.len());
|
|
||||||
for (i, e) in entities.iter().enumerate() {
|
|
||||||
entity_index.insert(e.id, i);
|
|
||||||
}
|
|
||||||
|
|
||||||
let mut by_entity: HashMap<u64, Vec<usize>> = HashMap::new();
|
|
||||||
for (i, r) in relations.iter().enumerate() {
|
|
||||||
by_entity.entry(r.src).or_default().push(i);
|
|
||||||
if r.tgt != r.src {
|
|
||||||
by_entity.entry(r.tgt).or_default().push(i);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
Self {
|
|
||||||
entity_index,
|
|
||||||
by_entity,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Indices into `relations` of every edge touching `entity_id`.
|
|
||||||
fn relations_touching(&self, entity_id: u64) -> &[usize] {
|
|
||||||
self.by_entity
|
|
||||||
.get(&entity_id)
|
|
||||||
.map(|v| v.as_slice())
|
|
||||||
.unwrap_or(&[])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
// KnowledgeCache
|
// KnowledgeCache
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
@@ -251,7 +198,6 @@ impl KnowledgeCache {
|
|||||||
self.entities.push(Entity {
|
self.entities.push(Entity {
|
||||||
id,
|
id,
|
||||||
name: name.to_owned(),
|
name: name.to_owned(),
|
||||||
name_lower: name.to_lowercase(),
|
|
||||||
entity_type: entity_type.to_owned(),
|
entity_type: entity_type.to_owned(),
|
||||||
embedding_idx,
|
embedding_idx,
|
||||||
properties: HashMap::new(),
|
properties: HashMap::new(),
|
||||||
@@ -364,22 +310,16 @@ impl KnowledgeCache {
|
|||||||
) -> (u64, bool) {
|
) -> (u64, bool) {
|
||||||
let lower_name = name.to_lowercase();
|
let lower_name = name.to_lowercase();
|
||||||
|
|
||||||
// Search for the closest existing entity, short-circuiting on an
|
// Search for the closest existing entity.
|
||||||
// exact match since no closer candidate can exist.
|
let best = self
|
||||||
let mut best: Option<(u64, usize)> = None;
|
.entities
|
||||||
for e in &self.entities {
|
.iter()
|
||||||
let dist = levenshtein(&lower_name, &e.name_lower);
|
.map(|e| {
|
||||||
if dist > max_distance {
|
let dist = levenshtein(&lower_name, &e.name.to_lowercase());
|
||||||
continue;
|
(e.id, dist)
|
||||||
}
|
})
|
||||||
if dist == 0 {
|
.filter(|&(_, dist)| dist <= max_distance)
|
||||||
best = Some((e.id, dist));
|
.min_by_key(|&(_, dist)| dist);
|
||||||
break;
|
|
||||||
}
|
|
||||||
if best.is_none_or(|(_, best_dist)| dist < best_dist) {
|
|
||||||
best = Some((e.id, dist));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if let Some((id, _)) = best {
|
if let Some((id, _)) = best {
|
||||||
return (id, false);
|
return (id, false);
|
||||||
@@ -397,7 +337,6 @@ impl KnowledgeCache {
|
|||||||
/// together with their discovered depth. The seed entity itself is NOT
|
/// together with their discovered depth. The seed entity itself is NOT
|
||||||
/// included. Traversal follows both outgoing and incoming relation edges.
|
/// included. Traversal follows both outgoing and incoming relation edges.
|
||||||
pub fn bfs_neighbors(&self, entity_id: u64, max_depth: usize) -> Vec<(Entity, usize)> {
|
pub fn bfs_neighbors(&self, entity_id: u64, max_depth: usize) -> Vec<(Entity, usize)> {
|
||||||
let idx = AdjacencyIndex::build(&self.entities, &self.relations);
|
|
||||||
let mut visited: HashSet<u64> = HashSet::new();
|
let mut visited: HashSet<u64> = HashSet::new();
|
||||||
let mut queue: VecDeque<(u64, usize)> = VecDeque::new();
|
let mut queue: VecDeque<(u64, usize)> = VecDeque::new();
|
||||||
let mut results: Vec<(Entity, usize)> = Vec::new();
|
let mut results: Vec<(Entity, usize)> = Vec::new();
|
||||||
@@ -410,13 +349,11 @@ impl KnowledgeCache {
|
|||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Collect neighbour IDs from outgoing and incoming edges touching
|
// Collect neighbour IDs from outgoing and incoming edges.
|
||||||
// this node only, instead of scanning every relation in the graph.
|
let neighbours: Vec<u64> = self
|
||||||
let neighbours: Vec<u64> = idx
|
.relations
|
||||||
.relations_touching(current_id)
|
|
||||||
.iter()
|
.iter()
|
||||||
.filter_map(|&i| {
|
.filter_map(|r| {
|
||||||
let r = &self.relations[i];
|
|
||||||
if r.src == current_id {
|
if r.src == current_id {
|
||||||
Some(r.tgt)
|
Some(r.tgt)
|
||||||
} else if r.tgt == current_id {
|
} else if r.tgt == current_id {
|
||||||
@@ -429,9 +366,9 @@ impl KnowledgeCache {
|
|||||||
|
|
||||||
for neighbour_id in neighbours {
|
for neighbour_id in neighbours {
|
||||||
if visited.insert(neighbour_id)
|
if visited.insert(neighbour_id)
|
||||||
&& let Some(&entity_idx) = idx.entity_index.get(&neighbour_id)
|
&& let Some(entity) = self.get_entity(neighbour_id)
|
||||||
{
|
{
|
||||||
results.push((self.entities[entity_idx].clone(), depth + 1));
|
results.push((entity.clone(), depth + 1));
|
||||||
queue.push_back((neighbour_id, depth + 1));
|
queue.push_back((neighbour_id, depth + 1));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -502,7 +439,6 @@ impl KnowledgeCache {
|
|||||||
min_activation: f32,
|
min_activation: f32,
|
||||||
max_steps: usize,
|
max_steps: usize,
|
||||||
) -> Vec<(u64, f32)> {
|
) -> Vec<(u64, f32)> {
|
||||||
let idx = AdjacencyIndex::build(&self.entities, &self.relations);
|
|
||||||
let mut activation: HashMap<u64, f32> = HashMap::new();
|
let mut activation: HashMap<u64, f32> = HashMap::new();
|
||||||
|
|
||||||
// Initialise seeds with activation 1.0.
|
// Initialise seeds with activation 1.0.
|
||||||
@@ -525,10 +461,8 @@ impl KnowledgeCache {
|
|||||||
let mut any_spread = false;
|
let mut any_spread = false;
|
||||||
|
|
||||||
for (source_id, source_score) in current {
|
for (source_id, source_score) in current {
|
||||||
// Spread only to edges touching this node, instead of
|
// Spread to all neighbours via outgoing and incoming edges.
|
||||||
// scanning every relation in the graph per active node.
|
for rel in &self.relations {
|
||||||
for &rel_idx in idx.relations_touching(source_id) {
|
|
||||||
let rel = &self.relations[rel_idx];
|
|
||||||
let neighbour_id = if rel.src == source_id {
|
let neighbour_id = if rel.src == source_id {
|
||||||
rel.tgt
|
rel.tgt
|
||||||
} else if rel.tgt == source_id {
|
} else if rel.tgt == source_id {
|
||||||
@@ -921,19 +855,6 @@ mod tests {
|
|||||||
assert_eq!(id, orig_id);
|
assert_eq!(id, orig_id);
|
||||||
}
|
}
|
||||||
|
|
||||||
/// An exact match must win even when a near-match with a smaller Levenshtein
|
|
||||||
/// distance-to-zero gap was scanned first — the early exit on dist == 0
|
|
||||||
/// must not skip past a later exact match.
|
|
||||||
#[test]
|
|
||||||
fn test_resolve_or_create_exact_match_beats_earlier_fuzzy_candidate() {
|
|
||||||
let mut cache = KnowledgeCache::new();
|
|
||||||
cache.add_entity("Alyce", "person", -1); // dist 1 from "Alice"
|
|
||||||
let exact_id = cache.add_entity("Alice", "person", -1); // dist 0
|
|
||||||
let (id, created) = cache.resolve_or_create("Alice", "person", -1, 2);
|
|
||||||
assert!(!created);
|
|
||||||
assert_eq!(id, exact_id);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_resolve_or_create_no_match_beyond_threshold() {
|
fn test_resolve_or_create_no_match_beyond_threshold() {
|
||||||
let mut cache = KnowledgeCache::new();
|
let mut cache = KnowledgeCache::new();
|
||||||
@@ -1114,30 +1035,6 @@ mod tests {
|
|||||||
assert!(b_score.unwrap() > 0.0);
|
assert!(b_score.unwrap() > 0.0);
|
||||||
}
|
}
|
||||||
|
|
||||||
/// A self-loop relation (src == tgt) must be visited exactly once by the
|
|
||||||
/// adjacency index, matching the pre-index behavior of iterating
|
|
||||||
/// `self.relations` directly (each relation processed once regardless of
|
|
||||||
/// how many of its endpoints match the current node).
|
|
||||||
#[test]
|
|
||||||
fn test_spreading_activation_self_loop_not_double_counted() {
|
|
||||||
let mut cache = KnowledgeCache::new();
|
|
||||||
let a = cache.add_entity("A", "node", -1);
|
|
||||||
cache.add_relation(a, a, "self", 1.0);
|
|
||||||
|
|
||||||
let result = cache.spreading_activation(&[a], 0.5, 0.0001, 1);
|
|
||||||
let a_score = result
|
|
||||||
.iter()
|
|
||||||
.find(|&&(id, _)| id == a)
|
|
||||||
.map(|&(_, s)| s)
|
|
||||||
.unwrap();
|
|
||||||
// Seed activation (1.0) plus exactly one spread contribution
|
|
||||||
// (1.0 * weight 1.0 * decay 0.5), not two.
|
|
||||||
assert!(
|
|
||||||
(a_score - 1.5).abs() < 1e-5,
|
|
||||||
"expected 1.5 (one self-loop contribution), got {a_score}"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_spreading_activation_decay_reduces_signal() {
|
fn test_spreading_activation_decay_reduces_signal() {
|
||||||
let mut cache = KnowledgeCache::new();
|
let mut cache = KnowledgeCache::new();
|
||||||
|
|||||||
+206
-725
File diff suppressed because it is too large
Load Diff
@@ -748,69 +748,6 @@ impl MemoryBackend for ClawhdfBackend {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// ─────────────────────────────────────────────────────────────────────────────
|
|
||||||
// Ephemeral tier methods on ClawhdfBackend
|
|
||||||
// ─────────────────────────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
impl ClawhdfBackend {
|
|
||||||
/// Enable the ephemeral (in-memory only) working memory tier.
|
|
||||||
pub fn enable_ephemeral(&mut self, config: crate::ephemeral::EphemeralConfig) {
|
|
||||||
self.memory.enable_ephemeral(config);
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Store a text value in ephemeral memory.
|
|
||||||
///
|
|
||||||
/// Returns an error string if the ephemeral tier has not been enabled.
|
|
||||||
pub fn ephemeral_set(
|
|
||||||
&mut self,
|
|
||||||
key: &str,
|
|
||||||
value: &str,
|
|
||||||
ttl_secs: Option<f64>,
|
|
||||||
) -> Result<(), String> {
|
|
||||||
match self.memory.ephemeral_mut() {
|
|
||||||
Some(s) => {
|
|
||||||
s.set_text(key, value, ttl_secs);
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
None => Err("ephemeral tier not enabled".to_string()),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Retrieve a text value from ephemeral memory.
|
|
||||||
///
|
|
||||||
/// Returns `None` if the tier is disabled, the key is absent, or the
|
|
||||||
/// entry has expired.
|
|
||||||
pub fn ephemeral_get(&mut self, key: &str) -> Option<String> {
|
|
||||||
self.memory
|
|
||||||
.ephemeral_mut()?
|
|
||||||
.get_text(key)
|
|
||||||
.map(|s| s.to_string())
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Delete a key from ephemeral memory.
|
|
||||||
///
|
|
||||||
/// Returns `true` if the key existed and was removed.
|
|
||||||
pub fn ephemeral_delete(&mut self, key: &str) -> bool {
|
|
||||||
self.memory.ephemeral_mut().is_some_and(|s| s.delete(key))
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Return a snapshot of ephemeral tier statistics, or `None` if the tier
|
|
||||||
/// is not enabled.
|
|
||||||
pub fn ephemeral_stats(&self) -> Option<crate::ephemeral::EphemeralStats> {
|
|
||||||
self.memory.ephemeral().map(|s| s.stats())
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Promote frequently-accessed ephemeral entries to persistent HDF5 storage.
|
|
||||||
///
|
|
||||||
/// Entries with `access_count >= min_access_count` are moved from the
|
|
||||||
/// ephemeral store into the persistent cache. Returns the count promoted.
|
|
||||||
pub fn promote_ephemeral(&mut self, min_access_count: u32) -> Result<usize, String> {
|
|
||||||
self.memory
|
|
||||||
.promote_ephemeral(min_access_count)
|
|
||||||
.map_err(|e| e.to_string())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ─────────────────────────────────────────────────────────────────────────────
|
// ─────────────────────────────────────────────────────────────────────────────
|
||||||
// Tests
|
// Tests
|
||||||
// ─────────────────────────────────────────────────────────────────────────────
|
// ─────────────────────────────────────────────────────────────────────────────
|
||||||
@@ -1396,3 +1333,66 @@ mod tests {
|
|||||||
assert!(out.starts_with("# Title"));
|
assert!(out.starts_with("# Title"));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ─────────────────────────────────────────────────────────────────────────────
|
||||||
|
// Ephemeral tier methods on ClawhdfBackend
|
||||||
|
// ─────────────────────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
impl ClawhdfBackend {
|
||||||
|
/// Enable the ephemeral (in-memory only) working memory tier.
|
||||||
|
pub fn enable_ephemeral(&mut self, config: crate::ephemeral::EphemeralConfig) {
|
||||||
|
self.memory.enable_ephemeral(config);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Store a text value in ephemeral memory.
|
||||||
|
///
|
||||||
|
/// Returns an error string if the ephemeral tier has not been enabled.
|
||||||
|
pub fn ephemeral_set(
|
||||||
|
&mut self,
|
||||||
|
key: &str,
|
||||||
|
value: &str,
|
||||||
|
ttl_secs: Option<f64>,
|
||||||
|
) -> Result<(), String> {
|
||||||
|
match self.memory.ephemeral_mut() {
|
||||||
|
Some(s) => {
|
||||||
|
s.set_text(key, value, ttl_secs);
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
None => Err("ephemeral tier not enabled".to_string()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Retrieve a text value from ephemeral memory.
|
||||||
|
///
|
||||||
|
/// Returns `None` if the tier is disabled, the key is absent, or the
|
||||||
|
/// entry has expired.
|
||||||
|
pub fn ephemeral_get(&mut self, key: &str) -> Option<String> {
|
||||||
|
self.memory
|
||||||
|
.ephemeral_mut()?
|
||||||
|
.get_text(key)
|
||||||
|
.map(|s| s.to_string())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Delete a key from ephemeral memory.
|
||||||
|
///
|
||||||
|
/// Returns `true` if the key existed and was removed.
|
||||||
|
pub fn ephemeral_delete(&mut self, key: &str) -> bool {
|
||||||
|
self.memory.ephemeral_mut().is_some_and(|s| s.delete(key))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Return a snapshot of ephemeral tier statistics, or `None` if the tier
|
||||||
|
/// is not enabled.
|
||||||
|
pub fn ephemeral_stats(&self) -> Option<crate::ephemeral::EphemeralStats> {
|
||||||
|
self.memory.ephemeral().map(|s| s.stats())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Promote frequently-accessed ephemeral entries to persistent HDF5 storage.
|
||||||
|
///
|
||||||
|
/// Entries with `access_count >= min_access_count` are moved from the
|
||||||
|
/// ephemeral store into the persistent cache. Returns the count promoted.
|
||||||
|
pub fn promote_ephemeral(&mut self, min_access_count: u32) -> Result<usize, String> {
|
||||||
|
self.memory
|
||||||
|
.promote_ephemeral(min_access_count)
|
||||||
|
.map_err(|e| e.to_string())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -105,23 +105,6 @@ impl ProvenanceStore {
|
|||||||
self.records.insert(provenance.record_id, provenance);
|
self.records.insert(provenance.record_id, provenance);
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Renumber records after the store was compacted. `index_map[old]` is
|
|
||||||
/// the record's new id, or `None` if it was removed. Without this, every
|
|
||||||
/// surviving record's hash ends up filed under some other record's id and
|
|
||||||
/// the next integrity check reports a bogus mismatch.
|
|
||||||
pub fn remap(&mut self, index_map: &[Option<usize>]) {
|
|
||||||
let old = std::mem::take(&mut self.records);
|
|
||||||
for (old_id, mut prov) in old {
|
|
||||||
let new_id = usize::try_from(old_id)
|
|
||||||
.ok()
|
|
||||||
.and_then(|i| index_map.get(i).copied().flatten());
|
|
||||||
if let Some(new_id) = new_id {
|
|
||||||
prov.record_id = new_id as u64;
|
|
||||||
self.records.insert(new_id as u64, prov);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Retrieve by record ID.
|
/// Retrieve by record ID.
|
||||||
pub fn get(&self, record_id: u64) -> Option<&MemoryProvenance> {
|
pub fn get(&self, record_id: u64) -> Option<&MemoryProvenance> {
|
||||||
self.records.get(&record_id)
|
self.records.get(&record_id)
|
||||||
|
|||||||
@@ -12,36 +12,16 @@ use crate::MemoryError;
|
|||||||
use crate::cache::MemoryCache;
|
use crate::cache::MemoryCache;
|
||||||
use crate::knowledge::KnowledgeCache;
|
use crate::knowledge::KnowledgeCache;
|
||||||
use crate::session::SessionCache;
|
use crate::session::SessionCache;
|
||||||
use crate::wal::WalMark;
|
|
||||||
|
|
||||||
pub const SCHEMA_VERSION: &str = "1.0";
|
pub const SCHEMA_VERSION: &str = "1.0";
|
||||||
pub const ZEROCLAW_VERSION: &str = "0.8.0";
|
pub const ZEROCLAW_VERSION: &str = "0.8.0";
|
||||||
|
|
||||||
/// `/meta` attributes holding the [`WalMark`] of the WAL prefix already folded
|
|
||||||
/// into this file. Absent on files written before the mark existed, and when
|
|
||||||
/// the checkpoint was taken with an empty WAL.
|
|
||||||
const WAL_APPLIED_LEN_ATTR: &str = "wal_applied_len";
|
|
||||||
const WAL_APPLIED_CRC_ATTR: &str = "wal_applied_crc";
|
|
||||||
|
|
||||||
/// Build a complete HDF5 file from the in-memory state.
|
/// Build a complete HDF5 file from the in-memory state.
|
||||||
pub fn build_hdf5_file(
|
pub fn build_hdf5_file(
|
||||||
config: &MemoryConfig,
|
config: &MemoryConfig,
|
||||||
cache: &MemoryCache,
|
cache: &MemoryCache,
|
||||||
sessions: &SessionCache,
|
sessions: &SessionCache,
|
||||||
knowledge: &KnowledgeCache,
|
knowledge: &KnowledgeCache,
|
||||||
) -> Result<Vec<u8>, MemoryError> {
|
|
||||||
build_hdf5_file_with_mark(config, cache, sessions, knowledge, None)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// [`build_hdf5_file`], recording which WAL prefix this state already
|
|
||||||
/// contains (see [`WalMark`]) so a crash before the WAL is truncated doesn't
|
|
||||||
/// replay those entries a second time.
|
|
||||||
pub fn build_hdf5_file_with_mark(
|
|
||||||
config: &MemoryConfig,
|
|
||||||
cache: &MemoryCache,
|
|
||||||
sessions: &SessionCache,
|
|
||||||
knowledge: &KnowledgeCache,
|
|
||||||
wal_applied: Option<WalMark>,
|
|
||||||
) -> Result<Vec<u8>, MemoryError> {
|
) -> Result<Vec<u8>, MemoryError> {
|
||||||
let mut builder = clawhdf5::FileBuilder::new();
|
let mut builder = clawhdf5::FileBuilder::new();
|
||||||
|
|
||||||
@@ -54,35 +34,10 @@ pub fn build_hdf5_file_with_mark(
|
|||||||
meta.set_attr("embedding_dim", AttrValue::I64(config.embedding_dim as i64));
|
meta.set_attr("embedding_dim", AttrValue::I64(config.embedding_dim as i64));
|
||||||
meta.set_attr("chunk_size", AttrValue::I64(config.chunk_size as i64));
|
meta.set_attr("chunk_size", AttrValue::I64(config.chunk_size as i64));
|
||||||
meta.set_attr("overlap", AttrValue::I64(config.overlap as i64));
|
meta.set_attr("overlap", AttrValue::I64(config.overlap as i64));
|
||||||
// Behavioural settings. These used to live only in memory, so reopening a
|
|
||||||
// store silently reset them to defaults — e.g. a compressed store was
|
|
||||||
// rewritten uncompressed by the first checkpoint after a reopen. Loaders
|
|
||||||
// treat each one as optional so older files keep opening.
|
|
||||||
meta.set_attr("float16", AttrValue::I64(config.float16.into()));
|
|
||||||
meta.set_attr("compression", AttrValue::I64(config.compression.into()));
|
|
||||||
meta.set_attr(
|
|
||||||
"compression_level",
|
|
||||||
AttrValue::I64(config.compression_level.into()),
|
|
||||||
);
|
|
||||||
meta.set_attr(
|
|
||||||
"compact_threshold",
|
|
||||||
AttrValue::F64(config.compact_threshold.into()),
|
|
||||||
);
|
|
||||||
meta.set_attr("hebbian_boost", AttrValue::F64(config.hebbian_boost.into()));
|
|
||||||
meta.set_attr("decay_factor", AttrValue::F64(config.decay_factor.into()));
|
|
||||||
meta.set_attr("wal_enabled", AttrValue::I64(config.wal_enabled.into()));
|
|
||||||
meta.set_attr(
|
|
||||||
"wal_max_entries",
|
|
||||||
AttrValue::I64(config.wal_max_entries as i64),
|
|
||||||
);
|
|
||||||
meta.set_attr(
|
meta.set_attr(
|
||||||
"edgehdf5_version",
|
"edgehdf5_version",
|
||||||
AttrValue::String(ZEROCLAW_VERSION.into()),
|
AttrValue::String(ZEROCLAW_VERSION.into()),
|
||||||
);
|
);
|
||||||
if let Some(mark) = wal_applied.filter(|m| m.len > 0) {
|
|
||||||
meta.set_attr(WAL_APPLIED_LEN_ATTR, AttrValue::I64(mark.len as i64));
|
|
||||||
meta.set_attr(WAL_APPLIED_CRC_ATTR, AttrValue::I64(i64::from(mark.crc)));
|
|
||||||
}
|
|
||||||
// Need at least one dataset in the group for it to be a proper group
|
// Need at least one dataset in the group for it to be a proper group
|
||||||
meta.create_dataset("_marker").with_u8_data(&[1]).compact();
|
meta.create_dataset("_marker").with_u8_data(&[1]).compact();
|
||||||
let finished_meta = meta.finish();
|
let finished_meta = meta.finish();
|
||||||
@@ -128,34 +83,16 @@ fn build_memory_group(
|
|||||||
let rows_per_chunk = (target_chunk_bytes / (d * 4)).max(1).min(n);
|
let rows_per_chunk = (target_chunk_bytes / (d * 4)).max(1).min(n);
|
||||||
ds.with_chunks(&[rows_per_chunk, d]);
|
ds.with_chunks(&[rows_per_chunk, d]);
|
||||||
|
|
||||||
// Compression. Shuffle is applied automatically (auto-shuffle
|
// Compression: Zstd for embeddings — faster than deflate at same ratio.
|
||||||
// pre-filter). Zstd is faster than deflate at the same ratio but
|
// Shuffle is applied automatically (auto-shuffle pre-filter).
|
||||||
// pulls in libzstd, so it is opt-in via the `zstd` feature; the
|
|
||||||
// default build uses deflate, which is always available. (This
|
|
||||||
// used to call `with_zstd` unconditionally, so without the
|
|
||||||
// feature every checkpoint of a compressed store failed with
|
|
||||||
// "unsupported filter: 32015".) Both are standard HDF5 filters;
|
|
||||||
// reading a zstd-compressed store needs a zstd-enabled build.
|
|
||||||
if config.compression {
|
if config.compression {
|
||||||
#[cfg(feature = "zstd")]
|
|
||||||
{
|
|
||||||
let level = if config.compression_level > 0 {
|
let level = if config.compression_level > 0 {
|
||||||
config.compression_level.min(22)
|
config.compression_level.min(22)
|
||||||
} else {
|
} else {
|
||||||
3 // fast + good ratio for f32 embeddings
|
3 // Zstd level 3: fast + good ratio for f32 embeddings
|
||||||
};
|
};
|
||||||
ds.with_zstd(level);
|
ds.with_zstd(level);
|
||||||
}
|
}
|
||||||
#[cfg(not(feature = "zstd"))]
|
|
||||||
{
|
|
||||||
let level = if config.compression_level > 0 {
|
|
||||||
config.compression_level.min(9)
|
|
||||||
} else {
|
|
||||||
4
|
|
||||||
};
|
|
||||||
ds.with_deflate(level);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Skip fill-value initialization — embeddings are fully written
|
// Skip fill-value initialization — embeddings are fully written
|
||||||
@@ -372,20 +309,6 @@ fn write_string_dataset(
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Validate an HDF5 file has the correct schema and load all data.
|
/// Validate an HDF5 file has the correct schema and load all data.
|
||||||
/// Read the checkpoint's [`WalMark`] from `/meta`, if it has one.
|
|
||||||
pub fn read_wal_mark(file: &clawhdf5::File) -> Option<WalMark> {
|
|
||||||
let attrs = file.group("meta").ok()?.attrs().ok()?;
|
|
||||||
let len = match attrs.get(WAL_APPLIED_LEN_ATTR)? {
|
|
||||||
AttrValue::I64(v) => u64::try_from(*v).ok()?,
|
|
||||||
_ => return None,
|
|
||||||
};
|
|
||||||
let crc = match attrs.get(WAL_APPLIED_CRC_ATTR)? {
|
|
||||||
AttrValue::I64(v) => u32::try_from(*v).ok()?,
|
|
||||||
_ => return None,
|
|
||||||
};
|
|
||||||
Some(WalMark { len, crc })
|
|
||||||
}
|
|
||||||
|
|
||||||
pub fn validate_and_load(
|
pub fn validate_and_load(
|
||||||
file: &clawhdf5::File,
|
file: &clawhdf5::File,
|
||||||
) -> Result<(MemoryConfig, MemoryCache, SessionCache, KnowledgeCache), MemoryError> {
|
) -> Result<(MemoryConfig, MemoryCache, SessionCache, KnowledgeCache), MemoryError> {
|
||||||
@@ -421,19 +344,15 @@ pub fn validate_and_load(
|
|||||||
embedding_dim,
|
embedding_dim,
|
||||||
chunk_size,
|
chunk_size,
|
||||||
overlap,
|
overlap,
|
||||||
float16: optional_bool_attr(&attrs, "float16", false),
|
float16: false,
|
||||||
compression: optional_bool_attr(&attrs, "compression", false),
|
compression: false,
|
||||||
compression_level: optional_i64_attr(&attrs, "compression_level")
|
compression_level: 0,
|
||||||
.and_then(|v| u32::try_from(v).ok())
|
compact_threshold: 0.3,
|
||||||
.unwrap_or(0),
|
hebbian_boost: 0.15,
|
||||||
compact_threshold: optional_f32_attr(&attrs, "compact_threshold", 0.3),
|
decay_factor: 0.98,
|
||||||
hebbian_boost: optional_f32_attr(&attrs, "hebbian_boost", 0.15),
|
|
||||||
decay_factor: optional_f32_attr(&attrs, "decay_factor", 0.98),
|
|
||||||
created_at,
|
created_at,
|
||||||
wal_enabled: optional_bool_attr(&attrs, "wal_enabled", true),
|
wal_enabled: true,
|
||||||
wal_max_entries: optional_i64_attr(&attrs, "wal_max_entries")
|
wal_max_entries: 500,
|
||||||
.and_then(|v| usize::try_from(v).ok())
|
|
||||||
.unwrap_or(500),
|
|
||||||
};
|
};
|
||||||
|
|
||||||
// Load /memory group
|
// Load /memory group
|
||||||
@@ -472,45 +391,19 @@ fn load_memory_group(
|
|||||||
let tags = read_string_dataset_from_group(&group, "tags")?;
|
let tags = read_string_dataset_from_group(&group, "tags")?;
|
||||||
let tombstones = read_u8_dataset(&group, "tombstones")?;
|
let tombstones = read_u8_dataset(&group, "tombstones")?;
|
||||||
|
|
||||||
// Every per-record dataset must describe exactly `n` records. Without
|
// Read norms if present, otherwise compute from embeddings
|
||||||
// this, a truncated or hand-edited file loads "successfully" and then
|
|
||||||
// panics on the first out-of-bounds index during search/delete.
|
|
||||||
if embedding_dim == 0 {
|
|
||||||
return Err(MemoryError::Schema(format!(
|
|
||||||
"/memory has {n} records but embedding_dim is 0"
|
|
||||||
)));
|
|
||||||
}
|
|
||||||
let expected_flat = n.checked_mul(embedding_dim).ok_or_else(|| {
|
|
||||||
MemoryError::Schema(format!("/memory size overflow: {n} x {embedding_dim}"))
|
|
||||||
})?;
|
|
||||||
let check_len = |name: &str, actual: usize, expected: usize| {
|
|
||||||
if actual == expected {
|
|
||||||
Ok(())
|
|
||||||
} else {
|
|
||||||
Err(MemoryError::Schema(format!(
|
|
||||||
"/memory/{name} has {actual} entries, expected {expected} \
|
|
||||||
({n} records)"
|
|
||||||
)))
|
|
||||||
}
|
|
||||||
};
|
|
||||||
check_len("embeddings", flat_embeddings.len(), expected_flat)?;
|
|
||||||
check_len("source_channel", source_channels.len(), n)?;
|
|
||||||
check_len("timestamps", timestamps.len(), n)?;
|
|
||||||
check_len("session_ids", session_ids.len(), n)?;
|
|
||||||
check_len("tags", tags.len(), n)?;
|
|
||||||
check_len("tombstones", tombstones.len(), n)?;
|
|
||||||
|
|
||||||
// Norms are derived data: use the stored ones only if they are present
|
|
||||||
// and the right length, otherwise recompute from the embeddings.
|
|
||||||
let norms = match read_f32_dataset(&group, "norms") {
|
let norms = match read_f32_dataset(&group, "norms") {
|
||||||
Ok(stored) if stored.len() == n => stored,
|
Ok(n) if n.len() == n.len() => n,
|
||||||
_ => flat_embeddings
|
_ => {
|
||||||
|
// Compute norms from flat embeddings
|
||||||
|
flat_embeddings
|
||||||
.chunks(embedding_dim)
|
.chunks(embedding_dim)
|
||||||
.map(|chunk| {
|
.map(|chunk| {
|
||||||
let sq_sum: f32 = chunk.iter().map(|x| x * x).sum();
|
let sq_sum: f32 = chunk.iter().map(|x| x * x).sum();
|
||||||
sq_sum.sqrt()
|
sq_sum.sqrt()
|
||||||
})
|
})
|
||||||
.collect(),
|
.collect()
|
||||||
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
// Unflatten embeddings
|
// Unflatten embeddings
|
||||||
@@ -534,7 +427,6 @@ fn load_memory_group(
|
|||||||
cache.tombstones = tombstones;
|
cache.tombstones = tombstones;
|
||||||
cache.norms = norms;
|
cache.norms = norms;
|
||||||
cache.activation_weights = activation_weights;
|
cache.activation_weights = activation_weights;
|
||||||
cache.rebuild_flat();
|
|
||||||
|
|
||||||
Ok(cache)
|
Ok(cache)
|
||||||
}
|
}
|
||||||
@@ -588,7 +480,6 @@ fn load_knowledge_group(file: &clawhdf5::File) -> Result<KnowledgeCache, MemoryE
|
|||||||
cache.entities.push(crate::knowledge::Entity {
|
cache.entities.push(crate::knowledge::Entity {
|
||||||
id: entity_ids[i] as u64,
|
id: entity_ids[i] as u64,
|
||||||
name: entity_names[i].clone(),
|
name: entity_names[i].clone(),
|
||||||
name_lower: entity_names[i].to_lowercase(),
|
|
||||||
entity_type: entity_types[i].clone(),
|
entity_type: entity_types[i].clone(),
|
||||||
embedding_idx: emb_idxs[i],
|
embedding_idx: emb_idxs[i],
|
||||||
..Default::default()
|
..Default::default()
|
||||||
@@ -638,27 +529,6 @@ fn extract_string_attr(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
type MetaAttrs = std::collections::HashMap<String, AttrValue>;
|
|
||||||
|
|
||||||
fn optional_i64_attr(attrs: &MetaAttrs, name: &str) -> Option<i64> {
|
|
||||||
match attrs.get(name) {
|
|
||||||
Some(AttrValue::I64(v)) => Some(*v),
|
|
||||||
_ => None,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn optional_bool_attr(attrs: &MetaAttrs, name: &str, default: bool) -> bool {
|
|
||||||
optional_i64_attr(attrs, name).map_or(default, |v| v != 0)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Finite values only: a NaN threshold/decay would poison every comparison.
|
|
||||||
fn optional_f32_attr(attrs: &MetaAttrs, name: &str, default: f32) -> f32 {
|
|
||||||
match attrs.get(name) {
|
|
||||||
Some(AttrValue::F64(v)) if v.is_finite() => *v as f32,
|
|
||||||
_ => default,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn extract_i64_attr(
|
fn extract_i64_attr(
|
||||||
attrs: &std::collections::HashMap<String, AttrValue>,
|
attrs: &std::collections::HashMap<String, AttrValue>,
|
||||||
name: &str,
|
name: &str,
|
||||||
@@ -744,108 +614,3 @@ fn read_u8_dataset(group: &clawhdf5::Group<'_>, name: &str) -> Result<Vec<u8>, M
|
|||||||
.map_err(|e| MemoryError::Hdf5(format!("cannot read u8 from {name}: {e}")))?;
|
.map_err(|e| MemoryError::Hdf5(format!("cannot read u8 from {name}: {e}")))?;
|
||||||
Ok(data.into_iter().map(|v| v as u8).collect())
|
Ok(data.into_iter().map(|v| v as u8).collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
fn config() -> MemoryConfig {
|
|
||||||
MemoryConfig::new(std::path::PathBuf::from("unused.h5"), "agent", 4)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn cache_with(n: usize) -> MemoryCache {
|
|
||||||
let mut cache = MemoryCache::new(4);
|
|
||||||
for i in 0..n {
|
|
||||||
cache.push(
|
|
||||||
format!("chunk {i}"),
|
|
||||||
vec![i as f32 + 1.0, 0.0, 0.0, 0.0],
|
|
||||||
"user".into(),
|
|
||||||
i as f64,
|
|
||||||
"s".into(),
|
|
||||||
"t".into(),
|
|
||||||
);
|
|
||||||
}
|
|
||||||
cache
|
|
||||||
}
|
|
||||||
|
|
||||||
fn roundtrip(cache: &MemoryCache) -> Result<MemoryCache, MemoryError> {
|
|
||||||
let bytes = build_hdf5_file(
|
|
||||||
&config(),
|
|
||||||
cache,
|
|
||||||
&SessionCache::new(),
|
|
||||||
&KnowledgeCache::new(),
|
|
||||||
)?;
|
|
||||||
let file =
|
|
||||||
clawhdf5::File::from_bytes(bytes).map_err(|e| MemoryError::Hdf5(e.to_string()))?;
|
|
||||||
validate_and_load(&file).map(|(_, cache, _, _)| cache)
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn behavioural_config_survives_a_reopen() {
|
|
||||||
let mut cfg = config();
|
|
||||||
cfg.compression = true;
|
|
||||||
cfg.compression_level = 7;
|
|
||||||
cfg.compact_threshold = 0.5;
|
|
||||||
cfg.hebbian_boost = 0.25;
|
|
||||||
cfg.decay_factor = 0.9;
|
|
||||||
cfg.wal_enabled = false;
|
|
||||||
cfg.wal_max_entries = 42;
|
|
||||||
let bytes = build_hdf5_file(
|
|
||||||
&cfg,
|
|
||||||
&cache_with(2),
|
|
||||||
&SessionCache::new(),
|
|
||||||
&KnowledgeCache::new(),
|
|
||||||
)
|
|
||||||
.unwrap();
|
|
||||||
let file = clawhdf5::File::from_bytes(bytes).unwrap();
|
|
||||||
let (loaded, loaded_cache, ..) = validate_and_load(&file).unwrap();
|
|
||||||
// The compressed embeddings must also read back intact.
|
|
||||||
assert_eq!(loaded_cache.embeddings, cache_with(2).embeddings);
|
|
||||||
assert!(loaded.compression);
|
|
||||||
assert_eq!(loaded.compression_level, 7);
|
|
||||||
assert_eq!(loaded.compact_threshold, 0.5);
|
|
||||||
assert_eq!(loaded.hebbian_boost, 0.25);
|
|
||||||
assert_eq!(loaded.decay_factor, 0.9);
|
|
||||||
assert!(!loaded.wal_enabled);
|
|
||||||
assert_eq!(loaded.wal_max_entries, 42);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn consistent_store_loads() {
|
|
||||||
let loaded = roundtrip(&cache_with(3)).unwrap();
|
|
||||||
assert_eq!(loaded.chunks.len(), 3);
|
|
||||||
assert_eq!(loaded.norms, vec![1.0, 2.0, 3.0]);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn wrong_length_norms_are_recomputed_not_trusted() {
|
|
||||||
// Regression: the guard used to be `n.len() == n.len()`, so a norms
|
|
||||||
// dataset of any length was accepted and corrupted every cosine score.
|
|
||||||
let mut cache = cache_with(3);
|
|
||||||
cache.norms = vec![99.0];
|
|
||||||
let loaded = roundtrip(&cache).unwrap();
|
|
||||||
assert_eq!(loaded.norms, vec![1.0, 2.0, 3.0]);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn mismatched_per_record_datasets_are_schema_errors() {
|
|
||||||
type Corrupt = fn(&mut MemoryCache);
|
|
||||||
let cases: [(&str, Corrupt); 5] = [
|
|
||||||
("tombstones", |c| c.tombstones.truncate(1)),
|
|
||||||
("timestamps", |c| c.timestamps.truncate(1)),
|
|
||||||
("tags", |c| c.tags.truncate(1)),
|
|
||||||
("session_ids", |c| c.session_ids.truncate(1)),
|
|
||||||
("source_channel", |c| c.source_channels.truncate(1)),
|
|
||||||
];
|
|
||||||
for (name, corrupt) in cases {
|
|
||||||
let mut cache = cache_with(3);
|
|
||||||
corrupt(&mut cache);
|
|
||||||
match roundtrip(&cache) {
|
|
||||||
Err(MemoryError::Schema(msg)) => {
|
|
||||||
assert!(msg.contains(name), "{name}: unexpected message {msg}")
|
|
||||||
}
|
|
||||||
other => panic!("{name}: expected Schema error, got {:?}", other.map(|_| ())),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -113,24 +113,13 @@ impl HDF5Memory {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
.collect();
|
.collect();
|
||||||
// Ties broken by index so results (and therefore which records get
|
|
||||||
// boosted) don't depend on HashMap iteration order upstream.
|
|
||||||
results.sort_by(|a, b| {
|
results.sort_by(|a, b| {
|
||||||
b.score
|
b.score
|
||||||
.partial_cmp(&a.score)
|
.partial_cmp(&a.score)
|
||||||
.unwrap_or(std::cmp::Ordering::Equal)
|
.unwrap_or(std::cmp::Ordering::Equal)
|
||||||
.then(a.index.cmp(&b.index))
|
|
||||||
});
|
});
|
||||||
|
|
||||||
// Only reinforce records that actually matched. When fewer than `k`
|
let hit_indices: Vec<usize> = results.iter().map(|r| r.index).collect();
|
||||||
// records are relevant, the rest of the list is zero-score filler;
|
|
||||||
// boosting it would teach the store that arbitrary records are
|
|
||||||
// important just because they were nearby in iteration order.
|
|
||||||
let hit_indices: Vec<usize> = results
|
|
||||||
.iter()
|
|
||||||
.filter(|r| r.score > 0.0)
|
|
||||||
.map(|r| r.index)
|
|
||||||
.collect();
|
|
||||||
self.apply_hebbian_boost(&hit_indices);
|
self.apply_hebbian_boost(&hit_indices);
|
||||||
self.flush().ok();
|
self.flush().ok();
|
||||||
|
|
||||||
|
|||||||
@@ -11,7 +11,6 @@ use crate::cache::MemoryCache;
|
|||||||
use crate::knowledge::KnowledgeCache;
|
use crate::knowledge::KnowledgeCache;
|
||||||
use crate::schema;
|
use crate::schema;
|
||||||
use crate::session::SessionCache;
|
use crate::session::SessionCache;
|
||||||
use crate::wal::WalMark;
|
|
||||||
|
|
||||||
/// Write all in-memory state to an HDF5 file on disk.
|
/// Write all in-memory state to an HDF5 file on disk.
|
||||||
pub fn write_to_disk(
|
pub fn write_to_disk(
|
||||||
@@ -21,20 +20,7 @@ pub fn write_to_disk(
|
|||||||
sessions: &SessionCache,
|
sessions: &SessionCache,
|
||||||
knowledge: &KnowledgeCache,
|
knowledge: &KnowledgeCache,
|
||||||
) -> Result<(), MemoryError> {
|
) -> Result<(), MemoryError> {
|
||||||
write_to_disk_with_mark(path, config, cache, sessions, knowledge, None)
|
let bytes = schema::build_hdf5_file(config, cache, sessions, knowledge)?;
|
||||||
}
|
|
||||||
|
|
||||||
/// [`write_to_disk`] for a checkpoint: `wal_applied` is the mark of the WAL
|
|
||||||
/// prefix whose entries `cache` already contains.
|
|
||||||
pub fn write_to_disk_with_mark(
|
|
||||||
path: &Path,
|
|
||||||
config: &MemoryConfig,
|
|
||||||
cache: &MemoryCache,
|
|
||||||
sessions: &SessionCache,
|
|
||||||
knowledge: &KnowledgeCache,
|
|
||||||
wal_applied: Option<WalMark>,
|
|
||||||
) -> Result<(), MemoryError> {
|
|
||||||
let bytes = schema::build_hdf5_file_with_mark(config, cache, sessions, knowledge, wal_applied)?;
|
|
||||||
|
|
||||||
if bytes.is_empty() {
|
if bytes.is_empty() {
|
||||||
return Err(MemoryError::Hdf5("build_hdf5_file produced 0 bytes".into()));
|
return Err(MemoryError::Hdf5("build_hdf5_file produced 0 bytes".into()));
|
||||||
@@ -42,41 +28,9 @@ pub fn write_to_disk_with_mark(
|
|||||||
|
|
||||||
// Write to a temp file first, then rename for atomicity
|
// Write to a temp file first, then rename for atomicity
|
||||||
let tmp_path = path.with_extension("h5.tmp");
|
let tmp_path = path.with_extension("h5.tmp");
|
||||||
write_synced(&tmp_path, &bytes)?;
|
std::fs::write(&tmp_path, &bytes).map_err(MemoryError::Io)?;
|
||||||
rename_synced(&tmp_path, path)
|
std::fs::rename(&tmp_path, path).map_err(MemoryError::Io)?;
|
||||||
}
|
|
||||||
|
|
||||||
/// Write `bytes` to `path` and flush them to stable storage.
|
|
||||||
fn write_synced(path: &Path, bytes: &[u8]) -> Result<(), MemoryError> {
|
|
||||||
use std::io::Write;
|
|
||||||
let mut f = std::fs::File::create(path).map_err(MemoryError::Io)?;
|
|
||||||
f.write_all(bytes).map_err(MemoryError::Io)?;
|
|
||||||
f.sync_all().map_err(MemoryError::Io)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Rename `from` over `to`, then sync the parent directory so the rename
|
|
||||||
/// itself survives a power loss. `from` must already be synced: without that,
|
|
||||||
/// the rename can reach disk before the data and leave an empty or partial
|
|
||||||
/// file under the final name.
|
|
||||||
///
|
|
||||||
/// This is per-checkpoint/snapshot cost only (each is already a full file
|
|
||||||
/// write). Individual WAL appends are deliberately not synced — see the
|
|
||||||
/// durability notes in the crate docs.
|
|
||||||
fn rename_synced(from: &Path, to: &Path) -> Result<(), MemoryError> {
|
|
||||||
std::fs::rename(from, to).map_err(MemoryError::Io)?;
|
|
||||||
#[cfg(unix)]
|
|
||||||
if let Some(dir) = to.parent() {
|
|
||||||
let dir = if dir.as_os_str().is_empty() {
|
|
||||||
Path::new(".")
|
|
||||||
} else {
|
|
||||||
dir
|
|
||||||
};
|
|
||||||
// Directory fsync is best-effort: some filesystems refuse it, and the
|
|
||||||
// rename has already happened.
|
|
||||||
if let Ok(d) = std::fs::File::open(dir) {
|
|
||||||
let _ = d.sync_all();
|
|
||||||
}
|
|
||||||
}
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -88,15 +42,6 @@ fn rename_synced(from: &Path, to: &Path) -> Result<(), MemoryError> {
|
|||||||
pub fn read_from_disk(
|
pub fn read_from_disk(
|
||||||
path: &Path,
|
path: &Path,
|
||||||
) -> Result<(MemoryConfig, MemoryCache, SessionCache, KnowledgeCache), MemoryError> {
|
) -> Result<(MemoryConfig, MemoryCache, SessionCache, KnowledgeCache), MemoryError> {
|
||||||
read_from_disk_with_mark(path).map(|(state, _mark)| state)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Everything [`read_from_disk`] returns.
|
|
||||||
pub type StoreState = (MemoryConfig, MemoryCache, SessionCache, KnowledgeCache);
|
|
||||||
|
|
||||||
/// [`read_from_disk`], plus the checkpoint's [`WalMark`] (if any) so the
|
|
||||||
/// caller can skip WAL entries this file already contains.
|
|
||||||
pub fn read_from_disk_with_mark(path: &Path) -> Result<(StoreState, Option<WalMark>), MemoryError> {
|
|
||||||
let mmap = clawhdf5_io::MmapReader::open(path).map_err(MemoryError::Io)?;
|
let mmap = clawhdf5_io::MmapReader::open(path).map_err(MemoryError::Io)?;
|
||||||
|
|
||||||
// Advise the OS we'll need the whole file for parsing
|
// Advise the OS we'll need the whole file for parsing
|
||||||
@@ -108,9 +53,8 @@ pub fn read_from_disk_with_mark(path: &Path) -> Result<(StoreState, Option<WalMa
|
|||||||
|
|
||||||
let (mut config, cache, sessions, knowledge) = schema::validate_and_load(&file)?;
|
let (mut config, cache, sessions, knowledge) = schema::validate_and_load(&file)?;
|
||||||
config.path = path.to_path_buf();
|
config.path = path.to_path_buf();
|
||||||
let wal_applied = schema::read_wal_mark(&file);
|
|
||||||
|
|
||||||
Ok(((config, cache, sessions, knowledge), wal_applied))
|
Ok((config, cache, sessions, knowledge))
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Copy an HDF5 file atomically to a destination.
|
/// Copy an HDF5 file atomically to a destination.
|
||||||
@@ -134,10 +78,7 @@ pub fn snapshot_file(src: &Path, dest: &Path) -> Result<std::path::PathBuf, Memo
|
|||||||
// Atomic copy: write to temp, then rename
|
// Atomic copy: write to temp, then rename
|
||||||
let tmp_path = dest_file.with_extension("h5.tmp");
|
let tmp_path = dest_file.with_extension("h5.tmp");
|
||||||
std::fs::copy(src, &tmp_path).map_err(MemoryError::Io)?;
|
std::fs::copy(src, &tmp_path).map_err(MemoryError::Io)?;
|
||||||
std::fs::File::open(&tmp_path)
|
std::fs::rename(&tmp_path, &dest_file).map_err(MemoryError::Io)?;
|
||||||
.and_then(|f| f.sync_all())
|
|
||||||
.map_err(MemoryError::Io)?;
|
|
||||||
rename_synced(&tmp_path, &dest_file)?;
|
|
||||||
|
|
||||||
Ok(dest_file)
|
Ok(dest_file)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,79 +0,0 @@
|
|||||||
//! Single-writer guard for a memory store.
|
|
||||||
//!
|
|
||||||
//! `HDF5Memory` keeps the whole store in memory and rewrites the `.h5` file at
|
|
||||||
//! every checkpoint, so two handles on one store (two processes, or two opens
|
|
||||||
//! in one process) silently destroy each other's data: whoever checkpoints
|
|
||||||
//! last wins, and both append to the same WAL with independent CRC chains.
|
|
||||||
//! The lock turns that into an immediate, explicit error.
|
|
||||||
|
|
||||||
use std::fs::{File, OpenOptions, TryLockError};
|
|
||||||
use std::path::{Path, PathBuf};
|
|
||||||
|
|
||||||
use crate::MemoryError;
|
|
||||||
|
|
||||||
const LOCK_RETRIES: u32 = 25;
|
|
||||||
const LOCK_RETRY_DELAY: std::time::Duration = std::time::Duration::from_millis(10);
|
|
||||||
|
|
||||||
/// An exclusive advisory lock on `<store>.h5.lock`, held for the lifetime of
|
|
||||||
/// the owning `HDF5Memory` and released when it is dropped (or when the
|
|
||||||
/// process dies — the OS drops the lock with the file descriptor, so a crash
|
|
||||||
/// never leaves a stale lock behind; the empty lock file itself is harmless).
|
|
||||||
#[derive(Debug)]
|
|
||||||
pub(crate) struct StoreLock {
|
|
||||||
_file: File,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl StoreLock {
|
|
||||||
pub(crate) fn lock_path(store: &Path) -> PathBuf {
|
|
||||||
store.with_extension("h5.lock")
|
|
||||||
}
|
|
||||||
|
|
||||||
pub(crate) fn acquire(store: &Path) -> Result<Self, MemoryError> {
|
|
||||||
let path = Self::lock_path(store);
|
|
||||||
let file = OpenOptions::new()
|
|
||||||
.create(true)
|
|
||||||
.truncate(false)
|
|
||||||
.write(true)
|
|
||||||
.open(&path)?;
|
|
||||||
// A previous owner may be mid-teardown (e.g. an `AsyncHDF5Memory`
|
|
||||||
// dropped without `shutdown()`: its background task releases the
|
|
||||||
// store a moment later), so give the lock a short, bounded grace
|
|
||||||
// period before reporting a genuine second writer.
|
|
||||||
let mut attempts_left = LOCK_RETRIES;
|
|
||||||
loop {
|
|
||||||
match file.try_lock() {
|
|
||||||
Ok(()) => return Ok(Self { _file: file }),
|
|
||||||
Err(TryLockError::WouldBlock) if attempts_left > 0 => {
|
|
||||||
attempts_left -= 1;
|
|
||||||
std::thread::sleep(LOCK_RETRY_DELAY);
|
|
||||||
}
|
|
||||||
Err(TryLockError::WouldBlock) => {
|
|
||||||
return Err(MemoryError::Locked(format!(
|
|
||||||
"{} is already open in this or another process (lock file {})",
|
|
||||||
store.display(),
|
|
||||||
path.display()
|
|
||||||
)));
|
|
||||||
}
|
|
||||||
Err(TryLockError::Error(e)) => return Err(MemoryError::Io(e)),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn second_acquire_fails_until_first_is_dropped() {
|
|
||||||
let dir = tempfile::TempDir::new().unwrap();
|
|
||||||
let store = dir.path().join("s.h5");
|
|
||||||
let first = StoreLock::acquire(&store).unwrap();
|
|
||||||
assert!(matches!(
|
|
||||||
StoreLock::acquire(&store),
|
|
||||||
Err(MemoryError::Locked(_))
|
|
||||||
));
|
|
||||||
drop(first);
|
|
||||||
StoreLock::acquire(&store).unwrap();
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -167,17 +167,10 @@ pub fn auto_select_strategy(num_vectors: usize, hw: &HardwareCapabilities) -> Se
|
|||||||
/// This dispatches to the appropriate search implementation based on the
|
/// This dispatches to the appropriate search implementation based on the
|
||||||
/// selected strategy. For IVF-PQ, an index must be provided externally
|
/// selected strategy. For IVF-PQ, an index must be provided externally
|
||||||
/// (this function uses brute-force fallback if no IVF-PQ index is available).
|
/// (this function uses brute-force fallback if no IVF-PQ index is available).
|
||||||
///
|
|
||||||
/// `vectors_flat` is `vectors` flattened into one contiguous `[N × dim]`
|
|
||||||
/// row-major buffer (e.g. `MemoryCache::embeddings_flat`, maintained
|
|
||||||
/// incrementally alongside `vectors`). It's only consulted by the
|
|
||||||
/// `Blas`/`Accelerate` strategies, which otherwise re-flatten the whole
|
|
||||||
/// corpus on every call — passing the already-flat buffer skips that copy.
|
|
||||||
#[allow(clippy::too_many_arguments)]
|
#[allow(clippy::too_many_arguments)]
|
||||||
pub fn search_with_metrics(
|
pub fn search_with_metrics(
|
||||||
query: &[f32],
|
query: &[f32],
|
||||||
vectors: &[Vec<f32>],
|
vectors: &[Vec<f32>],
|
||||||
vectors_flat: &[f32],
|
|
||||||
norms: &[f32],
|
norms: &[f32],
|
||||||
tombstones: &[u8],
|
tombstones: &[u8],
|
||||||
k: usize,
|
k: usize,
|
||||||
@@ -185,10 +178,6 @@ pub fn search_with_metrics(
|
|||||||
#[cfg(feature = "gpu")] gpu_backend: Option<&crate::gpu_search::GpuSearchBackend>,
|
#[cfg(feature = "gpu")] gpu_backend: Option<&crate::gpu_search::GpuSearchBackend>,
|
||||||
#[cfg(not(feature = "gpu"))] _gpu_backend: Option<&()>,
|
#[cfg(not(feature = "gpu"))] _gpu_backend: Option<&()>,
|
||||||
) -> (Vec<(usize, f32)>, SearchMetrics) {
|
) -> (Vec<(usize, f32)>, SearchMetrics) {
|
||||||
// Only read by the Blas/Accelerate arms below, which are themselves
|
|
||||||
// feature-gated — reference it unconditionally so a build with neither
|
|
||||||
// feature enabled doesn't warn about an unused parameter.
|
|
||||||
let _ = vectors_flat;
|
|
||||||
let start = Instant::now();
|
let start = Instant::now();
|
||||||
let active_count = tombstones.iter().filter(|&&t| t == 0).count();
|
let active_count = tombstones.iter().filter(|&&t| t == 0).count();
|
||||||
|
|
||||||
@@ -208,14 +197,7 @@ pub fn search_with_metrics(
|
|||||||
gpu_active = false;
|
gpu_active = false;
|
||||||
#[cfg(feature = "fast-math")]
|
#[cfg(feature = "fast-math")]
|
||||||
{
|
{
|
||||||
crate::blas_search::blas_cosine_batch_flat(
|
crate::blas_search::blas_cosine_batch(query, vectors, norms, tombstones, k)
|
||||||
query,
|
|
||||||
vectors_flat,
|
|
||||||
norms,
|
|
||||||
tombstones,
|
|
||||||
query.len(),
|
|
||||||
k,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
#[cfg(not(feature = "fast-math"))]
|
#[cfg(not(feature = "fast-math"))]
|
||||||
{
|
{
|
||||||
@@ -229,13 +211,8 @@ pub fn search_with_metrics(
|
|||||||
gpu_active = false;
|
gpu_active = false;
|
||||||
#[cfg(any(feature = "accelerate", feature = "openblas"))]
|
#[cfg(any(feature = "accelerate", feature = "openblas"))]
|
||||||
{
|
{
|
||||||
crate::accelerate_search::accelerate_cosine_batch(
|
crate::accelerate_search::accelerate_cosine_batch_vecs(
|
||||||
query,
|
query, vectors, norms, tombstones, k,
|
||||||
vectors_flat,
|
|
||||||
norms,
|
|
||||||
tombstones,
|
|
||||||
query.len(),
|
|
||||||
k,
|
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
#[cfg(not(any(feature = "accelerate", feature = "openblas")))]
|
#[cfg(not(any(feature = "accelerate", feature = "openblas")))]
|
||||||
@@ -348,10 +325,6 @@ mod tests {
|
|||||||
(0..n).map(|_| (0..dim).map(|_| next()).collect()).collect()
|
(0..n).map(|_| (0..dim).map(|_| next()).collect()).collect()
|
||||||
}
|
}
|
||||||
|
|
||||||
fn flatten(vectors: &[Vec<f32>]) -> Vec<f32> {
|
|
||||||
vectors.iter().flatten().copied().collect()
|
|
||||||
}
|
|
||||||
|
|
||||||
// --- auto_select_strategy tests ---
|
// --- auto_select_strategy tests ---
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -517,7 +490,6 @@ mod tests {
|
|||||||
let (results, metrics) = search_with_metrics(
|
let (results, metrics) = search_with_metrics(
|
||||||
&query,
|
&query,
|
||||||
&vectors,
|
&vectors,
|
||||||
&flatten(&vectors),
|
|
||||||
&norms,
|
&norms,
|
||||||
&tombstones,
|
&tombstones,
|
||||||
5,
|
5,
|
||||||
@@ -548,7 +520,6 @@ mod tests {
|
|||||||
let (results, metrics) = search_with_metrics(
|
let (results, metrics) = search_with_metrics(
|
||||||
&query,
|
&query,
|
||||||
&vectors,
|
&vectors,
|
||||||
&flatten(&vectors),
|
|
||||||
&norms,
|
&norms,
|
||||||
&tombstones,
|
&tombstones,
|
||||||
10,
|
10,
|
||||||
@@ -574,7 +545,6 @@ mod tests {
|
|||||||
let (_, metrics) = search_with_metrics(
|
let (_, metrics) = search_with_metrics(
|
||||||
&query,
|
&query,
|
||||||
&vectors,
|
&vectors,
|
||||||
&flatten(&vectors),
|
|
||||||
&norms,
|
&norms,
|
||||||
&tombstones,
|
&tombstones,
|
||||||
10,
|
10,
|
||||||
@@ -600,7 +570,6 @@ mod tests {
|
|||||||
let (results, _) = search_with_metrics(
|
let (results, _) = search_with_metrics(
|
||||||
&query,
|
&query,
|
||||||
&vectors,
|
&vectors,
|
||||||
&flatten(&vectors),
|
|
||||||
&norms,
|
&norms,
|
||||||
&tombstones,
|
&tombstones,
|
||||||
10,
|
10,
|
||||||
@@ -634,7 +603,6 @@ mod tests {
|
|||||||
let (results, metrics) = search_with_metrics(
|
let (results, metrics) = search_with_metrics(
|
||||||
&query,
|
&query,
|
||||||
&vectors,
|
&vectors,
|
||||||
&flatten(&vectors),
|
|
||||||
&norms,
|
&norms,
|
||||||
&tombstones,
|
&tombstones,
|
||||||
100,
|
100,
|
||||||
@@ -679,7 +647,6 @@ mod tests {
|
|||||||
let (_, metrics) = search_with_metrics(
|
let (_, metrics) = search_with_metrics(
|
||||||
&query,
|
&query,
|
||||||
&vectors,
|
&vectors,
|
||||||
&flatten(&vectors),
|
|
||||||
&norms,
|
&norms,
|
||||||
&tombstones,
|
&tombstones,
|
||||||
5,
|
5,
|
||||||
@@ -751,7 +718,6 @@ mod tests {
|
|||||||
let (results, metrics) = search_with_metrics(
|
let (results, metrics) = search_with_metrics(
|
||||||
&query,
|
&query,
|
||||||
&vectors,
|
&vectors,
|
||||||
&flatten(&vectors),
|
|
||||||
&norms,
|
&norms,
|
||||||
&tombstones,
|
&tombstones,
|
||||||
10,
|
10,
|
||||||
@@ -778,7 +744,6 @@ mod tests {
|
|||||||
let (results, metrics) = search_with_metrics(
|
let (results, metrics) = search_with_metrics(
|
||||||
&query,
|
&query,
|
||||||
&vectors,
|
&vectors,
|
||||||
&flatten(&vectors),
|
|
||||||
&norms,
|
&norms,
|
||||||
&tombstones,
|
&tombstones,
|
||||||
10,
|
10,
|
||||||
@@ -857,7 +822,6 @@ mod tests {
|
|||||||
let (results, metrics) = search_with_metrics(
|
let (results, metrics) = search_with_metrics(
|
||||||
&query,
|
&query,
|
||||||
&vectors,
|
&vectors,
|
||||||
&flatten(&vectors),
|
|
||||||
&norms,
|
&norms,
|
||||||
&tombstones,
|
&tombstones,
|
||||||
10,
|
10,
|
||||||
|
|||||||
@@ -13,57 +13,16 @@ use crate::MemoryError;
|
|||||||
|
|
||||||
const WAL_MAGIC: [u8; 4] = [0x45, 0x48, 0x57, 0x4C]; // "EHWL"
|
const WAL_MAGIC: [u8; 4] = [0x45, 0x48, 0x57, 0x4C]; // "EHWL"
|
||||||
|
|
||||||
/// Bytes before the first entry: [`WAL_MAGIC`] (4) + version (1) + entry
|
/// Current WAL format version: every entry ends with a 4-byte CRC32 trailer
|
||||||
/// count (4). Named so the offset arithmetic in `open()` — which decides
|
/// (see [`TeeReader`]) so a bit-flip is detected and replay stops there
|
||||||
/// where an append lands, and therefore whether it is replayable — reads as
|
/// instead of silently accepting corrupted data.
|
||||||
/// a header length rather than a bare 9.
|
const WAL_VERSION: u8 = 2;
|
||||||
const WAL_HEADER_LEN: u64 = WAL_MAGIC.len() as u64 + 1 + 4;
|
|
||||||
|
|
||||||
/// Current WAL format version: every entry's CRC32 trailer is computed over
|
/// The only other WAL version this crate still knows how to *read*: no
|
||||||
/// its own bytes *chained with the previous entry's stored CRC*
|
/// per-entry CRC trailer. Written by versions of this crate before the CRC32
|
||||||
/// (`crc32(entry_bytes ++ prev_crc.to_le_bytes())`, seeded with 0 for the
|
/// hardening. `WalFile::open` migrates a legacy file to [`WAL_VERSION`] by
|
||||||
/// first entry after a truncation). A per-entry CRC alone only detects a
|
/// recreating it fresh — safe because every real call site reads existing
|
||||||
/// bit-flip within that entry; chaining additionally detects entries being
|
/// entries via [`WalFile::read_entries`] before calling `open` (see
|
||||||
/// reordered, duplicated, or spliced (e.g. a Tombstone moved before/after
|
|
||||||
/// its target Save) — the moved/inserted entry's stored CRC was computed
|
|
||||||
/// against a different predecessor than the one now in front of it on disk,
|
|
||||||
/// so the chain breaks at that point and replay stops there.
|
|
||||||
const WAL_VERSION: u8 = 4;
|
|
||||||
|
|
||||||
/// The chained-CRC format before [`WalEntryType::Update`] records existed.
|
|
||||||
/// Byte-for-byte the same framing as [`WAL_VERSION`], so it is read by the
|
|
||||||
/// same code, and `WalFile::open` upgrades it in place by rewriting the
|
|
||||||
/// header's version byte (the header is not covered by the CRC chain).
|
|
||||||
///
|
|
||||||
/// The bump exists for *older binaries*: they don't know record type 0x04,
|
|
||||||
/// would treat it as a torn tail, and would truncate it — and everything
|
|
||||||
/// after it — away. An unknown header version makes them refuse the file
|
|
||||||
/// with a clear error instead.
|
|
||||||
const WAL_VERSION_CHAINED_NO_UPDATE: u8 = 3;
|
|
||||||
|
|
||||||
/// The previous WAL format version: still a CRC32 per entry (so a bit-flip
|
|
||||||
/// within one entry is caught), but not chained to the previous entry's CRC
|
|
||||||
/// (so reordering/splicing whole entries is not detected). Written by
|
|
||||||
/// versions of this crate before the chaining hardening. Fully supported for
|
|
||||||
/// reading via [`WalFile::read_entries`] — not restricted like
|
|
||||||
/// [`WAL_VERSION_LEGACY_NO_CRC`], since it still verifies each entry
|
|
||||||
/// individually. `WalFile::open` migrates it to [`WAL_VERSION`] by
|
|
||||||
/// recreating the file fresh, the same as the legacy-no-CRC migration below.
|
|
||||||
const WAL_VERSION_CRC_UNCHAINED: u8 = 2;
|
|
||||||
|
|
||||||
/// The oldest WAL version this crate still knows how to *read*: no
|
|
||||||
/// per-entry CRC trailer at all, so a bit-flip anywhere is silently
|
|
||||||
/// accepted. Written by versions of this crate before the CRC32 hardening.
|
|
||||||
/// Because of that — unlike [`WAL_VERSION_CRC_UNCHAINED`] — this version is
|
|
||||||
/// deliberately *not* reachable through the public [`WalFile::read_entries`]
|
|
||||||
/// API; only [`WalFile::read_entries_for_migration`] (used exclusively by
|
|
||||||
/// `HDF5Memory::open`'s one-time migration path) will parse it. Flipping a
|
|
||||||
/// version byte from 2/3 down to 1 no longer silently downgrades a file to
|
|
||||||
/// the fully-unverified parser for an arbitrary caller.
|
|
||||||
///
|
|
||||||
/// `WalFile::open` migrates a legacy file to [`WAL_VERSION`] by recreating
|
|
||||||
/// it fresh — safe because every real call site reads existing entries via
|
|
||||||
/// [`WalFile::read_entries_for_migration`] before calling `open` (see
|
|
||||||
/// `HDF5Memory::open`), so no data is lost.
|
/// `HDF5Memory::open`), so no data is lost.
|
||||||
const WAL_VERSION_LEGACY_NO_CRC: u8 = 1;
|
const WAL_VERSION_LEGACY_NO_CRC: u8 = 1;
|
||||||
|
|
||||||
@@ -78,10 +37,6 @@ pub enum WalEntryType {
|
|||||||
Save = 0x01,
|
Save = 0x01,
|
||||||
Tombstone = 0x02,
|
Tombstone = 0x02,
|
||||||
ActivationUpdate = 0x03,
|
ActivationUpdate = 0x03,
|
||||||
/// Replace the record at `update_index` in place (`save_or_update` hit).
|
|
||||||
/// Logged as a plain `Save` before this existed, so replay appended a
|
|
||||||
/// duplicate instead of updating.
|
|
||||||
Update = 0x04,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
impl WalEntryType {
|
impl WalEntryType {
|
||||||
@@ -90,7 +45,6 @@ impl WalEntryType {
|
|||||||
0x01 => Some(Self::Save),
|
0x01 => Some(Self::Save),
|
||||||
0x02 => Some(Self::Tombstone),
|
0x02 => Some(Self::Tombstone),
|
||||||
0x03 => Some(Self::ActivationUpdate),
|
0x03 => Some(Self::ActivationUpdate),
|
||||||
0x04 => Some(Self::Update),
|
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -107,8 +61,6 @@ pub struct WalEntry {
|
|||||||
pub tags: String,
|
pub tags: String,
|
||||||
/// For tombstone entries: the index of the entry to delete.
|
/// For tombstone entries: the index of the entry to delete.
|
||||||
pub tombstone_index: Option<usize>,
|
pub tombstone_index: Option<usize>,
|
||||||
/// For update entries: the index of the record to replace.
|
|
||||||
pub update_index: Option<usize>,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// How many entries to accumulate before updating the header entry_count.
|
/// How many entries to accumulate before updating the header entry_count.
|
||||||
@@ -125,77 +77,15 @@ pub struct WalFile {
|
|||||||
entry_count: u32,
|
entry_count: u32,
|
||||||
/// Entries written since the last header count update.
|
/// Entries written since the last header count update.
|
||||||
pending_header_sync: u32,
|
pending_header_sync: u32,
|
||||||
/// CRC32 chain state: the previous entry's stored CRC (0 if this file
|
|
||||||
/// has no entries yet), folded into the next entry's CRC computation.
|
|
||||||
/// Reset to 0 by `truncate()`/`create_fresh_wal_file`, and re-derived by
|
|
||||||
/// scanning existing entries when `open()` attaches to a non-empty file.
|
|
||||||
running_crc: u32,
|
|
||||||
/// Bytes of verified entries after the header (the length of the chain
|
|
||||||
/// `running_crc` covers). Together they form the [`WalMark`].
|
|
||||||
chain_len: u64,
|
|
||||||
}
|
|
||||||
|
|
||||||
/// What a WAL file's 9-byte header looks like, without reading any entries.
|
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
|
||||||
pub enum WalHeaderStatus {
|
|
||||||
/// A version this build can read (current or legacy).
|
|
||||||
Readable,
|
|
||||||
/// Shorter than a header — e.g. a crash while the file was being created.
|
|
||||||
/// It cannot contain entries.
|
|
||||||
Torn,
|
|
||||||
/// Not a WAL file at all.
|
|
||||||
BadMagic,
|
|
||||||
/// Well-formed header from a version this build doesn't know — most
|
|
||||||
/// likely written by a *newer* build. Never discard this: the entries are
|
|
||||||
/// probably fine, this binary just can't read them.
|
|
||||||
UnknownVersion(u8),
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Classify the header of the WAL at `path`.
|
|
||||||
pub fn wal_header_status(path: &Path) -> std::io::Result<WalHeaderStatus> {
|
|
||||||
let mut header = [0u8; WAL_HEADER_LEN as usize];
|
|
||||||
let mut f = File::open(path)?;
|
|
||||||
let mut filled = 0;
|
|
||||||
while filled < header.len() {
|
|
||||||
match f.read(&mut header[filled..])? {
|
|
||||||
0 => return Ok(WalHeaderStatus::Torn),
|
|
||||||
n => filled += n,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if header[0..4] != WAL_MAGIC {
|
|
||||||
return Ok(WalHeaderStatus::BadMagic);
|
|
||||||
}
|
|
||||||
Ok(match header[4] {
|
|
||||||
WAL_VERSION
|
|
||||||
| WAL_VERSION_CHAINED_NO_UPDATE
|
|
||||||
| WAL_VERSION_CRC_UNCHAINED
|
|
||||||
| WAL_VERSION_LEGACY_NO_CRC => WalHeaderStatus::Readable,
|
|
||||||
v => WalHeaderStatus::UnknownVersion(v),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
/// A position in a WAL's CRC chain: `len` bytes of entries after the header,
|
|
||||||
/// whose chained CRC is `crc`.
|
|
||||||
///
|
|
||||||
/// A checkpoint stores the mark of the WAL prefix it folded into the `.h5`
|
|
||||||
/// file. If the process dies after the new `.h5` is in place but before the
|
|
||||||
/// WAL is truncated, the next `open()` finds that exact prefix still in the
|
|
||||||
/// WAL and skips it instead of replaying it on top of data that already
|
|
||||||
/// contains it (which used to duplicate every pending entry).
|
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
|
||||||
pub struct WalMark {
|
|
||||||
pub len: u64,
|
|
||||||
pub crc: u32,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
impl WalFile {
|
impl WalFile {
|
||||||
/// Open or create a WAL file. If it exists, read the header and entry count.
|
/// Open or create a WAL file. If it exists, read the header and entry count.
|
||||||
///
|
///
|
||||||
/// A pre-chaining WAL file ([`WAL_VERSION_CRC_UNCHAINED`] or
|
/// A legacy (pre-CRC) WAL file is migrated to the current format by
|
||||||
/// [`WAL_VERSION_LEGACY_NO_CRC`]) is migrated to the current format by
|
/// recreating it fresh — see [`WAL_VERSION_LEGACY_NO_CRC`]. Callers that
|
||||||
/// recreating it fresh. Callers that need an existing file's entries must
|
/// need the legacy file's entries must call [`WalFile::read_entries`]
|
||||||
/// call [`WalFile::read_entries`] (or, for a legacy-no-CRC file,
|
/// first, before calling `open`.
|
||||||
/// [`WalFile::read_entries_for_migration`]) first, before calling `open`.
|
|
||||||
pub fn open(path: &Path) -> Result<Self, MemoryError> {
|
pub fn open(path: &Path) -> Result<Self, MemoryError> {
|
||||||
if path.exists() {
|
if path.exists() {
|
||||||
// Read existing header
|
// Read existing header
|
||||||
@@ -212,71 +102,20 @@ impl WalFile {
|
|||||||
let mut ver = [0u8; 1];
|
let mut ver = [0u8; 1];
|
||||||
f.read_exact(&mut ver)?;
|
f.read_exact(&mut ver)?;
|
||||||
match ver[0] {
|
match ver[0] {
|
||||||
WAL_VERSION | WAL_VERSION_CHAINED_NO_UPDATE => {
|
WAL_VERSION => {
|
||||||
if ver[0] == WAL_VERSION_CHAINED_NO_UPDATE {
|
|
||||||
// Same framing; stamp the current version so an older
|
|
||||||
// binary refuses this file rather than truncating an
|
|
||||||
// Update record it can't parse. See the constant.
|
|
||||||
f.seek(SeekFrom::Start(4))?;
|
|
||||||
f.write_all(&[WAL_VERSION])?;
|
|
||||||
f.seek(SeekFrom::Start(5))?;
|
|
||||||
}
|
|
||||||
let mut count_buf = [0u8; 4];
|
let mut count_buf = [0u8; 4];
|
||||||
f.read_exact(&mut count_buf)?;
|
f.read_exact(&mut count_buf)?;
|
||||||
let header_count = u32::from_le_bytes(count_buf);
|
let entry_count = u32::from_le_bytes(count_buf);
|
||||||
// Scan any existing entries to resume the CRC chain
|
// Seek to end for appending
|
||||||
// correctly for further appends (the header's count may
|
f.seek(SeekFrom::End(0))?;
|
||||||
// be stale from deferred group-commit sync, same
|
|
||||||
// tolerance `read_entries` already has, so the scanned
|
|
||||||
// count is also the more accurate of the two).
|
|
||||||
let (entries, running_crc, verified_bytes) =
|
|
||||||
read_chained_entries(&mut f, 0, None);
|
|
||||||
let entry_count = if entries.is_empty() {
|
|
||||||
header_count
|
|
||||||
} else {
|
|
||||||
entries.len() as u32
|
|
||||||
};
|
|
||||||
// Position the append at the end of the VERIFIED prefix,
|
|
||||||
// and drop anything after it.
|
|
||||||
//
|
|
||||||
// This used to `seek(End(0))`, which appends PAST a torn
|
|
||||||
// tail — the ordinary outcome of a crash mid-append. The
|
|
||||||
// new entry is then chained to the last good entry, but
|
|
||||||
// sits on disk behind the garbage:
|
|
||||||
//
|
|
||||||
// [1..N verified][torn bytes][N+1 chained to N]
|
|
||||||
//
|
|
||||||
// Replay stops at the torn bytes, so N+1 is unreachable
|
|
||||||
// FOREVER even though its `append` returned Ok and synced.
|
|
||||||
// That is silent data loss in the one situation a WAL
|
|
||||||
// exists for. Truncating to the verified end is the
|
|
||||||
// standard recovery: the torn tail was never acknowledged
|
|
||||||
// to any caller, so discarding it loses nothing, and the
|
|
||||||
// chain then continues from a byte offset that matches
|
|
||||||
// `running_crc`.
|
|
||||||
let verified_end = WAL_HEADER_LEN + verified_bytes;
|
|
||||||
let file_len = f.metadata()?.len();
|
|
||||||
if file_len > verified_end {
|
|
||||||
eprintln!(
|
|
||||||
"clawhdf5-agent: WAL {} has {} unverifiable byte(s) after entry {}; \
|
|
||||||
discarding them so appends stay replayable",
|
|
||||||
path.display(),
|
|
||||||
file_len - verified_end,
|
|
||||||
entries.len()
|
|
||||||
);
|
|
||||||
f.set_len(verified_end)?;
|
|
||||||
}
|
|
||||||
f.seek(SeekFrom::Start(verified_end))?;
|
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
path: path.to_path_buf(),
|
path: path.to_path_buf(),
|
||||||
file: Some(f),
|
file: Some(f),
|
||||||
entry_count,
|
entry_count,
|
||||||
pending_header_sync: 0,
|
pending_header_sync: 0,
|
||||||
running_crc,
|
|
||||||
chain_len: verified_bytes,
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
WAL_VERSION_CRC_UNCHAINED | WAL_VERSION_LEGACY_NO_CRC => {
|
WAL_VERSION_LEGACY_NO_CRC => {
|
||||||
drop(f);
|
drop(f);
|
||||||
let f = create_fresh_wal_file(path)?;
|
let f = create_fresh_wal_file(path)?;
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
@@ -284,8 +123,6 @@ impl WalFile {
|
|||||||
file: Some(f),
|
file: Some(f),
|
||||||
entry_count: 0,
|
entry_count: 0,
|
||||||
pending_header_sync: 0,
|
pending_header_sync: 0,
|
||||||
running_crc: 0,
|
|
||||||
chain_len: 0,
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
v => Err(MemoryError::Schema(format!("unsupported WAL version {v}"))),
|
v => Err(MemoryError::Schema(format!("unsupported WAL version {v}"))),
|
||||||
@@ -297,8 +134,6 @@ impl WalFile {
|
|||||||
file: Some(f),
|
file: Some(f),
|
||||||
entry_count: 0,
|
entry_count: 0,
|
||||||
pending_header_sync: 0,
|
pending_header_sync: 0,
|
||||||
running_crc: 0,
|
|
||||||
chain_len: 0,
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -322,20 +157,8 @@ impl WalFile {
|
|||||||
4 + entry.session_id.len() +
|
4 + entry.session_id.len() +
|
||||||
4 + entry.tags.len(),
|
4 + entry.tags.len(),
|
||||||
);
|
);
|
||||||
match entry.update_index {
|
|
||||||
Some(index) => {
|
|
||||||
let index = u32::try_from(index).map_err(|_| {
|
|
||||||
MemoryError::Schema(format!("WAL update index {index} exceeds u32"))
|
|
||||||
})?;
|
|
||||||
buf.push(WalEntryType::Update as u8);
|
|
||||||
buf.extend_from_slice(&entry.timestamp.to_le_bytes());
|
|
||||||
buf.extend_from_slice(&index.to_le_bytes());
|
|
||||||
}
|
|
||||||
None => {
|
|
||||||
buf.push(WalEntryType::Save as u8);
|
buf.push(WalEntryType::Save as u8);
|
||||||
buf.extend_from_slice(&entry.timestamp.to_le_bytes());
|
buf.extend_from_slice(&entry.timestamp.to_le_bytes());
|
||||||
}
|
|
||||||
}
|
|
||||||
serialize_str(&mut buf, &entry.chunk);
|
serialize_str(&mut buf, &entry.chunk);
|
||||||
buf.extend_from_slice(&(emb_len as u32).to_le_bytes());
|
buf.extend_from_slice(&(emb_len as u32).to_le_bytes());
|
||||||
for &val in &entry.embedding {
|
for &val in &entry.embedding {
|
||||||
@@ -345,10 +168,7 @@ impl WalFile {
|
|||||||
serialize_str(&mut buf, &entry.session_id);
|
serialize_str(&mut buf, &entry.session_id);
|
||||||
serialize_str(&mut buf, &entry.tags);
|
serialize_str(&mut buf, &entry.tags);
|
||||||
|
|
||||||
// Chain this entry's CRC to the previous one's so reordering/
|
let crc = crc32(&buf);
|
||||||
// splicing entries (not just flipping a bit within one) is detected
|
|
||||||
// on replay — see WAL_VERSION's doc comment.
|
|
||||||
let crc = chained_crc(&buf, self.running_crc);
|
|
||||||
buf.extend_from_slice(&crc.to_le_bytes());
|
buf.extend_from_slice(&crc.to_le_bytes());
|
||||||
|
|
||||||
let f = self
|
let f = self
|
||||||
@@ -356,9 +176,7 @@ impl WalFile {
|
|||||||
.as_mut()
|
.as_mut()
|
||||||
.ok_or_else(|| MemoryError::Io(std::io::Error::other("WAL file not open")))?;
|
.ok_or_else(|| MemoryError::Io(std::io::Error::other("WAL file not open")))?;
|
||||||
f.write_all(&buf)?;
|
f.write_all(&buf)?;
|
||||||
self.chain_len += buf.len() as u64;
|
|
||||||
|
|
||||||
self.running_crc = crc;
|
|
||||||
self.entry_count += 1;
|
self.entry_count += 1;
|
||||||
self.pending_header_sync += 1;
|
self.pending_header_sync += 1;
|
||||||
if self.pending_header_sync >= GROUP_COMMIT_SIZE {
|
if self.pending_header_sync >= GROUP_COMMIT_SIZE {
|
||||||
@@ -373,7 +191,7 @@ impl WalFile {
|
|||||||
buf[0] = WalEntryType::Tombstone as u8;
|
buf[0] = WalEntryType::Tombstone as u8;
|
||||||
buf[1..9].copy_from_slice(×tamp.to_le_bytes());
|
buf[1..9].copy_from_slice(×tamp.to_le_bytes());
|
||||||
buf[9..13].copy_from_slice(&(index as u32).to_le_bytes());
|
buf[9..13].copy_from_slice(&(index as u32).to_le_bytes());
|
||||||
let crc = chained_crc(&buf[..13], self.running_crc);
|
let crc = crc32(&buf[..13]);
|
||||||
buf[13..17].copy_from_slice(&crc.to_le_bytes());
|
buf[13..17].copy_from_slice(&crc.to_le_bytes());
|
||||||
|
|
||||||
let f = self
|
let f = self
|
||||||
@@ -381,9 +199,7 @@ impl WalFile {
|
|||||||
.as_mut()
|
.as_mut()
|
||||||
.ok_or_else(|| MemoryError::Io(std::io::Error::other("WAL file not open")))?;
|
.ok_or_else(|| MemoryError::Io(std::io::Error::other("WAL file not open")))?;
|
||||||
f.write_all(&buf)?;
|
f.write_all(&buf)?;
|
||||||
self.chain_len += buf.len() as u64;
|
|
||||||
|
|
||||||
self.running_crc = crc;
|
|
||||||
self.entry_count += 1;
|
self.entry_count += 1;
|
||||||
self.pending_header_sync += 1;
|
self.pending_header_sync += 1;
|
||||||
if self.pending_header_sync >= GROUP_COMMIT_SIZE {
|
if self.pending_header_sync >= GROUP_COMMIT_SIZE {
|
||||||
@@ -398,46 +214,9 @@ impl WalFile {
|
|||||||
/// (and may be stale if written with deferred group-commit updates). This
|
/// (and may be stale if written with deferred group-commit updates). This
|
||||||
/// tolerates both truncated files (crash mid-write) and stale header counts
|
/// tolerates both truncated files (crash mid-write) and stale header counts
|
||||||
/// (crash before the next group-commit header sync). On a `WAL_VERSION`
|
/// (crash before the next group-commit header sync). On a `WAL_VERSION`
|
||||||
/// file, a broken CRC chain (bit-flip, or an entry reordered/duplicated/
|
/// file, a CRC32 mismatch on an entry is treated the same way — replay
|
||||||
/// spliced in) is treated the same way — replay stops there rather than
|
/// stops there rather than accepting corrupted data.
|
||||||
/// accepting corrupted or tampered data. `WAL_VERSION_CRC_UNCHAINED`
|
|
||||||
/// files are read the same way minus the chain check (each entry's own
|
|
||||||
/// CRC is still verified).
|
|
||||||
///
|
|
||||||
/// Does **not** read [`WAL_VERSION_LEGACY_NO_CRC`] files — that format has
|
|
||||||
/// no integrity verification at all, so it's only reachable through
|
|
||||||
/// [`WalFile::read_entries_for_migration`], used exclusively by
|
|
||||||
/// `HDF5Memory::open`'s one-time migration path. Calling this on a
|
|
||||||
/// legacy-no-CRC file returns a typed error instead of silently
|
|
||||||
/// downgrading to the unverified parser.
|
|
||||||
pub fn read_entries(path: &Path) -> Result<Vec<WalEntry>, MemoryError> {
|
pub fn read_entries(path: &Path) -> Result<Vec<WalEntry>, MemoryError> {
|
||||||
Self::read_entries_impl(path, false, None)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Like [`WalFile::read_entries`], but also accepts
|
|
||||||
/// [`WAL_VERSION_LEGACY_NO_CRC`] files (no per-entry integrity check at
|
|
||||||
/// all). Restricted to `pub(crate)` and named accordingly: the only
|
|
||||||
/// legitimate caller is `HDF5Memory::open`'s one-time migration of a
|
|
||||||
/// pre-CRC WAL file, which immediately recreates it in the current
|
|
||||||
/// format afterward. Do not use this for anything else.
|
|
||||||
///
|
|
||||||
/// `applied` is the checkpoint mark read from the `.h5` file, if any: if
|
|
||||||
/// the WAL's chain passes through it (same byte length, same chained
|
|
||||||
/// CRC), everything up to that point is already in the `.h5` and is
|
|
||||||
/// dropped. If it never does — the normal case, because the WAL was
|
|
||||||
/// truncated after the checkpoint — every entry is returned.
|
|
||||||
pub(crate) fn read_entries_for_migration(
|
|
||||||
path: &Path,
|
|
||||||
applied: Option<WalMark>,
|
|
||||||
) -> Result<Vec<WalEntry>, MemoryError> {
|
|
||||||
Self::read_entries_impl(path, true, applied)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn read_entries_impl(
|
|
||||||
path: &Path,
|
|
||||||
allow_legacy_no_crc: bool,
|
|
||||||
applied: Option<WalMark>,
|
|
||||||
) -> Result<Vec<WalEntry>, MemoryError> {
|
|
||||||
if !path.exists() {
|
if !path.exists() {
|
||||||
return Ok(Vec::new());
|
return Ok(Vec::new());
|
||||||
}
|
}
|
||||||
@@ -450,16 +229,10 @@ impl WalFile {
|
|||||||
}
|
}
|
||||||
// entry_count is a pre-allocation hint only — we read until EOF.
|
// entry_count is a pre-allocation hint only — we read until EOF.
|
||||||
let entry_count_hint = u32::from_le_bytes([header[5], header[6], header[7], header[8]]);
|
let entry_count_hint = u32::from_le_bytes([header[5], header[6], header[7], header[8]]);
|
||||||
|
let mut entries = Vec::with_capacity(entry_count_hint as usize);
|
||||||
|
|
||||||
match header[4] {
|
match header[4] {
|
||||||
WAL_VERSION | WAL_VERSION_CHAINED_NO_UPDATE => {
|
WAL_VERSION => loop {
|
||||||
let (entries, _final_crc, _verified_bytes) =
|
|
||||||
read_chained_entries(&mut f, 0, applied);
|
|
||||||
Ok(entries)
|
|
||||||
}
|
|
||||||
WAL_VERSION_CRC_UNCHAINED => {
|
|
||||||
let mut entries = Vec::with_capacity(entry_count_hint as usize);
|
|
||||||
loop {
|
|
||||||
let raw_and_result = {
|
let raw_and_result = {
|
||||||
let mut tee = TeeReader::new(&mut f);
|
let mut tee = TeeReader::new(&mut f);
|
||||||
let result = read_one_entry(&mut tee);
|
let result = read_one_entry(&mut tee);
|
||||||
@@ -476,37 +249,27 @@ impl WalFile {
|
|||||||
}
|
}
|
||||||
let stored_crc = u32::from_le_bytes(crc_buf);
|
let stored_crc = u32::from_le_bytes(crc_buf);
|
||||||
if crc32(&raw) != stored_crc {
|
if crc32(&raw) != stored_crc {
|
||||||
// Corruption detected — stop replay here, same as a
|
// Corruption detected — stop replay here, same as a clean
|
||||||
// clean truncation/EOF, rather than accepting the bad
|
// truncation/EOF, rather than accepting the bad entry.
|
||||||
// entry.
|
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
if let Some(entry) = entry_opt {
|
if let Some(entry) = entry_opt {
|
||||||
entries.push(entry);
|
entries.push(entry);
|
||||||
}
|
}
|
||||||
}
|
},
|
||||||
Ok(entries)
|
WAL_VERSION_LEGACY_NO_CRC => loop {
|
||||||
}
|
|
||||||
WAL_VERSION_LEGACY_NO_CRC if allow_legacy_no_crc => {
|
|
||||||
let mut entries = Vec::with_capacity(entry_count_hint as usize);
|
|
||||||
loop {
|
|
||||||
match read_one_entry(&mut f) {
|
match read_one_entry(&mut f) {
|
||||||
Err(()) => break,
|
Err(()) => break,
|
||||||
Ok(Some(entry)) => entries.push(entry),
|
Ok(Some(entry)) => entries.push(entry),
|
||||||
Ok(None) => {}
|
Ok(None) => {}
|
||||||
}
|
}
|
||||||
|
},
|
||||||
|
v => {
|
||||||
|
return Err(MemoryError::Schema(format!("unsupported WAL version {v}")));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
Ok(entries)
|
Ok(entries)
|
||||||
}
|
}
|
||||||
WAL_VERSION_LEGACY_NO_CRC => Err(MemoryError::Schema(
|
|
||||||
"WAL file is in the legacy no-CRC format (version 1), which read_entries() no \
|
|
||||||
longer accepts — it has no per-entry integrity verification. Only the one-time \
|
|
||||||
migration path (WalFile::open) can read and upgrade it."
|
|
||||||
.into(),
|
|
||||||
)),
|
|
||||||
v => Err(MemoryError::Schema(format!("unsupported WAL version {v}"))),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Truncate the WAL (after merge into .h5).
|
/// Truncate the WAL (after merge into .h5).
|
||||||
pub fn truncate(&mut self) -> Result<(), MemoryError> {
|
pub fn truncate(&mut self) -> Result<(), MemoryError> {
|
||||||
@@ -516,20 +279,9 @@ impl WalFile {
|
|||||||
self.file = Some(f);
|
self.file = Some(f);
|
||||||
self.entry_count = 0;
|
self.entry_count = 0;
|
||||||
self.pending_header_sync = 0;
|
self.pending_header_sync = 0;
|
||||||
self.running_crc = 0;
|
|
||||||
self.chain_len = 0;
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
/// The mark covering every entry currently in this WAL. Store it with a
|
|
||||||
/// checkpoint taken from the state those entries produced.
|
|
||||||
pub fn mark(&self) -> WalMark {
|
|
||||||
WalMark {
|
|
||||||
len: self.chain_len,
|
|
||||||
crc: self.running_crc,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Number of pending entries.
|
/// Number of pending entries.
|
||||||
pub fn pending_count(&self) -> u32 {
|
pub fn pending_count(&self) -> u32 {
|
||||||
self.entry_count
|
self.entry_count
|
||||||
@@ -569,28 +321,6 @@ pub fn replay_into_cache(entries: &[WalEntry], cache: &mut crate::cache::MemoryC
|
|||||||
entry.tags.clone(),
|
entry.tags.clone(),
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
WalEntryType::Update => match entry.update_index {
|
|
||||||
// The index was valid when the record was written; if the
|
|
||||||
// store no longer has it, keep the data rather than drop it.
|
|
||||||
Some(idx) if idx < cache.len() => cache.update(
|
|
||||||
idx,
|
|
||||||
entry.chunk.clone(),
|
|
||||||
entry.embedding.clone(),
|
|
||||||
entry.source_channel.clone(),
|
|
||||||
entry.timestamp,
|
|
||||||
entry.session_id.clone(),
|
|
||||||
),
|
|
||||||
_ => {
|
|
||||||
cache.push(
|
|
||||||
entry.chunk.clone(),
|
|
||||||
entry.embedding.clone(),
|
|
||||||
entry.source_channel.clone(),
|
|
||||||
entry.timestamp,
|
|
||||||
entry.session_id.clone(),
|
|
||||||
entry.tags.clone(),
|
|
||||||
);
|
|
||||||
}
|
|
||||||
},
|
|
||||||
WalEntryType::Tombstone => {
|
WalEntryType::Tombstone => {
|
||||||
if let Some(idx) = entry.tombstone_index {
|
if let Some(idx) = entry.tombstone_index {
|
||||||
cache.mark_deleted(idx);
|
cache.mark_deleted(idx);
|
||||||
@@ -643,81 +373,6 @@ fn read_embedding<R: Read>(f: &mut R) -> Result<Vec<f32>, MemoryError> {
|
|||||||
Ok(vals)
|
Ok(vals)
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Compute the CRC32 trailer for a `WAL_VERSION` entry, chaining in the
|
|
||||||
/// previous entry's stored CRC (0 for the first entry after a truncation).
|
|
||||||
fn chained_crc(entry_bytes: &[u8], prev_crc: u32) -> u32 {
|
|
||||||
let mut chained = Vec::with_capacity(entry_bytes.len() + 4);
|
|
||||||
chained.extend_from_slice(entry_bytes);
|
|
||||||
chained.extend_from_slice(&prev_crc.to_le_bytes());
|
|
||||||
crc32(&chained)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Read and verify all entries from a `WAL_VERSION` (chained-CRC) stream
|
|
||||||
/// starting at the reader's current position, given the chain state to
|
|
||||||
/// resume from (0 for a stream starting at the beginning of a fresh WAL).
|
|
||||||
///
|
|
||||||
/// Returns the parsed entries, the final running CRC — the chain state to
|
|
||||||
/// continue from for further appends — and the number of BYTES consumed by
|
|
||||||
/// those verified entries. Stops (without erroring) at the first entry that
|
|
||||||
/// fails to parse or whose stored CRC doesn't match the expected chain value
|
|
||||||
/// — a bit-flip, truncation/EOF, or an entry having been
|
|
||||||
/// reordered/duplicated/spliced all produce a chain mismatch at that point,
|
|
||||||
/// and are all handled the same way: replay stops there.
|
|
||||||
///
|
|
||||||
/// The byte count is what lets `open()` position an append at the end of the
|
|
||||||
/// VERIFIED prefix rather than at end-of-file. Appending past a torn tail
|
|
||||||
/// writes entries that replay can never reach — see `open`.
|
|
||||||
///
|
|
||||||
/// `applied`, when given, is a checkpoint mark: once the chain reaches exactly
|
|
||||||
/// that position, the entries collected so far are discarded (they are
|
|
||||||
/// already in the `.h5` file). A zero-length mark matches nothing.
|
|
||||||
fn read_chained_entries<R: Read>(
|
|
||||||
f: &mut R,
|
|
||||||
start_crc: u32,
|
|
||||||
applied: Option<WalMark>,
|
|
||||||
) -> (Vec<WalEntry>, u32, u64) {
|
|
||||||
let applied = applied.filter(|m| m.len > 0);
|
|
||||||
let mut entries = Vec::new();
|
|
||||||
let mut running_crc = start_crc;
|
|
||||||
let mut verified_bytes: u64 = 0;
|
|
||||||
loop {
|
|
||||||
let raw_and_result = {
|
|
||||||
let mut tee = TeeReader::new(f);
|
|
||||||
let result = read_one_entry(&mut tee);
|
|
||||||
(tee.into_buf(), result)
|
|
||||||
};
|
|
||||||
let (raw, result) = raw_and_result;
|
|
||||||
let entry_opt = match result {
|
|
||||||
Err(()) => break,
|
|
||||||
Ok(v) => v,
|
|
||||||
};
|
|
||||||
let mut crc_buf = [0u8; 4];
|
|
||||||
if f.read_exact(&mut crc_buf).is_err() {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
let stored_crc = u32::from_le_bytes(crc_buf);
|
|
||||||
if chained_crc(&raw, running_crc) != stored_crc {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
running_crc = stored_crc;
|
|
||||||
// Only counted once the entry AND its CRC trailer verified, so the
|
|
||||||
// offset always points just past a complete, checked entry.
|
|
||||||
verified_bytes += raw.len() as u64 + crc_buf.len() as u64;
|
|
||||||
if let Some(entry) = entry_opt {
|
|
||||||
entries.push(entry);
|
|
||||||
}
|
|
||||||
if applied
|
|
||||||
== Some(WalMark {
|
|
||||||
len: verified_bytes,
|
|
||||||
crc: running_crc,
|
|
||||||
})
|
|
||||||
{
|
|
||||||
entries.clear();
|
|
||||||
}
|
|
||||||
}
|
|
||||||
(entries, running_crc, verified_bytes)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Create a fresh WAL file at `path` with the current-version header,
|
/// Create a fresh WAL file at `path` with the current-version header,
|
||||||
/// truncating/overwriting anything already there.
|
/// truncating/overwriting anything already there.
|
||||||
fn create_fresh_wal_file(path: &Path) -> Result<File, MemoryError> {
|
fn create_fresh_wal_file(path: &Path) -> Result<File, MemoryError> {
|
||||||
@@ -775,14 +430,7 @@ fn read_one_entry<R: Read>(r: &mut R) -> Result<Option<WalEntry>, ()> {
|
|||||||
let timestamp = f64::from_le_bytes(ts_buf);
|
let timestamp = f64::from_le_bytes(ts_buf);
|
||||||
|
|
||||||
match entry_type {
|
match entry_type {
|
||||||
WalEntryType::Save | WalEntryType::Update => {
|
WalEntryType::Save => {
|
||||||
let update_index = if entry_type == WalEntryType::Update {
|
|
||||||
let mut idx_buf = [0u8; 4];
|
|
||||||
r.read_exact(&mut idx_buf).map_err(|_| ())?;
|
|
||||||
Some(u32::from_le_bytes(idx_buf) as usize)
|
|
||||||
} else {
|
|
||||||
None
|
|
||||||
};
|
|
||||||
let chunk = read_len_prefixed_str(r).map_err(|_| ())?;
|
let chunk = read_len_prefixed_str(r).map_err(|_| ())?;
|
||||||
let embedding = read_embedding(r).map_err(|_| ())?;
|
let embedding = read_embedding(r).map_err(|_| ())?;
|
||||||
let source_channel = read_len_prefixed_str(r).map_err(|_| ())?;
|
let source_channel = read_len_prefixed_str(r).map_err(|_| ())?;
|
||||||
@@ -797,7 +445,6 @@ fn read_one_entry<R: Read>(r: &mut R) -> Result<Option<WalEntry>, ()> {
|
|||||||
session_id,
|
session_id,
|
||||||
tags,
|
tags,
|
||||||
tombstone_index: None,
|
tombstone_index: None,
|
||||||
update_index,
|
|
||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
WalEntryType::Tombstone => {
|
WalEntryType::Tombstone => {
|
||||||
@@ -813,7 +460,6 @@ fn read_one_entry<R: Read>(r: &mut R) -> Result<Option<WalEntry>, ()> {
|
|||||||
session_id: String::new(),
|
session_id: String::new(),
|
||||||
tags: String::new(),
|
tags: String::new(),
|
||||||
tombstone_index: Some(idx),
|
tombstone_index: Some(idx),
|
||||||
update_index: None,
|
|
||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
WalEntryType::ActivationUpdate => Ok(None),
|
WalEntryType::ActivationUpdate => Ok(None),
|
||||||
@@ -837,7 +483,6 @@ mod tests {
|
|||||||
session_id: "sess-001".to_string(),
|
session_id: "sess-001".to_string(),
|
||||||
tags: "tag1,tag2".to_string(),
|
tags: "tag1,tag2".to_string(),
|
||||||
tombstone_index: None,
|
tombstone_index: None,
|
||||||
update_index: None,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -956,7 +601,7 @@ mod tests {
|
|||||||
let dir = TempDir::new().unwrap();
|
let dir = TempDir::new().unwrap();
|
||||||
let wal_path = dir.path().join("test.h5.wal");
|
let wal_path = dir.path().join("test.h5.wal");
|
||||||
let unicode_chunk = "Hello 世界! 🌍 émojis & ünïcödé";
|
let unicode_chunk = "Hello 世界! 🌍 émojis & ünïcödé";
|
||||||
let embedding = vec![0.1, -0.2, 3.4567, f32::MAX, f32::MIN_POSITIVE];
|
let embedding = vec![0.1, -0.2, 3.14159, f32::MAX, f32::MIN_POSITIVE];
|
||||||
{
|
{
|
||||||
let mut wal = WalFile::open(&wal_path).unwrap();
|
let mut wal = WalFile::open(&wal_path).unwrap();
|
||||||
let entry = WalEntry {
|
let entry = WalEntry {
|
||||||
@@ -968,7 +613,6 @@ mod tests {
|
|||||||
session_id: "sess-öö-123".to_string(),
|
session_id: "sess-öö-123".to_string(),
|
||||||
tags: "α,β,γ".to_string(),
|
tags: "α,β,γ".to_string(),
|
||||||
tombstone_index: None,
|
tombstone_index: None,
|
||||||
update_index: None,
|
|
||||||
};
|
};
|
||||||
wal.append_save(&entry).unwrap();
|
wal.append_save(&entry).unwrap();
|
||||||
}
|
}
|
||||||
@@ -1103,148 +747,6 @@ mod tests {
|
|||||||
assert!(entries.is_empty());
|
assert!(entries.is_empty());
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Reopen `path` and return the stored chunks in order.
|
|
||||||
fn reopen_chunks(path: &std::path::Path) -> Vec<String> {
|
|
||||||
let mem = HDF5Memory::open(path).unwrap();
|
|
||||||
mem.cache.chunks.clone()
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn crash_between_checkpoint_and_wal_truncate_does_not_duplicate() {
|
|
||||||
// flush() writes the new .h5 and only then truncates the WAL. Dying in
|
|
||||||
// between leaves BOTH a .h5 that contains the pending entries and a
|
|
||||||
// WAL that still lists them; replaying blindly used to double them.
|
|
||||||
let dir = TempDir::new().unwrap();
|
|
||||||
let config = make_config(&dir);
|
|
||||||
let h5_path = config.path.clone();
|
|
||||||
let wal_path = h5_path.with_extension("h5.wal");
|
|
||||||
let stale_wal = dir.path().join("stale.wal");
|
|
||||||
|
|
||||||
{
|
|
||||||
let mut mem = HDF5Memory::create(config).unwrap();
|
|
||||||
for name in ["a", "b", "c"] {
|
|
||||||
mem.save(make_entry(name, &[1.0, 0.0, 0.0, 0.0])).unwrap();
|
|
||||||
}
|
|
||||||
assert_eq!(mem.wal_pending_count(), 3);
|
|
||||||
std::fs::copy(&wal_path, &stale_wal).unwrap();
|
|
||||||
mem.flush_wal().unwrap();
|
|
||||||
}
|
|
||||||
// Undo the truncate: this is the on-disk state right after the crash.
|
|
||||||
std::fs::copy(&stale_wal, &wal_path).unwrap();
|
|
||||||
assert_eq!(WalFile::read_entries(&wal_path).unwrap().len(), 3);
|
|
||||||
|
|
||||||
assert_eq!(reopen_chunks(&h5_path), ["a", "b", "c"]);
|
|
||||||
|
|
||||||
// Entries appended to that same WAL after recovery are still replayed.
|
|
||||||
{
|
|
||||||
let mut mem = HDF5Memory::open(&h5_path).unwrap();
|
|
||||||
mem.save(make_entry("d", &[0.0, 1.0, 0.0, 0.0])).unwrap();
|
|
||||||
}
|
|
||||||
assert_eq!(reopen_chunks(&h5_path), ["a", "b", "c", "d"]);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn entries_written_after_a_completed_checkpoint_are_all_replayed() {
|
|
||||||
// Normal case: the checkpoint's mark refers to a WAL that has since
|
|
||||||
// been truncated, so it must not suppress anything in the new one —
|
|
||||||
// including when the new WAL grows past the old mark's length.
|
|
||||||
let dir = TempDir::new().unwrap();
|
|
||||||
let config = make_config(&dir);
|
|
||||||
let h5_path = config.path.clone();
|
|
||||||
{
|
|
||||||
let mut mem = HDF5Memory::create(config).unwrap();
|
|
||||||
mem.save(make_entry("a", &[1.0, 0.0, 0.0, 0.0])).unwrap();
|
|
||||||
mem.flush_wal().unwrap();
|
|
||||||
for name in ["b", "c", "d"] {
|
|
||||||
mem.save(make_entry(name, &[1.0, 0.0, 0.0, 0.0])).unwrap();
|
|
||||||
}
|
|
||||||
}
|
|
||||||
assert_eq!(reopen_chunks(&h5_path), ["a", "b", "c", "d"]);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn save_or_update_replays_as_update_not_duplicate() {
|
|
||||||
let dir = TempDir::new().unwrap();
|
|
||||||
let config = make_config(&dir);
|
|
||||||
let h5_path = config.path.clone();
|
|
||||||
{
|
|
||||||
let mut mem = HDF5Memory::create(config).unwrap();
|
|
||||||
let mut first = make_entry("v1", &[1.0, 0.0, 0.0, 0.0]);
|
|
||||||
first.tags = "key".into();
|
|
||||||
let mut second = make_entry("v2", &[0.0, 1.0, 0.0, 0.0]);
|
|
||||||
second.tags = "key".into();
|
|
||||||
let a = mem.save_or_update(first).unwrap();
|
|
||||||
mem.save(make_entry("other", &[0.0, 0.0, 1.0, 0.0]))
|
|
||||||
.unwrap();
|
|
||||||
let b = mem.save_or_update(second).unwrap();
|
|
||||||
assert_eq!(a, b);
|
|
||||||
assert_eq!(mem.cache.chunks, ["v2", "other"]);
|
|
||||||
// Dropped without a checkpoint: all three records live in the WAL.
|
|
||||||
}
|
|
||||||
let mem = HDF5Memory::open(&h5_path).unwrap();
|
|
||||||
assert_eq!(mem.cache.chunks, ["v2", "other"]);
|
|
||||||
assert_eq!(mem.cache.embeddings[0], [0.0, 1.0, 0.0, 0.0]);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn v3_wal_is_read_and_upgraded_in_place() {
|
|
||||||
let dir = TempDir::new().unwrap();
|
|
||||||
let wal_path = dir.path().join("old.wal");
|
|
||||||
{
|
|
||||||
let mut wal = WalFile::open(&wal_path).unwrap();
|
|
||||||
wal.append_save(&make_wal_entry("kept", &[1.0])).unwrap();
|
|
||||||
}
|
|
||||||
// Rewrite the header as the pre-Update chained format.
|
|
||||||
let mut bytes = std::fs::read(&wal_path).unwrap();
|
|
||||||
bytes[4] = WAL_VERSION_CHAINED_NO_UPDATE;
|
|
||||||
std::fs::write(&wal_path, &bytes).unwrap();
|
|
||||||
|
|
||||||
assert_eq!(WalFile::read_entries(&wal_path).unwrap().len(), 1);
|
|
||||||
{
|
|
||||||
let mut wal = WalFile::open(&wal_path).unwrap();
|
|
||||||
assert_eq!(wal.pending_count(), 1);
|
|
||||||
wal.append_save(&make_wal_entry("new", &[2.0])).unwrap();
|
|
||||||
}
|
|
||||||
assert_eq!(std::fs::read(&wal_path).unwrap()[4], WAL_VERSION);
|
|
||||||
let chunks: Vec<_> = WalFile::read_entries(&wal_path)
|
|
||||||
.unwrap()
|
|
||||||
.into_iter()
|
|
||||||
.map(|e| e.chunk)
|
|
||||||
.collect();
|
|
||||||
assert_eq!(chunks, ["kept", "new"]);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn mark_matching_is_exact() {
|
|
||||||
let dir = TempDir::new().unwrap();
|
|
||||||
let wal_path = dir.path().join("m.wal");
|
|
||||||
let mut wal = WalFile::open(&wal_path).unwrap();
|
|
||||||
wal.append_save(&make_wal_entry("one", &[1.0])).unwrap();
|
|
||||||
let after_one = wal.mark();
|
|
||||||
wal.append_save(&make_wal_entry("two", &[2.0])).unwrap();
|
|
||||||
let after_two = wal.mark();
|
|
||||||
drop(wal);
|
|
||||||
|
|
||||||
let read = |m| {
|
|
||||||
WalFile::read_entries_for_migration(&wal_path, m)
|
|
||||||
.unwrap()
|
|
||||||
.into_iter()
|
|
||||||
.map(|e| e.chunk)
|
|
||||||
.collect::<Vec<_>>()
|
|
||||||
};
|
|
||||||
assert_eq!(read(None), ["one", "two"]);
|
|
||||||
assert_eq!(read(Some(after_one)), ["two"]);
|
|
||||||
assert!(read(Some(after_two)).is_empty());
|
|
||||||
// Right length, wrong CRC (a different WAL generation): skip nothing.
|
|
||||||
let foreign = WalMark {
|
|
||||||
crc: after_one.crc ^ 1,
|
|
||||||
..after_one
|
|
||||||
};
|
|
||||||
assert_eq!(read(Some(foreign)), ["one", "two"]);
|
|
||||||
// Reopening resumes the same mark.
|
|
||||||
assert_eq!(WalFile::open(&wal_path).unwrap().mark(), after_two);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_wal_replay_on_open() {
|
fn test_wal_replay_on_open() {
|
||||||
// Test WAL replay using read_entries + replay_into_cache directly,
|
// Test WAL replay using read_entries + replay_into_cache directly,
|
||||||
@@ -1410,157 +912,16 @@ mod tests {
|
|||||||
assert_eq!(entries[0].chunk, "first");
|
assert_eq!(entries[0].chunk, "first");
|
||||||
}
|
}
|
||||||
|
|
||||||
/// A crash mid-append leaves a torn final entry. Reopening the WAL must
|
|
||||||
/// place the next append at the end of the VERIFIED prefix, not at
|
|
||||||
/// end-of-file, or that append is written behind garbage the replay
|
|
||||||
/// scanner stops at — unreachable forever despite having returned Ok.
|
|
||||||
///
|
|
||||||
/// This is the ordinary crash case, so getting it wrong loses
|
|
||||||
/// acknowledged writes in exactly the situation a WAL exists for.
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_wal_append_after_torn_tail_stays_replayable() {
|
fn test_wal_reads_legacy_v1_format_without_crc() {
|
||||||
let dir = TempDir::new().unwrap();
|
let dir = TempDir::new().unwrap();
|
||||||
let wal_path = dir.path().join("test.h5.wal");
|
let wal_path = dir.path().join("legacy.h5.wal");
|
||||||
|
|
||||||
let mut wal = WalFile::open(&wal_path).unwrap();
|
|
||||||
wal.append_save(&make_wal_entry("first", &[1.0, 2.0]))
|
|
||||||
.unwrap();
|
|
||||||
drop(wal);
|
|
||||||
|
|
||||||
// Simulate the crash: a partial entry appended after the good one.
|
|
||||||
{
|
|
||||||
use std::io::Write;
|
|
||||||
let mut f = std::fs::OpenOptions::new()
|
|
||||||
.append(true)
|
|
||||||
.open(&wal_path)
|
|
||||||
.unwrap();
|
|
||||||
f.write_all(&[0xAB, 0xCD, 0xEF, 0x01, 0x02]).unwrap();
|
|
||||||
f.flush().unwrap();
|
|
||||||
}
|
|
||||||
|
|
||||||
// Reopen and append. The torn bytes must not survive between the
|
|
||||||
// verified prefix and the new entry.
|
|
||||||
let mut wal = WalFile::open(&wal_path).unwrap();
|
|
||||||
wal.append_save(&make_wal_entry("second", &[3.0, 4.0]))
|
|
||||||
.unwrap();
|
|
||||||
drop(wal);
|
|
||||||
|
|
||||||
let entries = WalFile::read_entries(&wal_path).unwrap();
|
|
||||||
assert_eq!(
|
|
||||||
entries.len(),
|
|
||||||
2,
|
|
||||||
"the append after a torn tail must be replayable; got {} entr(y/ies) — \
|
|
||||||
the post-crash write was silently lost",
|
|
||||||
entries.len()
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Reordering two entries on disk must break the CRC chain — the
|
|
||||||
/// second entry's stored CRC was computed against the first entry's
|
|
||||||
/// real CRC, not against the chain state a reader sees after swapping
|
|
||||||
/// them, so replay stops immediately instead of accepting the tampered
|
|
||||||
/// order (INT-09).
|
|
||||||
#[test]
|
|
||||||
fn test_wal_detects_reordered_entries() {
|
|
||||||
let dir = TempDir::new().unwrap();
|
|
||||||
let wal_path = dir.path().join("test.h5.wal");
|
|
||||||
let mut wal = WalFile::open(&wal_path).unwrap();
|
|
||||||
wal.append_save(&make_wal_entry("first", &[1.0, 2.0]))
|
|
||||||
.unwrap();
|
|
||||||
let len_after_first = std::fs::metadata(&wal_path).unwrap().len() as usize;
|
|
||||||
wal.append_save(&make_wal_entry("second", &[3.0, 4.0]))
|
|
||||||
.unwrap();
|
|
||||||
let len_after_second = std::fs::metadata(&wal_path).unwrap().len() as usize;
|
|
||||||
drop(wal);
|
|
||||||
|
|
||||||
let bytes = std::fs::read(&wal_path).unwrap();
|
|
||||||
let header_len = 9usize;
|
|
||||||
let entry1_bytes = bytes[header_len..len_after_first].to_vec();
|
|
||||||
let entry2_bytes = bytes[len_after_first..len_after_second].to_vec();
|
|
||||||
|
|
||||||
let mut spliced = bytes[..header_len].to_vec();
|
|
||||||
spliced.extend_from_slice(&entry2_bytes);
|
|
||||||
spliced.extend_from_slice(&entry1_bytes);
|
|
||||||
std::fs::write(&wal_path, &spliced).unwrap();
|
|
||||||
|
|
||||||
let entries = WalFile::read_entries(&wal_path).unwrap();
|
|
||||||
assert!(
|
|
||||||
entries.is_empty(),
|
|
||||||
"reordered entries must break the CRC chain and stop replay, got {} entries",
|
|
||||||
entries.len()
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Splicing a third-party entry in between two legitimate entries (e.g.
|
|
||||||
/// moving a Tombstone in front of the Save it's meant to follow) must
|
|
||||||
/// also break the chain for everything after the splice point.
|
|
||||||
#[test]
|
|
||||||
fn test_wal_detects_spliced_entry() {
|
|
||||||
let dir = TempDir::new().unwrap();
|
|
||||||
let wal_path = dir.path().join("test.h5.wal");
|
|
||||||
let mut wal = WalFile::open(&wal_path).unwrap();
|
|
||||||
wal.append_save(&make_wal_entry("first", &[1.0])).unwrap();
|
|
||||||
let len_after_first = std::fs::metadata(&wal_path).unwrap().len() as usize;
|
|
||||||
wal.append_save(&make_wal_entry("second", &[2.0])).unwrap();
|
|
||||||
let len_after_second = std::fs::metadata(&wal_path).unwrap().len() as usize;
|
|
||||||
wal.append_save(&make_wal_entry("third", &[3.0])).unwrap();
|
|
||||||
drop(wal);
|
|
||||||
|
|
||||||
let bytes = std::fs::read(&wal_path).unwrap();
|
|
||||||
let entry2_bytes = bytes[len_after_first..len_after_second].to_vec();
|
|
||||||
|
|
||||||
// Duplicate "second" right after itself: [first][second][second][third]
|
|
||||||
let mut spliced = bytes[..len_after_second].to_vec();
|
|
||||||
spliced.extend_from_slice(&entry2_bytes);
|
|
||||||
spliced.extend_from_slice(&bytes[len_after_second..]);
|
|
||||||
std::fs::write(&wal_path, &spliced).unwrap();
|
|
||||||
|
|
||||||
let entries = WalFile::read_entries(&wal_path).unwrap();
|
|
||||||
assert_eq!(
|
|
||||||
entries.len(),
|
|
||||||
2,
|
|
||||||
"replay must stop at the spliced duplicate, keeping only the entries before it"
|
|
||||||
);
|
|
||||||
assert_eq!(entries[0].chunk, "first");
|
|
||||||
assert_eq!(entries[1].chunk, "second");
|
|
||||||
}
|
|
||||||
|
|
||||||
/// A WAL closed (without truncating) and reopened must continue the CRC
|
|
||||||
/// chain correctly for newly appended entries — this is the normal
|
|
||||||
/// crash-restart-without-flush scenario (`HDF5Memory::open` replays
|
|
||||||
/// existing entries, then reopens the same file for further appends
|
|
||||||
/// without clearing it), and must not produce a false "reordering"
|
|
||||||
/// detection for its own legitimately-appended entries.
|
|
||||||
#[test]
|
|
||||||
fn test_wal_chain_continues_across_reopen() {
|
|
||||||
let dir = TempDir::new().unwrap();
|
|
||||||
let wal_path = dir.path().join("test.h5.wal");
|
|
||||||
|
|
||||||
let mut wal = WalFile::open(&wal_path).unwrap();
|
|
||||||
wal.append_save(&make_wal_entry("first", &[1.0])).unwrap();
|
|
||||||
drop(wal); // simulate a restart without ever truncating the WAL
|
|
||||||
|
|
||||||
let mut wal2 = WalFile::open(&wal_path).unwrap();
|
|
||||||
wal2.append_save(&make_wal_entry("second", &[2.0])).unwrap();
|
|
||||||
drop(wal2);
|
|
||||||
|
|
||||||
let entries = WalFile::read_entries(&wal_path).unwrap();
|
|
||||||
assert_eq!(
|
|
||||||
entries.len(),
|
|
||||||
2,
|
|
||||||
"both pre- and post-reopen entries must replay cleanly"
|
|
||||||
);
|
|
||||||
assert_eq!(entries[0].chunk, "first");
|
|
||||||
assert_eq!(entries[1].chunk, "second");
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Build a legacy (WAL_VERSION_LEGACY_NO_CRC) WAL file containing one
|
|
||||||
/// Save entry, with no trailing CRC32.
|
|
||||||
fn build_legacy_v1_wal_bytes() -> Vec<u8> {
|
|
||||||
let mut buf = Vec::new();
|
let mut buf = Vec::new();
|
||||||
buf.extend_from_slice(&WAL_MAGIC);
|
buf.extend_from_slice(&WAL_MAGIC);
|
||||||
buf.push(WAL_VERSION_LEGACY_NO_CRC);
|
buf.push(WAL_VERSION_LEGACY_NO_CRC);
|
||||||
buf.extend_from_slice(&1u32.to_le_bytes());
|
buf.extend_from_slice(&1u32.to_le_bytes());
|
||||||
|
// One Save entry in the old format: type + timestamp + fields, with
|
||||||
|
// no trailing CRC32.
|
||||||
buf.push(WalEntryType::Save as u8);
|
buf.push(WalEntryType::Save as u8);
|
||||||
buf.extend_from_slice(&42.0f64.to_le_bytes());
|
buf.extend_from_slice(&42.0f64.to_le_bytes());
|
||||||
serialize_str(&mut buf, "legacy-chunk");
|
serialize_str(&mut buf, "legacy-chunk");
|
||||||
@@ -1572,39 +933,14 @@ mod tests {
|
|||||||
serialize_str(&mut buf, "chan");
|
serialize_str(&mut buf, "chan");
|
||||||
serialize_str(&mut buf, "sess");
|
serialize_str(&mut buf, "sess");
|
||||||
serialize_str(&mut buf, "tags");
|
serialize_str(&mut buf, "tags");
|
||||||
buf
|
std::fs::write(&wal_path, &buf).unwrap();
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
let entries = WalFile::read_entries(&wal_path).unwrap();
|
||||||
fn test_wal_reads_legacy_v1_format_without_crc() {
|
|
||||||
let dir = TempDir::new().unwrap();
|
|
||||||
let wal_path = dir.path().join("legacy.h5.wal");
|
|
||||||
std::fs::write(&wal_path, build_legacy_v1_wal_bytes()).unwrap();
|
|
||||||
|
|
||||||
// Only the migration-only reader may read a legacy no-CRC file.
|
|
||||||
let entries = WalFile::read_entries_for_migration(&wal_path, None).unwrap();
|
|
||||||
assert_eq!(entries.len(), 1);
|
assert_eq!(entries.len(), 1);
|
||||||
assert_eq!(entries[0].chunk, "legacy-chunk");
|
assert_eq!(entries[0].chunk, "legacy-chunk");
|
||||||
assert_eq!(entries[0].embedding, vec![1.0, 2.0]);
|
assert_eq!(entries[0].embedding, vec![1.0, 2.0]);
|
||||||
}
|
}
|
||||||
|
|
||||||
/// The public `read_entries` must reject a legacy no-CRC file instead of
|
|
||||||
/// silently downgrading to the fully-unverified parser (INT-09) — flipping
|
|
||||||
/// a version byte from 2/3 down to 1 must not be a way to bypass every
|
|
||||||
/// integrity check for an arbitrary caller of the public API.
|
|
||||||
#[test]
|
|
||||||
fn test_wal_read_entries_rejects_legacy_v1_format() {
|
|
||||||
let dir = TempDir::new().unwrap();
|
|
||||||
let wal_path = dir.path().join("legacy.h5.wal");
|
|
||||||
std::fs::write(&wal_path, build_legacy_v1_wal_bytes()).unwrap();
|
|
||||||
|
|
||||||
let result = WalFile::read_entries(&wal_path);
|
|
||||||
assert!(
|
|
||||||
result.is_err(),
|
|
||||||
"read_entries() must reject a legacy no-CRC WAL file, not silently parse it"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_wal_open_migrates_legacy_v1_to_current_version() {
|
fn test_wal_open_migrates_legacy_v1_to_current_version() {
|
||||||
let dir = TempDir::new().unwrap();
|
let dir = TempDir::new().unwrap();
|
||||||
|
|||||||
@@ -1,187 +0,0 @@
|
|||||||
//! Crash-recovery matrix for `HDF5Memory`.
|
|
||||||
//!
|
|
||||||
//! A process crash leaves whatever reached the OS on disk. These tests build
|
|
||||||
//! the on-disk images such a crash can leave behind — after every operation,
|
|
||||||
//! inside the checkpoint window (new `.h5` in place, WAL not yet truncated),
|
|
||||||
//! and with the WAL torn at every possible length — then reopen each image
|
|
||||||
//! and check the recovered store against a model of what was acknowledged.
|
|
||||||
//!
|
|
||||||
//! Invariants:
|
|
||||||
//! * never a duplicated or invented record;
|
|
||||||
//! * an image taken between operations recovers *exactly* the acknowledged
|
|
||||||
//! state;
|
|
||||||
//! * a torn WAL recovers the last checkpoint plus a prefix of the operations
|
|
||||||
//! logged since.
|
|
||||||
|
|
||||||
use std::path::{Path, PathBuf};
|
|
||||||
|
|
||||||
use clawhdf5_agent::{AgentMemory, HDF5Memory, MemoryConfig, MemoryEntry};
|
|
||||||
use tempfile::TempDir;
|
|
||||||
|
|
||||||
struct Rng(u64);
|
|
||||||
|
|
||||||
impl Rng {
|
|
||||||
fn next(&mut self) -> u64 {
|
|
||||||
self.0 = self.0.wrapping_add(0x9E37_79B9_7F4A_7C15);
|
|
||||||
let mut z = self.0;
|
|
||||||
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
|
|
||||||
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
|
|
||||||
z ^ (z >> 31)
|
|
||||||
}
|
|
||||||
fn below(&mut self, n: usize) -> usize {
|
|
||||||
(self.next() % n.max(1) as u64) as usize
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn entry(chunk: &str, tags: &str) -> MemoryEntry {
|
|
||||||
MemoryEntry {
|
|
||||||
chunk: chunk.to_string(),
|
|
||||||
embedding: vec![1.0, 0.0, 0.0, 0.0],
|
|
||||||
source_channel: "test".into(),
|
|
||||||
timestamp: 1.0,
|
|
||||||
session_id: "s".into(),
|
|
||||||
tags: tags.to_string(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn wal_path(h5: &Path) -> PathBuf {
|
|
||||||
h5.with_extension("h5.wal")
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Copy the store (`.h5` + WAL) into a fresh directory, as a crash image.
|
|
||||||
fn image(h5: &Path, into: &TempDir, name: &str) -> PathBuf {
|
|
||||||
let dest = into.path().join(format!("{name}.h5"));
|
|
||||||
std::fs::copy(h5, &dest).unwrap();
|
|
||||||
if wal_path(h5).exists() {
|
|
||||||
std::fs::copy(wal_path(h5), wal_path(&dest)).unwrap();
|
|
||||||
}
|
|
||||||
dest
|
|
||||||
}
|
|
||||||
|
|
||||||
fn recovered(h5: &Path) -> Vec<String> {
|
|
||||||
// Read-only: the image must not be modified, and no lock is needed.
|
|
||||||
HDF5Memory::open_read_only(h5).unwrap().cache.chunks.clone()
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Apply one random operation to the store and to the model.
|
|
||||||
fn step(mem: &mut HDF5Memory, model: &mut Vec<String>, rng: &mut Rng, n: usize) {
|
|
||||||
match rng.below(6) {
|
|
||||||
0 => mem.flush_wal().unwrap(),
|
|
||||||
1 if !model.is_empty() => {
|
|
||||||
// Update an existing record in place, addressed by its tag.
|
|
||||||
let idx = rng.below(model.len());
|
|
||||||
let chunk = format!("u{n}");
|
|
||||||
assert_eq!(
|
|
||||||
mem.save_or_update(entry(&chunk, &format!("tag{idx}")))
|
|
||||||
.unwrap(),
|
|
||||||
idx
|
|
||||||
);
|
|
||||||
model[idx] = chunk;
|
|
||||||
}
|
|
||||||
_ => {
|
|
||||||
let chunk = format!("c{n}");
|
|
||||||
mem.save(entry(&chunk, &format!("tag{}", model.len())))
|
|
||||||
.unwrap();
|
|
||||||
model.push(chunk);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn image_after_every_operation_recovers_the_acknowledged_state() {
|
|
||||||
for seed in 0..40u64 {
|
|
||||||
let mut rng = Rng(seed);
|
|
||||||
let dir = TempDir::new().unwrap();
|
|
||||||
let images = TempDir::new().unwrap();
|
|
||||||
let mut config = MemoryConfig::new(dir.path().join("store.h5"), "agent", 4);
|
|
||||||
config.wal_enabled = true;
|
|
||||||
config.wal_max_entries = 1 + rng.below(6); // force frequent checkpoints
|
|
||||||
let h5 = config.path.clone();
|
|
||||||
let mut mem = HDF5Memory::create(config).unwrap();
|
|
||||||
let mut model = Vec::new();
|
|
||||||
|
|
||||||
for n in 0..30 {
|
|
||||||
step(&mut mem, &mut model, &mut rng, n);
|
|
||||||
let img = image(&h5, &images, &format!("s{seed}-{n}"));
|
|
||||||
assert_eq!(recovered(&img), model, "seed {seed}, after op {n}");
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn crash_inside_the_checkpoint_window_never_duplicates() {
|
|
||||||
for seed in 0..40u64 {
|
|
||||||
let mut rng = Rng(seed ^ 0xABCD);
|
|
||||||
let dir = TempDir::new().unwrap();
|
|
||||||
let images = TempDir::new().unwrap();
|
|
||||||
let mut config = MemoryConfig::new(dir.path().join("store.h5"), "agent", 4);
|
|
||||||
config.wal_enabled = true;
|
|
||||||
config.wal_max_entries = 1000; // checkpoints only when we ask
|
|
||||||
let h5 = config.path.clone();
|
|
||||||
let mut mem = HDF5Memory::create(config).unwrap();
|
|
||||||
let mut model = Vec::new();
|
|
||||||
|
|
||||||
for round in 0..4 {
|
|
||||||
for n in 0..(1 + rng.below(6)) {
|
|
||||||
step(&mut mem, &mut model, &mut rng, round * 100 + n);
|
|
||||||
}
|
|
||||||
// The WAL as it is just before the checkpoint...
|
|
||||||
let stale_wal = images.path().join(format!("stale-{seed}-{round}.wal"));
|
|
||||||
if wal_path(&h5).exists() {
|
|
||||||
std::fs::copy(wal_path(&h5), &stale_wal).unwrap();
|
|
||||||
}
|
|
||||||
mem.flush_wal().unwrap();
|
|
||||||
// ...put back next to the NEW .h5: the crash-in-the-window image.
|
|
||||||
let img = image(&h5, &images, &format!("w{seed}-{round}"));
|
|
||||||
if stale_wal.exists() {
|
|
||||||
std::fs::copy(&stale_wal, wal_path(&img)).unwrap();
|
|
||||||
}
|
|
||||||
assert_eq!(recovered(&img), model, "seed {seed}, round {round}");
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn torn_wal_recovers_checkpoint_plus_a_prefix() {
|
|
||||||
let dir = TempDir::new().unwrap();
|
|
||||||
let images = TempDir::new().unwrap();
|
|
||||||
let mut config = MemoryConfig::new(dir.path().join("store.h5"), "agent", 4);
|
|
||||||
config.wal_enabled = true;
|
|
||||||
config.wal_max_entries = 1000;
|
|
||||||
let h5 = config.path.clone();
|
|
||||||
let mut mem = HDF5Memory::create(config).unwrap();
|
|
||||||
|
|
||||||
for name in ["a", "b"] {
|
|
||||||
mem.save(entry(name, name)).unwrap();
|
|
||||||
}
|
|
||||||
mem.flush_wal().unwrap();
|
|
||||||
let checkpointed = vec!["a".to_string(), "b".to_string()];
|
|
||||||
|
|
||||||
// States the store passes through as each later op is logged.
|
|
||||||
let mut states = vec![checkpointed.clone()];
|
|
||||||
let mut model = checkpointed.clone();
|
|
||||||
mem.save(entry("c", "c")).unwrap();
|
|
||||||
model.push("c".into());
|
|
||||||
states.push(model.clone());
|
|
||||||
mem.save_or_update(entry("a2", "a")).unwrap();
|
|
||||||
model[0] = "a2".into();
|
|
||||||
states.push(model.clone());
|
|
||||||
mem.save(entry("d", "d")).unwrap();
|
|
||||||
model.push("d".into());
|
|
||||||
states.push(model.clone());
|
|
||||||
|
|
||||||
let full_wal = std::fs::read(wal_path(&h5)).unwrap();
|
|
||||||
let mut seen = std::collections::BTreeSet::new();
|
|
||||||
for len in 0..=full_wal.len() {
|
|
||||||
let img = image(&h5, &images, &format!("t{len}"));
|
|
||||||
std::fs::write(wal_path(&img), &full_wal[..len]).unwrap();
|
|
||||||
let got = recovered(&img);
|
|
||||||
let which = states
|
|
||||||
.iter()
|
|
||||||
.position(|s| *s == got)
|
|
||||||
.unwrap_or_else(|| panic!("WAL torn at {len} bytes recovered {got:?}"));
|
|
||||||
seen.insert(which);
|
|
||||||
}
|
|
||||||
// Every intermediate state is reachable, and the full WAL gives the last.
|
|
||||||
assert_eq!(seen.into_iter().collect::<Vec<_>>(), [0, 1, 2, 3]);
|
|
||||||
}
|
|
||||||
@@ -196,7 +196,7 @@ fn test_migration_round_trip() {
|
|||||||
mem.add_relation(e1, e2, "discusses", 0.8).unwrap();
|
mem.add_relation(e1, e2, "discusses", 0.8).unwrap();
|
||||||
|
|
||||||
// Verify all data transferred by reopening
|
// Verify all data transferred by reopening
|
||||||
let reopened = HDF5Memory::open_read_only(&path).unwrap();
|
let reopened = HDF5Memory::open(&path).unwrap();
|
||||||
assert_eq!(reopened.count(), 500);
|
assert_eq!(reopened.count(), 500);
|
||||||
|
|
||||||
// Verify sessions
|
// Verify sessions
|
||||||
@@ -266,7 +266,7 @@ fn test_knowledge_graph_workflow() {
|
|||||||
assert_eq!(entity.entity_type, "library");
|
assert_eq!(entity.entity_type, "library");
|
||||||
|
|
||||||
// Persistence
|
// Persistence
|
||||||
let reopened = HDF5Memory::open_read_only(&path).unwrap();
|
let reopened = HDF5Memory::open(&path).unwrap();
|
||||||
assert_eq!(reopened.knowledge().entities.len(), 4);
|
assert_eq!(reopened.knowledge().entities.len(), 4);
|
||||||
assert_eq!(reopened.knowledge().relations.len(), 4);
|
assert_eq!(reopened.knowledge().relations.len(), 4);
|
||||||
|
|
||||||
@@ -316,7 +316,7 @@ fn test_multi_session_workflow() {
|
|||||||
assert_eq!(mem.count(), 100); // 5 sessions * 20 entries
|
assert_eq!(mem.count(), 100); // 5 sessions * 20 entries
|
||||||
|
|
||||||
// Reopen and verify sessions
|
// Reopen and verify sessions
|
||||||
let reopened = HDF5Memory::open_read_only(&path).unwrap();
|
let reopened = HDF5Memory::open(&path).unwrap();
|
||||||
for sess in 0..5 {
|
for sess in 0..5 {
|
||||||
let summary = reopened
|
let summary = reopened
|
||||||
.get_session_summary(&format!("sess_{sess}"))
|
.get_session_summary(&format!("sess_{sess}"))
|
||||||
@@ -460,7 +460,7 @@ fn test_snapshot_and_continue() {
|
|||||||
assert_eq!(snap_mem.count(), 50);
|
assert_eq!(snap_mem.count(), 50);
|
||||||
|
|
||||||
// Original should have 100
|
// Original should have 100
|
||||||
let orig_mem = HDF5Memory::open_read_only(&path).unwrap();
|
let orig_mem = HDF5Memory::open(&path).unwrap();
|
||||||
assert_eq!(orig_mem.count(), 100);
|
assert_eq!(orig_mem.count(), 100);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -483,7 +483,7 @@ fn test_config_persistence_across_ops() {
|
|||||||
mem.add_session("s1", 0, 0, "ch", "summary").unwrap();
|
mem.add_session("s1", 0, 0, "ch", "summary").unwrap();
|
||||||
mem.add_entity("Entity", "type", -1).unwrap();
|
mem.add_entity("Entity", "type", -1).unwrap();
|
||||||
|
|
||||||
let reopened = HDF5Memory::open_read_only(&path).unwrap();
|
let reopened = HDF5Memory::open(&path).unwrap();
|
||||||
assert_eq!(reopened.config().embedding_dim, 128);
|
assert_eq!(reopened.config().embedding_dim, 128);
|
||||||
assert_eq!(reopened.config().embedder, "custom:my-embedder-v2");
|
assert_eq!(reopened.config().embedder, "custom:my-embedder-v2");
|
||||||
assert_eq!(reopened.config().chunk_size, 2048);
|
assert_eq!(reopened.config().chunk_size, 2048);
|
||||||
@@ -695,7 +695,7 @@ fn test_large_text_chunks() {
|
|||||||
mem.save_batch(entries).unwrap();
|
mem.save_batch(entries).unwrap();
|
||||||
|
|
||||||
// Reopen and verify
|
// Reopen and verify
|
||||||
let reopened = HDF5Memory::open_read_only(&path).unwrap();
|
let reopened = HDF5Memory::open(&path).unwrap();
|
||||||
assert_eq!(reopened.count(), 10);
|
assert_eq!(reopened.count(), 10);
|
||||||
|
|
||||||
let (_, cache, _, _) = read_cache(&path);
|
let (_, cache, _, _) = read_cache(&path);
|
||||||
@@ -752,7 +752,7 @@ fn test_interleaved_sessions_entries() {
|
|||||||
mem.flush_wal().unwrap();
|
mem.flush_wal().unwrap();
|
||||||
|
|
||||||
// Verify
|
// Verify
|
||||||
let reopened = HDF5Memory::open_read_only(&path).unwrap();
|
let reopened = HDF5Memory::open(&path).unwrap();
|
||||||
assert_eq!(reopened.count(), 6);
|
assert_eq!(reopened.count(), 6);
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
reopened.get_session_summary("s1").unwrap().as_deref(),
|
reopened.get_session_summary("s1").unwrap().as_deref(),
|
||||||
@@ -806,7 +806,7 @@ fn test_knowledge_graph_with_embeddings() {
|
|||||||
mem.add_relation(e_python, e_hdf5, "reads", 0.9).unwrap();
|
mem.add_relation(e_python, e_hdf5, "reads", 0.9).unwrap();
|
||||||
|
|
||||||
// Verify entity-embedding linkage persists
|
// Verify entity-embedding linkage persists
|
||||||
let reopened = HDF5Memory::open_read_only(&path).unwrap();
|
let reopened = HDF5Memory::open(&path).unwrap();
|
||||||
let rust_entity = reopened.knowledge().get_entity(e_rust).unwrap();
|
let rust_entity = reopened.knowledge().get_entity(e_rust).unwrap();
|
||||||
assert_eq!(rust_entity.embedding_idx, idx0 as i64);
|
assert_eq!(rust_entity.embedding_idx, idx0 as i64);
|
||||||
|
|
||||||
@@ -1048,7 +1048,7 @@ fn test_gpu_l2_fallback_works() {
|
|||||||
let tombstones = vec![0u8; 3];
|
let tombstones = vec![0u8; 3];
|
||||||
|
|
||||||
let gpu = clawhdf5_agent::gpu_search::GpuSearchBackend::try_init(&vectors, &norms, 2, 1);
|
let gpu = clawhdf5_agent::gpu_search::GpuSearchBackend::try_init(&vectors, &norms, 2, 1);
|
||||||
let results = gpu.search_l2(&[0.0, 0.0], &vectors, &tombstones, 3);
|
let results = gpu.search_l2(&vec![0.0, 0.0], &vectors, &tombstones, 3);
|
||||||
|
|
||||||
assert_eq!(results.len(), 3);
|
assert_eq!(results.len(), 3);
|
||||||
assert_eq!(results[0].0, 0);
|
assert_eq!(results[0].0, 0);
|
||||||
@@ -1099,7 +1099,7 @@ fn test_mmap_reader_direct_access() {
|
|||||||
|
|
||||||
// Open via MmapReader directly
|
// Open via MmapReader directly
|
||||||
let mmap = clawhdf5_io::MmapReader::open(&path).unwrap();
|
let mmap = clawhdf5_io::MmapReader::open(&path).unwrap();
|
||||||
assert!(!mmap.is_empty());
|
assert!(mmap.len() > 0);
|
||||||
// Verify we can read bytes at specific offsets
|
// Verify we can read bytes at specific offsets
|
||||||
let bytes = mmap.read_at(0, 8);
|
let bytes = mmap.read_at(0, 8);
|
||||||
assert!(bytes.is_some());
|
assert!(bytes.is_some());
|
||||||
@@ -1144,11 +1144,9 @@ fn test_strategy_reports_backend() {
|
|||||||
let tombstones = vec![0u8; n];
|
let tombstones = vec![0u8; n];
|
||||||
let query = vectors[0].clone();
|
let query = vectors[0].clone();
|
||||||
|
|
||||||
let flat: Vec<f32> = vectors.iter().flatten().copied().collect();
|
|
||||||
let (_, metrics) = strategy::search_with_metrics(
|
let (_, metrics) = strategy::search_with_metrics(
|
||||||
&query,
|
&query,
|
||||||
&vectors,
|
&vectors,
|
||||||
&flat,
|
|
||||||
&norms,
|
&norms,
|
||||||
&tombstones,
|
&tombstones,
|
||||||
5,
|
5,
|
||||||
|
|||||||
@@ -137,12 +137,12 @@ fn bench_hit_at_1_1014_records() {
|
|||||||
0.3,
|
0.3,
|
||||||
1,
|
1,
|
||||||
);
|
);
|
||||||
if let Some((top_idx, _)) = results.first()
|
if let Some((top_idx, _)) = results.first() {
|
||||||
&& *top_idx == target_indices[qi]
|
if *top_idx == target_indices[qi] {
|
||||||
{
|
|
||||||
hits += 1;
|
hits += 1;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
let hit_at_1 = hits as f64 / NUM_QUERIES as f64;
|
let hit_at_1 = hits as f64 / NUM_QUERIES as f64;
|
||||||
println!(
|
println!(
|
||||||
|
|||||||
@@ -105,7 +105,7 @@ fn test_heavy_tombstoning() {
|
|||||||
assert_eq!(mem.count_active(), 5000);
|
assert_eq!(mem.count_active(), 5000);
|
||||||
|
|
||||||
// Verify persistence
|
// Verify persistence
|
||||||
let reopened = HDF5Memory::open_read_only(&path).unwrap();
|
let reopened = HDF5Memory::open(&path).unwrap();
|
||||||
assert_eq!(reopened.count(), 5000);
|
assert_eq!(reopened.count(), 5000);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -163,7 +163,7 @@ fn test_large_embeddings_1536() {
|
|||||||
assert_eq!(mem.count(), 10_000);
|
assert_eq!(mem.count(), 10_000);
|
||||||
|
|
||||||
// Verify persistence
|
// Verify persistence
|
||||||
let reopened = HDF5Memory::open_read_only(&path).unwrap();
|
let reopened = HDF5Memory::open(&path).unwrap();
|
||||||
assert_eq!(reopened.count(), 10_000);
|
assert_eq!(reopened.count(), 10_000);
|
||||||
|
|
||||||
// Verify search works on large dims
|
// Verify search works on large dims
|
||||||
@@ -545,7 +545,7 @@ fn test_delete_all_entries() {
|
|||||||
assert_eq!(mem.count(), 0);
|
assert_eq!(mem.count(), 0);
|
||||||
|
|
||||||
// Verify persistence
|
// Verify persistence
|
||||||
let reopened = HDF5Memory::open_read_only(&path).unwrap();
|
let reopened = HDF5Memory::open(&path).unwrap();
|
||||||
assert_eq!(reopened.count(), 0);
|
assert_eq!(reopened.count(), 0);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -639,7 +639,7 @@ fn test_unicode_content() {
|
|||||||
];
|
];
|
||||||
mem.save_batch(entries).unwrap();
|
mem.save_batch(entries).unwrap();
|
||||||
|
|
||||||
let reopened = HDF5Memory::open_read_only(&path).unwrap();
|
let reopened = HDF5Memory::open(&path).unwrap();
|
||||||
assert_eq!(reopened.count(), 3);
|
assert_eq!(reopened.count(), 3);
|
||||||
|
|
||||||
let (_, cache, _, _) = clawhdf5_agent::storage::read_from_disk(&path).unwrap();
|
let (_, cache, _, _) = clawhdf5_agent::storage::read_from_disk(&path).unwrap();
|
||||||
@@ -685,6 +685,6 @@ fn test_rapid_save_delete_cycles() {
|
|||||||
assert_eq!(removed, 250);
|
assert_eq!(removed, 250);
|
||||||
assert_eq!(mem.count(), 250);
|
assert_eq!(mem.count(), 250);
|
||||||
|
|
||||||
let reopened = HDF5Memory::open_read_only(&path).unwrap();
|
let reopened = HDF5Memory::open(&path).unwrap();
|
||||||
assert_eq!(reopened.count(), 250);
|
assert_eq!(reopened.count(), 250);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,213 +0,0 @@
|
|||||||
//! Property tests for the write-ahead log.
|
|
||||||
//!
|
|
||||||
//! A deterministic generator (no external crates, reproducible from the seed
|
|
||||||
//! printed on failure) drives thousands of cases through two properties:
|
|
||||||
//!
|
|
||||||
//! 1. **Round trip** — whatever was appended is read back, in order, intact.
|
|
||||||
//! 2. **Prefix under corruption** — after *any* damage to the file (bit flips,
|
|
||||||
//! truncation, inserted or deleted bytes, duplicated or reordered regions),
|
|
||||||
//! reading never panics and yields an exact *prefix* of what was written.
|
|
||||||
//! This is the guarantee the chained CRC exists to provide: replay may stop
|
|
||||||
//! early, but it never returns a corrupted, reordered, or invented entry.
|
|
||||||
|
|
||||||
use clawhdf5_agent::wal::{WalEntry, WalEntryType, WalFile};
|
|
||||||
|
|
||||||
/// SplitMix64: tiny, well-distributed, and fully determined by its seed.
|
|
||||||
struct Rng(u64);
|
|
||||||
|
|
||||||
impl Rng {
|
|
||||||
fn next(&mut self) -> u64 {
|
|
||||||
self.0 = self.0.wrapping_add(0x9E37_79B9_7F4A_7C15);
|
|
||||||
let mut z = self.0;
|
|
||||||
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
|
|
||||||
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
|
|
||||||
z ^ (z >> 31)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn below(&mut self, n: usize) -> usize {
|
|
||||||
(self.next() % n.max(1) as u64) as usize
|
|
||||||
}
|
|
||||||
|
|
||||||
fn string(&mut self, max_len: usize) -> String {
|
|
||||||
const ALPHABET: &[char] = &['a', 'Z', '0', ' ', '\n', '\0', 'é', '漢', '🦀', '"'];
|
|
||||||
(0..self.below(max_len + 1))
|
|
||||||
.map(|_| ALPHABET[self.below(ALPHABET.len())])
|
|
||||||
.collect()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// What a test appended, in a form comparable with what is read back.
|
|
||||||
#[derive(Debug, Clone, PartialEq)]
|
|
||||||
enum Logged {
|
|
||||||
Save(String, Vec<u32>, String, String, String, u64),
|
|
||||||
Update(usize, String, Vec<u32>, u64),
|
|
||||||
Tombstone(usize, u64),
|
|
||||||
}
|
|
||||||
|
|
||||||
fn logged(entry: &WalEntry) -> Logged {
|
|
||||||
// Compare floats by bit pattern so NaN payloads and -0.0 count as intact.
|
|
||||||
let bits: Vec<u32> = entry.embedding.iter().map(|f| f.to_bits()).collect();
|
|
||||||
let ts = entry.timestamp.to_bits();
|
|
||||||
match entry.entry_type {
|
|
||||||
WalEntryType::Save => Logged::Save(
|
|
||||||
entry.chunk.clone(),
|
|
||||||
bits,
|
|
||||||
entry.source_channel.clone(),
|
|
||||||
entry.session_id.clone(),
|
|
||||||
entry.tags.clone(),
|
|
||||||
ts,
|
|
||||||
),
|
|
||||||
WalEntryType::Update => {
|
|
||||||
Logged::Update(entry.update_index.unwrap(), entry.chunk.clone(), bits, ts)
|
|
||||||
}
|
|
||||||
WalEntryType::Tombstone => Logged::Tombstone(entry.tombstone_index.unwrap(), ts),
|
|
||||||
WalEntryType::ActivationUpdate => unreachable!("never written by these tests"),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Append a random mix of records; return what was written.
|
|
||||||
fn write_random_wal(path: &std::path::Path, rng: &mut Rng) -> Vec<Logged> {
|
|
||||||
let mut wal = WalFile::open(path).unwrap();
|
|
||||||
let mut written = Vec::new();
|
|
||||||
for _ in 0..rng.below(12) {
|
|
||||||
let timestamp = f64::from_bits(rng.next());
|
|
||||||
if rng.below(5) == 0 {
|
|
||||||
let index = rng.below(1000);
|
|
||||||
wal.append_tombstone(index, timestamp).unwrap();
|
|
||||||
written.push(Logged::Tombstone(index, timestamp.to_bits()));
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
let update_index = (rng.below(4) == 0).then(|| rng.below(1000));
|
|
||||||
let entry = WalEntry {
|
|
||||||
entry_type: if update_index.is_some() {
|
|
||||||
WalEntryType::Update
|
|
||||||
} else {
|
|
||||||
WalEntryType::Save
|
|
||||||
},
|
|
||||||
timestamp,
|
|
||||||
chunk: rng.string(40),
|
|
||||||
embedding: (0..rng.below(9))
|
|
||||||
.map(|_| f32::from_bits(rng.next() as u32))
|
|
||||||
.collect(),
|
|
||||||
source_channel: rng.string(8),
|
|
||||||
session_id: rng.string(8),
|
|
||||||
tags: rng.string(8),
|
|
||||||
tombstone_index: None,
|
|
||||||
update_index,
|
|
||||||
};
|
|
||||||
wal.append_save(&entry).unwrap();
|
|
||||||
written.push(logged(&entry));
|
|
||||||
}
|
|
||||||
written
|
|
||||||
}
|
|
||||||
|
|
||||||
fn read_back(path: &std::path::Path) -> Option<Vec<Logged>> {
|
|
||||||
WalFile::read_entries(path)
|
|
||||||
.ok()
|
|
||||||
.map(|entries| entries.iter().map(logged).collect())
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn everything_appended_is_read_back_intact() {
|
|
||||||
let dir = tempfile::TempDir::new().unwrap();
|
|
||||||
for seed in 0..300u64 {
|
|
||||||
let path = dir.path().join(format!("rt-{seed}.wal"));
|
|
||||||
let written = write_random_wal(&path, &mut Rng(seed));
|
|
||||||
assert_eq!(read_back(&path).unwrap(), written, "seed {seed}");
|
|
||||||
// Reopening (which scans and repositions) must not disturb anything.
|
|
||||||
drop(WalFile::open(&path).unwrap());
|
|
||||||
assert_eq!(
|
|
||||||
read_back(&path).unwrap(),
|
|
||||||
written,
|
|
||||||
"seed {seed} after reopen"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Damage `bytes` in one of several ways.
|
|
||||||
fn corrupt(bytes: &mut Vec<u8>, rng: &mut Rng) {
|
|
||||||
if bytes.is_empty() {
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
match rng.below(7) {
|
|
||||||
0 => {
|
|
||||||
let i = rng.below(bytes.len());
|
|
||||||
bytes[i] ^= 1 << rng.below(8);
|
|
||||||
}
|
|
||||||
1 => bytes.truncate(rng.below(bytes.len())),
|
|
||||||
2 => {
|
|
||||||
let i = rng.below(bytes.len() + 1);
|
|
||||||
bytes.insert(i, rng.next() as u8);
|
|
||||||
}
|
|
||||||
3 => {
|
|
||||||
let i = rng.below(bytes.len());
|
|
||||||
bytes.remove(i);
|
|
||||||
}
|
|
||||||
4 => {
|
|
||||||
// Duplicate a region in place (a replayed/duplicated entry).
|
|
||||||
let a = rng.below(bytes.len());
|
|
||||||
let b = a + rng.below(bytes.len() - a);
|
|
||||||
let region = bytes[a..b].to_vec();
|
|
||||||
let at = rng.below(bytes.len() + 1);
|
|
||||||
bytes.splice(at..at, region);
|
|
||||||
}
|
|
||||||
5 => {
|
|
||||||
// Swap two regions (reordered entries).
|
|
||||||
let mid = rng.below(bytes.len());
|
|
||||||
bytes.rotate_left(mid);
|
|
||||||
}
|
|
||||||
_ => {
|
|
||||||
let i = rng.below(bytes.len());
|
|
||||||
let n = rng.below(bytes.len() - i + 1);
|
|
||||||
for b in &mut bytes[i..i + n] {
|
|
||||||
*b = rng.next() as u8;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn any_corruption_yields_a_prefix_never_a_wrong_entry() {
|
|
||||||
let dir = tempfile::TempDir::new().unwrap();
|
|
||||||
let mut shortened = 0u32;
|
|
||||||
for seed in 0..1500u64 {
|
|
||||||
let mut rng = Rng(seed ^ 0xC0FF_EE00);
|
|
||||||
let path = dir.path().join("c.wal");
|
|
||||||
let _ = std::fs::remove_file(&path);
|
|
||||||
let written = write_random_wal(&path, &mut rng);
|
|
||||||
|
|
||||||
let mut bytes = std::fs::read(&path).unwrap();
|
|
||||||
for _ in 0..=rng.below(3) {
|
|
||||||
corrupt(&mut bytes, &mut rng);
|
|
||||||
}
|
|
||||||
std::fs::write(&path, &bytes).unwrap();
|
|
||||||
|
|
||||||
// An unreadable header is a clean error; anything else is a prefix.
|
|
||||||
if let Some(read) = read_back(&path) {
|
|
||||||
assert!(
|
|
||||||
read.len() <= written.len() && read[..] == written[..read.len()],
|
|
||||||
"seed {seed}: read {read:?}\nis not a prefix of {written:?}"
|
|
||||||
);
|
|
||||||
if read.len() < written.len() {
|
|
||||||
shortened += 1;
|
|
||||||
}
|
|
||||||
// Opening for append repairs the tail; what was readable stays so,
|
|
||||||
// and a new entry lands right after it.
|
|
||||||
if let Ok(mut wal) = WalFile::open(&path) {
|
|
||||||
wal.append_tombstone(7, 1.0).unwrap();
|
|
||||||
drop(wal);
|
|
||||||
let mut expected = read.clone();
|
|
||||||
expected.push(Logged::Tombstone(7, 1.0f64.to_bits()));
|
|
||||||
assert_eq!(
|
|
||||||
read_back(&path).unwrap(),
|
|
||||||
expected,
|
|
||||||
"seed {seed} after repair"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
assert!(
|
|
||||||
shortened > 100,
|
|
||||||
"corruption rarely took effect: {shortened}"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
@@ -1,6 +1,6 @@
|
|||||||
[package]
|
[package]
|
||||||
name = "clawhdf5-android"
|
name = "clawhdf5-android"
|
||||||
version = "2.3.0"
|
version = "2.1.0"
|
||||||
edition = "2024"
|
edition = "2024"
|
||||||
description = "Android JNI bridge for edgehdf5-memory HDF5 backend"
|
description = "Android JNI bridge for edgehdf5-memory HDF5 backend"
|
||||||
license = "MIT"
|
license = "MIT"
|
||||||
|
|||||||
@@ -1,18 +1,17 @@
|
|||||||
[package]
|
[package]
|
||||||
name = "clawhdf5-ann"
|
name = "clawhdf5-ann"
|
||||||
version = "2.3.0"
|
version = "2.1.0"
|
||||||
edition = "2024"
|
edition = "2024"
|
||||||
description = "HNSW approximate nearest neighbor index stored as HDF5"
|
description = "HNSW approximate nearest neighbor index stored as HDF5"
|
||||||
license = "MIT"
|
license = "MIT"
|
||||||
repository = "https://git.redclaw.dev/quantumclaw/clawhdf5"
|
repository = "https://github.com/redclawsystems/clawhdf5"
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
keywords = ["hdf5", "ann", "hnsw", "nearest-neighbor"]
|
keywords = ["hdf5", "ann", "hnsw", "nearest-neighbor"]
|
||||||
categories = ["algorithms", "science"]
|
categories = ["algorithms", "science"]
|
||||||
|
|
||||||
[dependencies]
|
[dependencies]
|
||||||
clawhdf5-format = { path = "../clawhdf5-format", version = "2.3.0" }
|
clawhdf5-format = { path = "../clawhdf5-format", version = "2.1.0" }
|
||||||
clawhdf5-io = { path = "../clawhdf5-io", version = "2.3.0" }
|
clawhdf5-io = { path = "../clawhdf5-io", version = "2.1.0" }
|
||||||
clawhdf5-accel = { path = "../clawhdf5-accel", version = "2.3.0" }
|
|
||||||
rayon = { version = "1", optional = true }
|
rayon = { version = "1", optional = true }
|
||||||
|
|
||||||
[features]
|
[features]
|
||||||
|
|||||||
@@ -44,14 +44,32 @@ impl DistanceMetric {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Compute distance between two vectors using the given metric.
|
/// Compute distance between two vectors using the given metric.
|
||||||
///
|
|
||||||
/// Delegates to `clawhdf5-accel`'s runtime-dispatched SIMD kernels (AVX2 on
|
|
||||||
/// x86_64, NEON on aarch64, portable scalar fallback elsewhere) — this is
|
|
||||||
/// the hottest loop in both HNSW build and every `hybrid_search` query.
|
|
||||||
fn compute_distance(a: &[f32], b: &[f32], metric: DistanceMetric) -> f32 {
|
fn compute_distance(a: &[f32], b: &[f32], metric: DistanceMetric) -> f32 {
|
||||||
match metric {
|
match metric {
|
||||||
DistanceMetric::L2 => clawhdf5_accel::l2_distance(a, b),
|
DistanceMetric::L2 => {
|
||||||
DistanceMetric::Cosine => 1.0 - clawhdf5_accel::cosine_similarity(a, b),
|
let mut sum = 0.0f32;
|
||||||
|
for i in 0..a.len() {
|
||||||
|
let d = a[i] - b[i];
|
||||||
|
sum += d * d;
|
||||||
|
}
|
||||||
|
sum.sqrt()
|
||||||
|
}
|
||||||
|
DistanceMetric::Cosine => {
|
||||||
|
let mut dot = 0.0f32;
|
||||||
|
let mut norm_a = 0.0f32;
|
||||||
|
let mut norm_b = 0.0f32;
|
||||||
|
for i in 0..a.len() {
|
||||||
|
dot += a[i] * b[i];
|
||||||
|
norm_a += a[i] * a[i];
|
||||||
|
norm_b += b[i] * b[i];
|
||||||
|
}
|
||||||
|
let denom = norm_a.sqrt() * norm_b.sqrt();
|
||||||
|
if denom < f32::EPSILON {
|
||||||
|
1.0
|
||||||
|
} else {
|
||||||
|
1.0 - (dot / denom)
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1300,18 +1318,6 @@ mod tests {
|
|||||||
assert!((d - 1.0).abs() < 1e-6); // zero vector -> distance 1
|
assert!((d - 1.0).abs() < 1e-6); // zero vector -> distance 1
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn cosine_near_zero_vector() {
|
|
||||||
// Tiny-but-nonzero, identical-direction vectors: denom is well
|
|
||||||
// below f32::EPSILON but not exactly 0.0. Must still be treated
|
|
||||||
// as a degenerate/unreliable direction (distance 1, "maximally
|
|
||||||
// dissimilar"), not as an exact match (distance 0).
|
|
||||||
let a = vec![1e-4, 1e-4];
|
|
||||||
let b = vec![1e-4, 1e-4];
|
|
||||||
let d = compute_distance(&a, &b, DistanceMetric::Cosine);
|
|
||||||
assert!((d - 1.0).abs() < 1e-6);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn insert_into_empty_index() {
|
fn insert_into_empty_index() {
|
||||||
let mut index = HnswIndex::new(4, 16, DistanceMetric::L2);
|
let mut index = HnswIndex::new(4, 16, DistanceMetric::L2);
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
[package]
|
[package]
|
||||||
name = "clawhdf5-bench"
|
name = "clawhdf5-bench"
|
||||||
version = "2.3.0"
|
version = "2.1.0"
|
||||||
edition = "2024"
|
edition = "2024"
|
||||||
description = "Benchmark harnesses for clawhdf5-agent (Track 8)"
|
description = "Benchmark harnesses for clawhdf5-agent (Track 8)"
|
||||||
license = "MIT"
|
license = "MIT"
|
||||||
|
|||||||
@@ -22,9 +22,7 @@
|
|||||||
use std::time::Instant;
|
use std::time::Instant;
|
||||||
|
|
||||||
use clawhdf5_agent::bm25::BM25Index;
|
use clawhdf5_agent::bm25::BM25Index;
|
||||||
use clawhdf5_agent::consolidation::{
|
use clawhdf5_agent::consolidation::{ConsolidationConfig, ConsolidationEngine, MemorySource};
|
||||||
ConsolidationConfig, ConsolidationEngine, TrustedSource, UntrustedSource,
|
|
||||||
};
|
|
||||||
use clawhdf5_agent::hybrid::hybrid_search;
|
use clawhdf5_agent::hybrid::hybrid_search;
|
||||||
|
|
||||||
const EMBEDDING_DIM: usize = 384;
|
const EMBEDDING_DIM: usize = 384;
|
||||||
@@ -234,7 +232,7 @@ fn run_quality_benchmark() {
|
|||||||
for i in 0..SIGNAL_KEYWORDS.len() {
|
for i in 0..SIGNAL_KEYWORDS.len() {
|
||||||
let chunk = make_signal_content(i);
|
let chunk = make_signal_content(i);
|
||||||
let embedding = make_embedding(i * 1000);
|
let embedding = make_embedding(i * 1000);
|
||||||
let id = engine.add_trusted_memory(chunk, embedding, TrustedSource::Correction, now);
|
let id = engine.add_memory(chunk, embedding, MemorySource::Correction, now);
|
||||||
signal_ids.push(id);
|
signal_ids.push(id);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -242,12 +240,7 @@ fn run_quality_benchmark() {
|
|||||||
for i in 0..990 {
|
for i in 0..990 {
|
||||||
let chunk = make_noise_content(i);
|
let chunk = make_noise_content(i);
|
||||||
let embedding = make_embedding(i + 100);
|
let embedding = make_embedding(i + 100);
|
||||||
engine.add_trusted_memory(
|
engine.add_memory(chunk, embedding, MemorySource::System, now + i as f64 * 0.1);
|
||||||
chunk,
|
|
||||||
embedding,
|
|
||||||
TrustedSource::System,
|
|
||||||
now + i as f64 * 0.1,
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
println!(" → Inserted {} records total", engine.records().len());
|
println!(" → Inserted {} records total", engine.records().len());
|
||||||
@@ -340,7 +333,7 @@ fn run_cycle_time_benchmark() {
|
|||||||
for i in 0..n {
|
for i in 0..n {
|
||||||
let chunk = make_noise_content(i);
|
let chunk = make_noise_content(i);
|
||||||
let embedding = make_embedding(i);
|
let embedding = make_embedding(i);
|
||||||
engine.add_memory(chunk, embedding, UntrustedSource::User, now + i as f64);
|
engine.add_memory(chunk, embedding, MemorySource::User, now + i as f64);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Warmup
|
// Warmup
|
||||||
@@ -351,7 +344,7 @@ fn run_cycle_time_benchmark() {
|
|||||||
for i in n..(n * 2) {
|
for i in n..(n * 2) {
|
||||||
let chunk = make_noise_content(i);
|
let chunk = make_noise_content(i);
|
||||||
let embedding = make_embedding(i);
|
let embedding = make_embedding(i);
|
||||||
engine.add_memory(chunk, embedding, UntrustedSource::User, now + i as f64);
|
engine.add_memory(chunk, embedding, MemorySource::User, now + i as f64);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Timed consolidation
|
// Timed consolidation
|
||||||
@@ -417,13 +410,13 @@ fn run_memory_reduction_benchmark() {
|
|||||||
for i in 0..signal_count {
|
for i in 0..signal_count {
|
||||||
let chunk = make_signal_content(i % SIGNAL_KEYWORDS.len());
|
let chunk = make_signal_content(i % SIGNAL_KEYWORDS.len());
|
||||||
let emb = make_embedding(i * 999);
|
let emb = make_embedding(i * 999);
|
||||||
let id = engine.add_trusted_memory(chunk, emb, TrustedSource::Correction, now);
|
let id = engine.add_memory(chunk, emb, MemorySource::Correction, now);
|
||||||
signal_ids.push(id);
|
signal_ids.push(id);
|
||||||
}
|
}
|
||||||
for i in 0..noise_count {
|
for i in 0..noise_count {
|
||||||
let chunk = make_noise_content(i);
|
let chunk = make_noise_content(i);
|
||||||
let emb = make_embedding(i + 200);
|
let emb = make_embedding(i + 200);
|
||||||
engine.add_trusted_memory(chunk, emb, TrustedSource::System, now + i as f64 * 0.1);
|
engine.add_memory(chunk, emb, MemorySource::System, now + i as f64 * 0.1);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Access signal records heavily
|
// Access signal records heavily
|
||||||
|
|||||||
@@ -1,10 +1,10 @@
|
|||||||
[package]
|
[package]
|
||||||
name = "clawhdf5-cli"
|
name = "clawhdf5-cli"
|
||||||
version = "2.3.0"
|
version = "2.1.0"
|
||||||
edition = "2024"
|
edition = "2024"
|
||||||
license = "MIT"
|
license = "MIT"
|
||||||
description = "CLI for clawhdf5 agent memory — create, save, search, recall, stats"
|
description = "CLI for clawhdf5 agent memory — create, save, search, recall, stats"
|
||||||
repository = "https://git.redclaw.dev/quantumclaw/clawhdf5"
|
repository = "https://github.com/redclawsystems/clawhdf5"
|
||||||
keywords = ["hdf5", "ai", "memory", "agent", "cli"]
|
keywords = ["hdf5", "ai", "memory", "agent", "cli"]
|
||||||
categories = ["command-line-utilities", "science"]
|
categories = ["command-line-utilities", "science"]
|
||||||
readme = "../../README.md"
|
readme = "../../README.md"
|
||||||
@@ -14,7 +14,7 @@ name = "clawhdf5"
|
|||||||
path = "src/main.rs"
|
path = "src/main.rs"
|
||||||
|
|
||||||
[dependencies]
|
[dependencies]
|
||||||
clawhdf5-agent = { path = "../clawhdf5-agent", version = "2.3.0" }
|
clawhdf5-agent = { path = "../clawhdf5-agent", version = "2.1.0" }
|
||||||
clap = { version = "4", features = ["derive", "env"] }
|
clap = { version = "4", features = ["derive", "env"] }
|
||||||
serde_json = "1"
|
serde_json = "1"
|
||||||
serde = { workspace = true }
|
serde = { workspace = true }
|
||||||
|
|||||||
@@ -146,7 +146,7 @@ fn run(cli: Cli) -> Result<(), Box<dyn std::error::Error>> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
Commands::Recall { index } => {
|
Commands::Recall { index } => {
|
||||||
let mem = HDF5Memory::open_read_only(&cli.path)?;
|
let mem = HDF5Memory::open(&cli.path)?;
|
||||||
match mem.get_chunk(index) {
|
match mem.get_chunk(index) {
|
||||||
Some(content) => {
|
Some(content) => {
|
||||||
let j = serde_json::json!({ "index": index, "chunk": content });
|
let j = serde_json::json!({ "index": index, "chunk": content });
|
||||||
@@ -160,7 +160,7 @@ fn run(cli: Cli) -> Result<(), Box<dyn std::error::Error>> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
Commands::Stats => {
|
Commands::Stats => {
|
||||||
let mem = HDF5Memory::open_read_only(&cli.path)?;
|
let mem = HDF5Memory::open(&cli.path)?;
|
||||||
let cfg = mem.config();
|
let cfg = mem.config();
|
||||||
let j = serde_json::json!({
|
let j = serde_json::json!({
|
||||||
"path": cli.path.display().to_string(),
|
"path": cli.path.display().to_string(),
|
||||||
@@ -187,7 +187,7 @@ fn run(cli: Cli) -> Result<(), Box<dyn std::error::Error>> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
Commands::AgentsMd { output } => {
|
Commands::AgentsMd { output } => {
|
||||||
let mem = HDF5Memory::open_read_only(&cli.path)?;
|
let mem = HDF5Memory::open(&cli.path)?;
|
||||||
let md = mem.generate_agents_md();
|
let md = mem.generate_agents_md();
|
||||||
match output {
|
match output {
|
||||||
Some(p) => {
|
Some(p) => {
|
||||||
@@ -199,7 +199,7 @@ fn run(cli: Cli) -> Result<(), Box<dyn std::error::Error>> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
Commands::Export => {
|
Commands::Export => {
|
||||||
let mem = HDF5Memory::open_read_only(&cli.path)?;
|
let mem = HDF5Memory::open(&cli.path)?;
|
||||||
for i in 0..mem.count() {
|
for i in 0..mem.count() {
|
||||||
if let Some(chunk) = mem.get_chunk(i) {
|
if let Some(chunk) = mem.get_chunk(i) {
|
||||||
let j = serde_json::json!({ "index": i, "chunk": chunk });
|
let j = serde_json::json!({ "index": i, "chunk": chunk });
|
||||||
|
|||||||
@@ -1,10 +1,10 @@
|
|||||||
[package]
|
[package]
|
||||||
name = "clawhdf5-derive"
|
name = "clawhdf5-derive"
|
||||||
version = "2.3.0"
|
version = "2.1.0"
|
||||||
edition = "2024"
|
edition = "2024"
|
||||||
description = "Derive macros for rustyhdf5 HDF5 traits"
|
description = "Derive macros for rustyhdf5 HDF5 traits"
|
||||||
license = "MIT"
|
license = "MIT"
|
||||||
repository = "https://git.redclaw.dev/quantumclaw/clawhdf5"
|
repository = "https://github.com/redclawsystems/clawhdf5"
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
keywords = ["hdf5", "derive", "macros", "science"]
|
keywords = ["hdf5", "derive", "macros", "science"]
|
||||||
categories = ["development-tools::procedural-macro-helpers"]
|
categories = ["development-tools::procedural-macro-helpers"]
|
||||||
|
|||||||
@@ -1,10 +1,10 @@
|
|||||||
[package]
|
[package]
|
||||||
name = "clawhdf5-filters"
|
name = "clawhdf5-filters"
|
||||||
version = "2.3.0"
|
version = "2.1.0"
|
||||||
edition = "2024"
|
edition = "2024"
|
||||||
description = "Filter and compression pipeline for clawhdf5"
|
description = "Filter and compression pipeline for clawhdf5"
|
||||||
license = "MIT"
|
license = "MIT"
|
||||||
repository = "https://git.redclaw.dev/quantumclaw/clawhdf5"
|
repository = "https://github.com/redclawsystems/clawhdf5"
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
keywords = ["hdf5", "compression", "deflate", "filters"]
|
keywords = ["hdf5", "compression", "deflate", "filters"]
|
||||||
categories = ["compression", "science"]
|
categories = ["compression", "science"]
|
||||||
|
|||||||
@@ -1,10 +1,10 @@
|
|||||||
[package]
|
[package]
|
||||||
name = "clawhdf5-format"
|
name = "clawhdf5-format"
|
||||||
version = "2.3.0"
|
version = "2.1.0"
|
||||||
edition = "2024"
|
edition = "2024"
|
||||||
description = "Pure-Rust HDF5 binary format parsing and writing — no C dependencies"
|
description = "Pure-Rust HDF5 binary format parsing and writing — no C dependencies"
|
||||||
license = "MIT"
|
license = "MIT"
|
||||||
repository = "https://git.redclaw.dev/quantumclaw/clawhdf5"
|
repository = "https://github.com/redclawsystems/clawhdf5"
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
keywords = ["hdf5", "science", "data", "binary", "no-std"]
|
keywords = ["hdf5", "science", "data", "binary", "no-std"]
|
||||||
categories = ["parser-implementations", "science", "encoding", "no-std"]
|
categories = ["parser-implementations", "science", "encoding", "no-std"]
|
||||||
@@ -25,7 +25,7 @@ pco = { version = "1.0", optional = true }
|
|||||||
[dev-dependencies]
|
[dev-dependencies]
|
||||||
serde_json = "1"
|
serde_json = "1"
|
||||||
criterion = { workspace = true }
|
criterion = { workspace = true }
|
||||||
clawhdf5-derive = { path = "../clawhdf5-derive", version = "2.3.0" }
|
clawhdf5-derive = { path = "../clawhdf5-derive", version = "2.1.0" }
|
||||||
|
|
||||||
[[bench]]
|
[[bench]]
|
||||||
name = "bench"
|
name = "bench"
|
||||||
|
|||||||
Binary file not shown.
Binary file not shown.
@@ -1,9 +1,7 @@
|
|||||||
//! HDF5 Attribute message parsing (message type 0x000C).
|
//! HDF5 Attribute message parsing (message type 0x000C).
|
||||||
|
|
||||||
#[cfg(not(feature = "std"))]
|
#[cfg(not(feature = "std"))]
|
||||||
use alloc::{borrow::Cow, string::String, vec::Vec};
|
use alloc::{string::String, vec::Vec};
|
||||||
#[cfg(feature = "std")]
|
|
||||||
use std::borrow::Cow;
|
|
||||||
|
|
||||||
use crate::attribute_info::AttributeInfoMessage;
|
use crate::attribute_info::AttributeInfoMessage;
|
||||||
use crate::btree_v2::{BTreeV2Header, collect_btree_v2_records};
|
use crate::btree_v2::{BTreeV2Header, collect_btree_v2_records};
|
||||||
@@ -50,64 +48,17 @@ impl AttributeMessage {
|
|||||||
///
|
///
|
||||||
/// `length_size` is needed for dataspace dimension parsing.
|
/// `length_size` is needed for dataspace dimension parsing.
|
||||||
pub fn parse(data: &[u8], length_size: u8) -> Result<AttributeMessage, FormatError> {
|
pub fn parse(data: &[u8], length_size: u8) -> Result<AttributeMessage, FormatError> {
|
||||||
Self::parse_impl(data, length_size, None)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// [`AttributeMessage::parse`] with access to the rest of the file, which
|
|
||||||
/// is needed when the attribute's datatype or dataspace is *shared* (v2/v3
|
|
||||||
/// flag bits 0/1) — e.g. an attribute created with a committed datatype.
|
|
||||||
/// In that case the embedded bytes are a reference to the real message,
|
|
||||||
/// not the message. Without file access such an attribute is an error
|
|
||||||
/// rather than a garbage datatype.
|
|
||||||
pub fn parse_in_file(
|
|
||||||
data: &[u8],
|
|
||||||
file_data: &[u8],
|
|
||||||
offset_size: u8,
|
|
||||||
length_size: u8,
|
|
||||||
) -> Result<AttributeMessage, FormatError> {
|
|
||||||
Self::parse_impl(data, length_size, Some((file_data, offset_size)))
|
|
||||||
}
|
|
||||||
|
|
||||||
fn parse_impl(
|
|
||||||
data: &[u8],
|
|
||||||
length_size: u8,
|
|
||||||
file: Option<(&[u8], u8)>,
|
|
||||||
) -> Result<AttributeMessage, FormatError> {
|
|
||||||
ensure_len(data, 0, 2)?;
|
ensure_len(data, 0, 2)?;
|
||||||
let version = data[0];
|
let version = data[0];
|
||||||
|
|
||||||
match version {
|
match version {
|
||||||
1 => Self::parse_v1(data, length_size),
|
1 => Self::parse_v1(data, length_size),
|
||||||
2 => Self::parse_v2(data, length_size, file),
|
2 => Self::parse_v2(data, length_size),
|
||||||
3 => Self::parse_v3(data, length_size, file),
|
3 => Self::parse_v3(data, length_size),
|
||||||
_ => Err(FormatError::InvalidAttributeVersion(version)),
|
_ => Err(FormatError::InvalidAttributeVersion(version)),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// The bytes of an embedded datatype/dataspace message, following the
|
|
||||||
/// shared-message reference when `shared` is set.
|
|
||||||
fn embedded_message<'a>(
|
|
||||||
bytes: &'a [u8],
|
|
||||||
shared: bool,
|
|
||||||
msg_type: MessageType,
|
|
||||||
length_size: u8,
|
|
||||||
file: Option<(&[u8], u8)>,
|
|
||||||
) -> Result<Cow<'a, [u8]>, FormatError> {
|
|
||||||
if !shared {
|
|
||||||
return Ok(Cow::Borrowed(bytes));
|
|
||||||
}
|
|
||||||
let (file_data, offset_size) = file.ok_or(FormatError::UnresolvedSharedMessage)?;
|
|
||||||
let shared_ref = shared_message::parse_shared_ref(bytes, offset_size)?;
|
|
||||||
shared_message::resolve_shared_message(
|
|
||||||
file_data,
|
|
||||||
&shared_ref,
|
|
||||||
msg_type,
|
|
||||||
offset_size,
|
|
||||||
length_size,
|
|
||||||
)
|
|
||||||
.map(Cow::Owned)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn parse_v1(data: &[u8], length_size: u8) -> Result<AttributeMessage, FormatError> {
|
fn parse_v1(data: &[u8], length_size: u8) -> Result<AttributeMessage, FormatError> {
|
||||||
// version(1) + reserved(1) + name_size(2) + datatype_size(2) + dataspace_size(2) = 8
|
// version(1) + reserved(1) + name_size(2) + datatype_size(2) + dataspace_size(2) = 8
|
||||||
ensure_len(data, 0, 8)?;
|
ensure_len(data, 0, 8)?;
|
||||||
@@ -143,13 +94,7 @@ impl AttributeMessage {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
fn parse_v2(
|
fn parse_v2(data: &[u8], length_size: u8) -> Result<AttributeMessage, FormatError> {
|
||||||
data: &[u8],
|
|
||||||
length_size: u8,
|
|
||||||
file: Option<(&[u8], u8)>,
|
|
||||||
) -> Result<AttributeMessage, FormatError> {
|
|
||||||
// Flags: bit 0 = datatype is shared, bit 1 = dataspace is shared.
|
|
||||||
let flags = data.get(1).copied().unwrap_or(0);
|
|
||||||
// version(1) + flags(1) + name_size(2) + datatype_size(2) + dataspace_size(2) = 8
|
// version(1) + flags(1) + name_size(2) + datatype_size(2) + dataspace_size(2) = 8
|
||||||
ensure_len(data, 0, 8)?;
|
ensure_len(data, 0, 8)?;
|
||||||
let name_size = u16::from_le_bytes([data[2], data[3]]) as usize;
|
let name_size = u16::from_le_bytes([data[2], data[3]]) as usize;
|
||||||
@@ -165,26 +110,12 @@ impl AttributeMessage {
|
|||||||
|
|
||||||
// Datatype (NO padding)
|
// Datatype (NO padding)
|
||||||
ensure_len(data, pos, datatype_size)?;
|
ensure_len(data, pos, datatype_size)?;
|
||||||
let dt_bytes = Self::embedded_message(
|
let (datatype, _) = Datatype::parse(&data[pos..pos + datatype_size])?;
|
||||||
&data[pos..pos + datatype_size],
|
|
||||||
flags & 0x01 != 0,
|
|
||||||
MessageType::Datatype,
|
|
||||||
length_size,
|
|
||||||
file,
|
|
||||||
)?;
|
|
||||||
let (datatype, _) = Datatype::parse(&dt_bytes)?;
|
|
||||||
pos += datatype_size;
|
pos += datatype_size;
|
||||||
|
|
||||||
// Dataspace (NO padding)
|
// Dataspace (NO padding)
|
||||||
ensure_len(data, pos, dataspace_size)?;
|
ensure_len(data, pos, dataspace_size)?;
|
||||||
let ds_bytes = Self::embedded_message(
|
let dataspace = Dataspace::parse(&data[pos..pos + dataspace_size], length_size)?;
|
||||||
&data[pos..pos + dataspace_size],
|
|
||||||
flags & 0x02 != 0,
|
|
||||||
MessageType::Dataspace,
|
|
||||||
length_size,
|
|
||||||
file,
|
|
||||||
)?;
|
|
||||||
let dataspace = Dataspace::parse(&ds_bytes, length_size)?;
|
|
||||||
pos += dataspace_size;
|
pos += dataspace_size;
|
||||||
|
|
||||||
let raw_data = compute_raw_data(data, pos, &dataspace, &datatype);
|
let raw_data = compute_raw_data(data, pos, &dataspace, &datatype);
|
||||||
@@ -197,13 +128,7 @@ impl AttributeMessage {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
fn parse_v3(
|
fn parse_v3(data: &[u8], length_size: u8) -> Result<AttributeMessage, FormatError> {
|
||||||
data: &[u8],
|
|
||||||
length_size: u8,
|
|
||||||
file: Option<(&[u8], u8)>,
|
|
||||||
) -> Result<AttributeMessage, FormatError> {
|
|
||||||
// Flags: bit 0 = datatype is shared, bit 1 = dataspace is shared.
|
|
||||||
let flags = data.get(1).copied().unwrap_or(0);
|
|
||||||
// version(1) + flags(1) + name_size(2) + datatype_size(2) + dataspace_size(2) + encoding(1) = 9
|
// version(1) + flags(1) + name_size(2) + datatype_size(2) + dataspace_size(2) + encoding(1) = 9
|
||||||
ensure_len(data, 0, 9)?;
|
ensure_len(data, 0, 9)?;
|
||||||
let name_size = u16::from_le_bytes([data[2], data[3]]) as usize;
|
let name_size = u16::from_le_bytes([data[2], data[3]]) as usize;
|
||||||
@@ -220,26 +145,12 @@ impl AttributeMessage {
|
|||||||
|
|
||||||
// Datatype (NO padding)
|
// Datatype (NO padding)
|
||||||
ensure_len(data, pos, datatype_size)?;
|
ensure_len(data, pos, datatype_size)?;
|
||||||
let dt_bytes = Self::embedded_message(
|
let (datatype, _) = Datatype::parse(&data[pos..pos + datatype_size])?;
|
||||||
&data[pos..pos + datatype_size],
|
|
||||||
flags & 0x01 != 0,
|
|
||||||
MessageType::Datatype,
|
|
||||||
length_size,
|
|
||||||
file,
|
|
||||||
)?;
|
|
||||||
let (datatype, _) = Datatype::parse(&dt_bytes)?;
|
|
||||||
pos += datatype_size;
|
pos += datatype_size;
|
||||||
|
|
||||||
// Dataspace (NO padding)
|
// Dataspace (NO padding)
|
||||||
ensure_len(data, pos, dataspace_size)?;
|
ensure_len(data, pos, dataspace_size)?;
|
||||||
let ds_bytes = Self::embedded_message(
|
let dataspace = Dataspace::parse(&data[pos..pos + dataspace_size], length_size)?;
|
||||||
&data[pos..pos + dataspace_size],
|
|
||||||
flags & 0x02 != 0,
|
|
||||||
MessageType::Dataspace,
|
|
||||||
length_size,
|
|
||||||
file,
|
|
||||||
)?;
|
|
||||||
let dataspace = Dataspace::parse(&ds_bytes, length_size)?;
|
|
||||||
pos += dataspace_size;
|
pos += dataspace_size;
|
||||||
|
|
||||||
let raw_data = compute_raw_data(data, pos, &dataspace, &datatype);
|
let raw_data = compute_raw_data(data, pos, &dataspace, &datatype);
|
||||||
@@ -415,20 +326,10 @@ pub fn extract_attributes_full(
|
|||||||
offset_size,
|
offset_size,
|
||||||
length_size,
|
length_size,
|
||||||
)?;
|
)?;
|
||||||
let attr = AttributeMessage::parse_in_file(
|
let attr = AttributeMessage::parse(&resolved_data, length_size)?;
|
||||||
&resolved_data,
|
|
||||||
file_data,
|
|
||||||
offset_size,
|
|
||||||
length_size,
|
|
||||||
)?;
|
|
||||||
attrs.push(attr);
|
attrs.push(attr);
|
||||||
} else {
|
} else {
|
||||||
let attr = AttributeMessage::parse_in_file(
|
let attr = AttributeMessage::parse(&msg.data, length_size)?;
|
||||||
&msg.data,
|
|
||||||
file_data,
|
|
||||||
offset_size,
|
|
||||||
length_size,
|
|
||||||
)?;
|
|
||||||
attrs.push(attr);
|
attrs.push(attr);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -498,8 +399,7 @@ fn extract_dense_attributes(
|
|||||||
let attr_data = fh.read_managed_object(file_data, id_bytes, offset_size)?;
|
let attr_data = fh.read_managed_object(file_data, id_bytes, offset_size)?;
|
||||||
|
|
||||||
// The data in the heap is a complete attribute message
|
// The data in the heap is a complete attribute message
|
||||||
let attr =
|
let attr = AttributeMessage::parse(&attr_data, length_size)?;
|
||||||
AttributeMessage::parse_in_file(&attr_data, file_data, offset_size, length_size)?;
|
|
||||||
attrs.push(attr);
|
attrs.push(attr);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -572,13 +472,14 @@ mod tests {
|
|||||||
|
|
||||||
// Name padded to 8 bytes
|
// Name padded to 8 bytes
|
||||||
data.extend_from_slice(name);
|
data.extend_from_slice(name);
|
||||||
if data.len() % 8 != 0 || data.len() == 8 {
|
while data.len() % 8 != 0 || data.len() == 8 {
|
||||||
// Pad name to 8-byte boundary from start of name
|
// Pad name to 8-byte boundary from start of name
|
||||||
let name_start = 8;
|
let name_start = 8;
|
||||||
let name_padded = pad8(name_size);
|
let name_padded = pad8(name_size);
|
||||||
while data.len() < name_start + name_padded {
|
while data.len() < name_start + name_padded {
|
||||||
data.push(0);
|
data.push(0);
|
||||||
}
|
}
|
||||||
|
break;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Datatype padded to 8 bytes
|
// Datatype padded to 8 bytes
|
||||||
@@ -848,11 +749,11 @@ mod tests {
|
|||||||
data.extend_from_slice(name);
|
data.extend_from_slice(name);
|
||||||
data.extend_from_slice(&dt_bytes);
|
data.extend_from_slice(&dt_bytes);
|
||||||
data.extend_from_slice(&ds_bytes);
|
data.extend_from_slice(&ds_bytes);
|
||||||
data.extend_from_slice(&3.25f64.to_le_bytes());
|
data.extend_from_slice(&3.14f64.to_le_bytes());
|
||||||
|
|
||||||
let attr = AttributeMessage::parse(&data, 8).unwrap();
|
let attr = AttributeMessage::parse(&data, 8).unwrap();
|
||||||
let vals = attr.read_as_f64().unwrap();
|
let vals = attr.read_as_f64().unwrap();
|
||||||
assert_eq!(vals, vec![3.25]);
|
assert_eq!(vals, vec![3.14]);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
|
|||||||
@@ -416,7 +416,6 @@ fn header_max_total_records(max_leaf_nrec: u64, depth: u16) -> u64 {
|
|||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
|
|
||||||
#[allow(clippy::too_many_arguments)]
|
|
||||||
fn build_btree_v2_header(
|
fn build_btree_v2_header(
|
||||||
tree_type: u8,
|
tree_type: u8,
|
||||||
node_size: u32,
|
node_size: u32,
|
||||||
|
|||||||
@@ -132,47 +132,6 @@ fn ensure_len(data: &[u8], offset: usize, needed: usize) -> Result<(), FormatErr
|
|||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
/// `elements * elem_size` for sizes that come from the file. Dataspace and
|
|
||||||
/// chunk dimensions are untrusted 64-bit fields, so a crafted file can make
|
|
||||||
/// the plain product wrap to a small number (or to something enormous).
|
|
||||||
pub(crate) fn checked_byte_len(elements: u64, elem_size: usize) -> Result<usize, FormatError> {
|
|
||||||
usize::try_from(elements)
|
|
||||||
.ok()
|
|
||||||
.and_then(|n| n.checked_mul(elem_size))
|
|
||||||
.ok_or_else(|| {
|
|
||||||
FormatError::Overflow(format!(
|
|
||||||
"{elements} elements of {elem_size} bytes exceeds the addressable size"
|
|
||||||
))
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Product of chunk dimensions times the element size, overflow-checked.
|
|
||||||
pub(crate) fn checked_chunk_byte_len(
|
|
||||||
chunk_dims: &[usize],
|
|
||||||
elem_size: usize,
|
|
||||||
) -> Result<usize, FormatError> {
|
|
||||||
chunk_dims
|
|
||||||
.iter()
|
|
||||||
.try_fold(elem_size, |acc, &d| acc.checked_mul(d))
|
|
||||||
.ok_or_else(|| {
|
|
||||||
FormatError::Overflow(format!(
|
|
||||||
"chunk dimensions {chunk_dims:?} x {elem_size} bytes exceeds the addressable size"
|
|
||||||
))
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
/// A zero-filled output buffer of `len` bytes. `vec![0; len]` aborts the
|
|
||||||
/// process when the allocation fails; a size taken from the file must surface
|
|
||||||
/// as an error instead.
|
|
||||||
pub(crate) fn alloc_output(len: usize) -> Result<Vec<u8>, FormatError> {
|
|
||||||
let mut out = Vec::new();
|
|
||||||
out.try_reserve_exact(len).map_err(|_| {
|
|
||||||
FormatError::Overflow(format!("cannot allocate {len} bytes for dataset output"))
|
|
||||||
})?;
|
|
||||||
out.resize(len, 0);
|
|
||||||
Ok(out)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn read_offset(data: &[u8], pos: usize, size: u8) -> Result<u64, FormatError> {
|
fn read_offset(data: &[u8], pos: usize, size: u8) -> Result<u64, FormatError> {
|
||||||
let s = size as usize;
|
let s = size as usize;
|
||||||
if pos.checked_add(s).is_none_or(|end| end > data.len()) {
|
if pos.checked_add(s).is_none_or(|end| end > data.len()) {
|
||||||
@@ -362,17 +321,15 @@ pub fn generate_implicit_chunks(
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Read a chunked dataset, decompressing chunks as needed.
|
/// Read a chunked dataset, decompressing chunks as needed.
|
||||||
/// Every allocated chunk of a chunked dataset, for any supported chunk index,
|
pub fn read_chunked_data(
|
||||||
/// plus the spatial chunk dimensions. Chunks the file never allocated (sparse
|
|
||||||
/// datasets) are simply absent from the list.
|
|
||||||
pub fn list_chunks(
|
|
||||||
file_data: &[u8],
|
file_data: &[u8],
|
||||||
layout: &DataLayout,
|
layout: &DataLayout,
|
||||||
dataspace: &Dataspace,
|
dataspace: &Dataspace,
|
||||||
elem_size: usize,
|
datatype: &Datatype,
|
||||||
|
pipeline: Option<&FilterPipeline>,
|
||||||
offset_size: u8,
|
offset_size: u8,
|
||||||
length_size: u8,
|
length_size: u8,
|
||||||
) -> Result<(Vec<ChunkInfo>, Vec<usize>), FormatError> {
|
) -> Result<Vec<u8>, FormatError> {
|
||||||
let (
|
let (
|
||||||
chunk_dimensions,
|
chunk_dimensions,
|
||||||
version,
|
version,
|
||||||
@@ -406,6 +363,8 @@ pub fn list_chunks(
|
|||||||
let addr = addr_opt
|
let addr = addr_opt
|
||||||
.ok_or_else(|| FormatError::ChunkedReadError("no address for chunked layout".into()))?;
|
.ok_or_else(|| FormatError::ChunkedReadError("no address for chunked layout".into()))?;
|
||||||
|
|
||||||
|
let elem_size = datatype.type_size() as usize;
|
||||||
|
|
||||||
// Both v3 and v4 include element size as last dim (rank+1)
|
// Both v3 and v4 include element size as last dim (rank+1)
|
||||||
let ndims = chunk_dimensions.len();
|
let ndims = chunk_dimensions.len();
|
||||||
let rank = ndims
|
let rank = ndims
|
||||||
@@ -434,7 +393,7 @@ pub fn list_chunks(
|
|||||||
}
|
}
|
||||||
(4, Some(1)) => {
|
(4, Some(1)) => {
|
||||||
// Single chunk — one chunk covering the entire dataset
|
// Single chunk — one chunk covering the entire dataset
|
||||||
let chunk_byte_size = checked_chunk_byte_len(&chunk_dims, elem_size)?;
|
let chunk_byte_size: usize = chunk_dims.iter().product::<usize>() * elem_size;
|
||||||
let (csize, fmask) = if let Some(fs) = single_filtered_size {
|
let (csize, fmask) = if let Some(fs) = single_filtered_size {
|
||||||
(fs as u32, single_filter_mask.unwrap_or(0))
|
(fs as u32, single_filter_mask.unwrap_or(0))
|
||||||
} else {
|
} else {
|
||||||
@@ -494,38 +453,10 @@ pub fn list_chunks(
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
Ok((chunks, chunk_dims))
|
|
||||||
}
|
|
||||||
|
|
||||||
pub fn read_chunked_data(
|
|
||||||
file_data: &[u8],
|
|
||||||
layout: &DataLayout,
|
|
||||||
dataspace: &Dataspace,
|
|
||||||
datatype: &Datatype,
|
|
||||||
pipeline: Option<&FilterPipeline>,
|
|
||||||
offset_size: u8,
|
|
||||||
length_size: u8,
|
|
||||||
) -> Result<Vec<u8>, FormatError> {
|
|
||||||
let elem_size = datatype.type_size() as usize;
|
|
||||||
let (chunks, chunk_dims) = list_chunks(
|
|
||||||
file_data,
|
|
||||||
layout,
|
|
||||||
dataspace,
|
|
||||||
elem_size,
|
|
||||||
offset_size,
|
|
||||||
length_size,
|
|
||||||
)?;
|
|
||||||
let rank = chunk_dims.len();
|
|
||||||
let ds_dims: Vec<usize> = dataspace.dimensions.iter().map(|&d| d as usize).collect();
|
|
||||||
|
|
||||||
// Assemble output
|
// Assemble output
|
||||||
let total_bytes = checked_byte_len(dataspace.checked_num_elements()?, elem_size)?;
|
let total_elements = dataspace.num_elements() as usize;
|
||||||
if total_bytes == 0 {
|
let total_bytes = total_elements * elem_size;
|
||||||
// Also keeps the stride products below in range: with a zero-sized
|
let mut output = vec![0u8; total_bytes];
|
||||||
// dimension the total is 0 even if other dimensions are huge.
|
|
||||||
return Ok(Vec::new());
|
|
||||||
}
|
|
||||||
let mut output = alloc_output(total_bytes)?;
|
|
||||||
|
|
||||||
let mut ds_strides = vec![1usize; rank];
|
let mut ds_strides = vec![1usize; rank];
|
||||||
for i in (0..rank.saturating_sub(1)).rev() {
|
for i in (0..rank.saturating_sub(1)).rev() {
|
||||||
@@ -537,7 +468,8 @@ pub fn read_chunked_data(
|
|||||||
chunk_strides[i] = chunk_strides[i + 1] * chunk_dims[i + 1];
|
chunk_strides[i] = chunk_strides[i + 1] * chunk_dims[i + 1];
|
||||||
}
|
}
|
||||||
|
|
||||||
let chunk_total_bytes = checked_chunk_byte_len(&chunk_dims, elem_size)?;
|
let chunk_total_elements: usize = chunk_dims.iter().product();
|
||||||
|
let chunk_total_bytes = chunk_total_elements * elem_size;
|
||||||
|
|
||||||
// Fast path: no filters — copy directly from file_data without intermediate alloc
|
// Fast path: no filters — copy directly from file_data without intermediate alloc
|
||||||
if pipeline.is_none() {
|
if pipeline.is_none() {
|
||||||
@@ -691,7 +623,7 @@ pub fn read_chunked_data_cached(
|
|||||||
let chunks = match (version, chunk_index_type) {
|
let chunks = match (version, chunk_index_type) {
|
||||||
(3, _) => collect_chunk_info(file_data, addr, ndims, offset_size, length_size)?,
|
(3, _) => collect_chunk_info(file_data, addr, ndims, offset_size, length_size)?,
|
||||||
(4, Some(1)) => {
|
(4, Some(1)) => {
|
||||||
let chunk_byte_size = checked_chunk_byte_len(&chunk_dims, elem_size)?;
|
let chunk_byte_size: usize = chunk_dims.iter().product::<usize>() * elem_size;
|
||||||
let (csize, fmask) = if let Some(fs) = single_filtered_size {
|
let (csize, fmask) = if let Some(fs) = single_filtered_size {
|
||||||
(fs as u32, single_filter_mask.unwrap_or(0))
|
(fs as u32, single_filter_mask.unwrap_or(0))
|
||||||
} else {
|
} else {
|
||||||
@@ -757,13 +689,9 @@ pub fn read_chunked_data_cached(
|
|||||||
let chunks = cache.all_indexed_chunks().unwrap_or_default();
|
let chunks = cache.all_indexed_chunks().unwrap_or_default();
|
||||||
|
|
||||||
// Assemble output
|
// Assemble output
|
||||||
let total_bytes = checked_byte_len(dataspace.checked_num_elements()?, elem_size)?;
|
let total_elements = dataspace.num_elements() as usize;
|
||||||
if total_bytes == 0 {
|
let total_bytes = total_elements * elem_size;
|
||||||
// Also keeps the stride products below in range: with a zero-sized
|
let mut output = vec![0u8; total_bytes];
|
||||||
// dimension the total is 0 even if other dimensions are huge.
|
|
||||||
return Ok(Vec::new());
|
|
||||||
}
|
|
||||||
let mut output = alloc_output(total_bytes)?;
|
|
||||||
|
|
||||||
let mut ds_strides = vec![1usize; rank];
|
let mut ds_strides = vec![1usize; rank];
|
||||||
for i in (0..rank.saturating_sub(1)).rev() {
|
for i in (0..rank.saturating_sub(1)).rev() {
|
||||||
@@ -775,7 +703,8 @@ pub fn read_chunked_data_cached(
|
|||||||
chunk_strides[i] = chunk_strides[i + 1] * chunk_dims[i + 1];
|
chunk_strides[i] = chunk_strides[i + 1] * chunk_dims[i + 1];
|
||||||
}
|
}
|
||||||
|
|
||||||
let chunk_total_bytes = checked_chunk_byte_len(&chunk_dims, elem_size)?;
|
let chunk_total_elements: usize = chunk_dims.iter().product();
|
||||||
|
let chunk_total_bytes = chunk_total_elements * elem_size;
|
||||||
|
|
||||||
for chunk_info in &chunks {
|
for chunk_info in &chunks {
|
||||||
let coord: Vec<u64> = chunk_info.offsets.iter().take(rank).copied().collect();
|
let coord: Vec<u64> = chunk_info.offsets.iter().take(rank).copied().collect();
|
||||||
@@ -1047,7 +976,7 @@ pub fn read_chunked_data_sweep(
|
|||||||
let chunks = match (version, chunk_index_type) {
|
let chunks = match (version, chunk_index_type) {
|
||||||
(3, _) => collect_chunk_info(file_data, addr, ndims, offset_size, length_size)?,
|
(3, _) => collect_chunk_info(file_data, addr, ndims, offset_size, length_size)?,
|
||||||
(4, Some(1)) => {
|
(4, Some(1)) => {
|
||||||
let chunk_byte_size = checked_chunk_byte_len(&chunk_dims, elem_size)?;
|
let chunk_byte_size: usize = chunk_dims.iter().product::<usize>() * elem_size;
|
||||||
let (csize, fmask) = if let Some(fs) = single_filtered_size {
|
let (csize, fmask) = if let Some(fs) = single_filtered_size {
|
||||||
(fs as u32, single_filter_mask.unwrap_or(0))
|
(fs as u32, single_filter_mask.unwrap_or(0))
|
||||||
} else {
|
} else {
|
||||||
@@ -1113,13 +1042,9 @@ pub fn read_chunked_data_sweep(
|
|||||||
let chunks = cache.all_indexed_chunks().unwrap_or_default();
|
let chunks = cache.all_indexed_chunks().unwrap_or_default();
|
||||||
|
|
||||||
// Assemble output
|
// Assemble output
|
||||||
let total_bytes = checked_byte_len(dataspace.checked_num_elements()?, elem_size)?;
|
let total_elements = dataspace.num_elements() as usize;
|
||||||
if total_bytes == 0 {
|
let total_bytes = total_elements * elem_size;
|
||||||
// Also keeps the stride products below in range: with a zero-sized
|
let mut output = vec![0u8; total_bytes];
|
||||||
// dimension the total is 0 even if other dimensions are huge.
|
|
||||||
return Ok(Vec::new());
|
|
||||||
}
|
|
||||||
let mut output = alloc_output(total_bytes)?;
|
|
||||||
|
|
||||||
let mut ds_strides = vec![1usize; rank];
|
let mut ds_strides = vec![1usize; rank];
|
||||||
for i in (0..rank.saturating_sub(1)).rev() {
|
for i in (0..rank.saturating_sub(1)).rev() {
|
||||||
@@ -1131,7 +1056,8 @@ pub fn read_chunked_data_sweep(
|
|||||||
chunk_strides[i] = chunk_strides[i + 1] * chunk_dims[i + 1];
|
chunk_strides[i] = chunk_strides[i + 1] * chunk_dims[i + 1];
|
||||||
}
|
}
|
||||||
|
|
||||||
let chunk_total_bytes = checked_chunk_byte_len(&chunk_dims, elem_size)?;
|
let chunk_total_elements: usize = chunk_dims.iter().product();
|
||||||
|
let chunk_total_bytes = chunk_total_elements * elem_size;
|
||||||
|
|
||||||
for chunk_info in &chunks {
|
for chunk_info in &chunks {
|
||||||
let coord: Vec<u64> = chunk_info.offsets.iter().take(rank).copied().collect();
|
let coord: Vec<u64> = chunk_info.offsets.iter().take(rank).copied().collect();
|
||||||
@@ -1273,7 +1199,7 @@ pub fn read_chunked_data_indexed(
|
|||||||
let chunks = match (version, chunk_index_type) {
|
let chunks = match (version, chunk_index_type) {
|
||||||
(3, _) => collect_chunk_info(file_data, addr, ndims, offset_size, length_size)?,
|
(3, _) => collect_chunk_info(file_data, addr, ndims, offset_size, length_size)?,
|
||||||
(4, Some(1)) => {
|
(4, Some(1)) => {
|
||||||
let chunk_byte_size = checked_chunk_byte_len(&chunk_dims, elem_size)?;
|
let chunk_byte_size: usize = chunk_dims.iter().product::<usize>() * elem_size;
|
||||||
let (csize, fmask) = if let Some(fs) = single_filtered_size {
|
let (csize, fmask) = if let Some(fs) = single_filtered_size {
|
||||||
(fs as u32, single_filter_mask.unwrap_or(0))
|
(fs as u32, single_filter_mask.unwrap_or(0))
|
||||||
} else {
|
} else {
|
||||||
@@ -1537,64 +1463,6 @@ fn copy_chunk_to_output(
|
|||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
|
|
||||||
fn simple_space(dimensions: Vec<u64>) -> Dataspace {
|
|
||||||
Dataspace {
|
|
||||||
space_type: crate::dataspace::DataspaceType::Simple,
|
|
||||||
rank: dimensions.len() as u8,
|
|
||||||
dimensions,
|
|
||||||
max_dimensions: None,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn crafted_dimensions_are_errors_not_wraparound() {
|
|
||||||
// 2^63 * 2 wraps to 0 with a plain product; 2^40 * 2^40 wraps too.
|
|
||||||
for dims in [
|
|
||||||
vec![1u64 << 63, 2],
|
|
||||||
vec![1 << 40, 1 << 40],
|
|
||||||
vec![u64::MAX, u64::MAX],
|
|
||||||
] {
|
|
||||||
let space = simple_space(dims.clone());
|
|
||||||
assert!(
|
|
||||||
matches!(space.checked_num_elements(), Err(FormatError::Overflow(_))),
|
|
||||||
"{dims:?}"
|
|
||||||
);
|
|
||||||
// The infallible accessor saturates instead of wrapping.
|
|
||||||
assert_eq!(space.num_elements(), u64::MAX, "{dims:?}");
|
|
||||||
}
|
|
||||||
assert_eq!(simple_space(vec![3, 4]).checked_num_elements().unwrap(), 12);
|
|
||||||
// A zero-sized dimension makes the whole product 0, not an overflow.
|
|
||||||
assert_eq!(
|
|
||||||
simple_space(vec![0, 1 << 40, 1 << 40])
|
|
||||||
.checked_num_elements()
|
|
||||||
.unwrap(),
|
|
||||||
0
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn byte_length_helpers_check_overflow() {
|
|
||||||
assert_eq!(checked_byte_len(10, 8).unwrap(), 80);
|
|
||||||
assert!(matches!(
|
|
||||||
checked_byte_len(u64::MAX, 8),
|
|
||||||
Err(FormatError::Overflow(_))
|
|
||||||
));
|
|
||||||
assert_eq!(checked_chunk_byte_len(&[10, 10], 4).unwrap(), 400);
|
|
||||||
assert!(matches!(
|
|
||||||
checked_chunk_byte_len(&[usize::MAX, 2], 4),
|
|
||||||
Err(FormatError::Overflow(_))
|
|
||||||
));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn unallocatable_output_is_an_error_not_an_abort() {
|
|
||||||
assert_eq!(alloc_output(16).unwrap(), vec![0u8; 16]);
|
|
||||||
assert!(matches!(
|
|
||||||
alloc_output(usize::MAX / 2),
|
|
||||||
Err(FormatError::Overflow(_))
|
|
||||||
));
|
|
||||||
}
|
|
||||||
|
|
||||||
fn write_offset(buf: &mut Vec<u8>, val: u64, size: u8) {
|
fn write_offset(buf: &mut Vec<u8>, val: u64, size: u8) {
|
||||||
match size {
|
match size {
|
||||||
4 => buf.extend_from_slice(&(val as u32).to_le_bytes()),
|
4 => buf.extend_from_slice(&(val as u32).to_le_bytes()),
|
||||||
@@ -1789,9 +1657,9 @@ mod tests {
|
|||||||
let chunk_bytes = chunk_size_elems * elem_size; // full chunk allocation
|
let chunk_bytes = chunk_size_elems * elem_size; // full chunk allocation
|
||||||
|
|
||||||
// Write chunk data (full chunk size, padding with zeros)
|
// Write chunk data (full chunk size, padding with zeros)
|
||||||
for (i, value) in values.iter().enumerate().take(end).skip(start) {
|
for i in start..end {
|
||||||
let byte_offset = data_offset + (i - start) * elem_size;
|
let byte_offset = data_offset + (i - start) * elem_size;
|
||||||
file_data[byte_offset..byte_offset + 8].copy_from_slice(&value.to_le_bytes());
|
file_data[byte_offset..byte_offset + 8].copy_from_slice(&values[i].to_le_bytes());
|
||||||
}
|
}
|
||||||
|
|
||||||
chunk_infos.push(ChunkInfo {
|
chunk_infos.push(ChunkInfo {
|
||||||
@@ -1969,8 +1837,8 @@ mod tests {
|
|||||||
for chunk_idx in 0..2 {
|
for chunk_idx in 0..2 {
|
||||||
let start = chunk_idx * chunk_elems;
|
let start = chunk_idx * chunk_elems;
|
||||||
let mut chunk_bytes = Vec::new();
|
let mut chunk_bytes = Vec::new();
|
||||||
for value in values.iter().skip(start).take(chunk_elems) {
|
for i in start..start + chunk_elems {
|
||||||
chunk_bytes.extend_from_slice(&value.to_le_bytes());
|
chunk_bytes.extend_from_slice(&values[i].to_le_bytes());
|
||||||
}
|
}
|
||||||
let compressed = compress_chunk(&chunk_bytes, &pipeline, elem_size as u32).unwrap();
|
let compressed = compress_chunk(&chunk_bytes, &pipeline, elem_size as u32).unwrap();
|
||||||
|
|
||||||
|
|||||||
@@ -143,6 +143,10 @@ pub fn parse_vds_mappings(
|
|||||||
let source_selection = read_selection(heap_data, &mut pos)?;
|
let source_selection = read_selection(heap_data, &mut pos)?;
|
||||||
let virtual_selection = read_selection(heap_data, &mut pos)?;
|
let virtual_selection = read_selection(heap_data, &mut pos)?;
|
||||||
|
|
||||||
|
// Validate external file name to prevent directory traversal attacks
|
||||||
|
// (Dataset paths within files can use absolute HDF5 paths like "/data")
|
||||||
|
validate_vds_file_name(&source_file)?;
|
||||||
|
|
||||||
mappings.push(VdsMapping {
|
mappings.push(VdsMapping {
|
||||||
source_file,
|
source_file,
|
||||||
source_dataset,
|
source_dataset,
|
||||||
@@ -154,6 +158,37 @@ pub fn parse_vds_mappings(
|
|||||||
Ok(mappings)
|
Ok(mappings)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Validate external file names to prevent directory traversal.
|
||||||
|
/// Dataset paths within files can use absolute HDF5 paths (starting with /),
|
||||||
|
/// but external file names must not escape the file tree via .. or absolute paths.
|
||||||
|
fn validate_vds_file_name(filename: &str) -> Result<(), FormatError> {
|
||||||
|
if filename.is_empty() {
|
||||||
|
return Ok(());
|
||||||
|
}
|
||||||
|
|
||||||
|
// "." means same file - always OK
|
||||||
|
if filename == "." {
|
||||||
|
return Ok(());
|
||||||
|
}
|
||||||
|
|
||||||
|
// Filesystem paths cannot start with / (absolute filesystem path)
|
||||||
|
if filename.starts_with('/') {
|
||||||
|
return Err(FormatError::FilterError(
|
||||||
|
"VDS file name cannot be an absolute filesystem path".into(),
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reject directory traversal (..)
|
||||||
|
if filename.contains("..") {
|
||||||
|
return Err(FormatError::FilterError(
|
||||||
|
"VDS file name contains illegal traversal sequence (..)".into(),
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Relative filesystem paths are OK
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
/// Read a null-terminated UTF-8 string from data starting at `pos`.
|
/// Read a null-terminated UTF-8 string from data starting at `pos`.
|
||||||
fn read_null_terminated_string(data: &[u8], pos: &mut usize) -> Result<String, FormatError> {
|
fn read_null_terminated_string(data: &[u8], pos: &mut usize) -> Result<String, FormatError> {
|
||||||
let start = *pos;
|
let start = *pos;
|
||||||
@@ -862,4 +897,68 @@ mod tests {
|
|||||||
let blob = [0x01u8, 0, 0, 0, 0, 0, 0, 0, 0];
|
let blob = [0x01u8, 0, 0, 0, 0, 0, 0, 0, 0];
|
||||||
assert!(parse_vds_mappings(&blob, 8).unwrap().is_empty());
|
assert!(parse_vds_mappings(&blob, 8).unwrap().is_empty());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn parse_vds_mappings_rejects_path_traversal() {
|
||||||
|
// INT-06: Verify that VDS file names containing ".." are rejected
|
||||||
|
let blob = [
|
||||||
|
0x00u8, // version 0 (with explicit file name)
|
||||||
|
0x01, 0, 0, 0, 0, 0, 0, 0, // nused = 1
|
||||||
|
0x2e, 0x2e, 0x2f, 0x65, 0x74, 0x63, 0x2f, 0x70, 0x61, 0x73, 0x73, 0x77, 0x64, 0x00, // "../etc/passwd | ||||||