tidaldb/tidal/src/entities/multi_preference.rs
jx12n 6a937fc4bc feat(m12): multi-vector user preference modeling + ANN candidate-gen
Add multi-vector preference entity (per-signal-type preference vectors with
event-time decay) feeding ANN candidate generation in the query executor.

- entities: multi_preference vectors + event-time-aware preference updates
- query/executor: ANN candidate-gen + personalization/pipeline integration
- storage/keys, db ops, state_rebuild: persist & rebuild multi-vector prefs
- ranking: profile + builtins support for multi-vector scoring
- tidal-server/config: expose multi-preference knobs
- tests/bench: m12_preference_event_time integration + multi_preference bench
- docs: multi-vector-preference research, ROADMAP/ARCHITECTURE refresh,
  legal/tidaldb-patent-proposal
- .codex/agents: codex agent definitions
- chore: gitignore tool-regenerated .agents/ mirror (doc-guard rejects it)
2026-06-23 09:52:36 -06:00

1780 lines
77 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

//! Multi-vector (PinnerSage-style) per-user preference.
//!
//! Online sequential k-means with a DP-means threshold split, forward-decayed
//! cluster importance, and a cold-start fall-back to the single adaptive-LR
//! vector.
//!
//! This is the **Approach B** of `docs/research/multi-vector-preference.md`
//! (settled design). A user who engages with hiking, cooking, and cars no longer
//! collapses to one EMA centroid that represents none of them — each coherent
//! interest gets its own cluster centroid, and query time fans out across the
//! top-M clusters by current (decayed) importance.
//!
//! # Tiers
//!
//! - **Cold start (`interaction_count < COLD_START_N`).** The tier gate is the
//! per-user *total* interaction count (`interaction_counts`), distinct from a
//! cluster's per-cluster adaptive-LR `update_count`. Below the threshold the
//! user has too few positive interactions for clustering to be meaningful, so
//! updates flow into the *existing* [`PreferenceVectors`] single adaptive-LR
//! vector unchanged (the documented `ann_for_tidaldb.md:112` cold-start tier,
//! K=1 floor). On crossing the threshold the single vector seeds the first
//! cluster — cold start is the natural K=1 case, not a separate algorithm.
//! - **Warm (multi-cluster).** Each positive engagement is assigned to the
//! nearest existing centroid by cosine; if the best cosine is `>= split_threshold`
//! (τ) the embedding is blended into that centroid via the per-cluster adaptive
//! EMA (`alpha = base / (1 + ln_1p(count))`, preserved exactly per cluster),
//! else a new centroid is opened up to `K_max`. Over the cap, the embedding is
//! assigned to the nearest centroid — never evicted (eviction loses an
//! interest; merging is the deferred Approach-C recluster's job).
//!
//! # Cluster importance (forward-decay composition)
//!
//! Each cluster carries `(importance_at_anchor, anchor_ts)` and is forward-decayed
//! with tidalDB's canonical [`forward_decay_step`] kernel on read — the exact O(1)
//! primitive the signal ledger uses. A stale interest naturally falls out of the
//! queried top-M as its decayed importance drops below the third-ranked cluster,
//! without ever being deleted (a re-engagement re-boosts it).
//!
//! # Persistence
//!
//! Checkpoint/restore live under [`Tag::Preference`] (no new tag — keeps the
//! single prefix-scan restore). The value is **version-tagged** and
//! **backward-compatible**: a legacy single-vector row restores as a cold-start
//! single vector (a zero-migration upgrade). Discrimination is structural — see
//! [`FORMAT_VERSION`] — so a legacy row whose `update_count` low byte aliases the
//! sentinel is still decoded correctly, not dropped.
//!
//! The per-cluster value layout already carries the `anchor_ts` /
//! `importance_at_anchor` decay state that Approach C (the deferred periodic
//! medoid recluster) resets, so that field-level change needs no migration.
//! Approach C *also* needs a second persisted structure — a bounded per-user
//! interaction-embedding window — whose tag ([`Tag::PreferenceWindow`]) is
//! reserved but not yet populated; adding that window is an additive tag, never a
//! rewrite of existing `Tag::Preference` rows.
//!
//! Every centroid is re-normalized + NaN-neutralized at the load boundary, and a
//! torn cluster is dropped individually rather than failing the whole user —
//! matching the single-vector `restore` contract.
use dashmap::DashMap;
use crate::signals::decay::forward_decay_step;
pub use crate::entities::preference::PreferenceVectors;
/// Cold-start interaction threshold N.
///
/// Below this many positive interactions a user stays on the single-vector
/// cold-start tier (`ann_for_tidaldb.md:112`, `multi-vector-preference.md` §2).
/// The single vector *is* the `K=1` case; on crossing this it seeds cluster 0.
pub const COLD_START_N: u64 = 5;
/// Maximum interest clusters per user.
///
/// `ARCHITECTURE.md`'s "3-10" (`multi-vector-preference.md` §2). Over the cap the
/// split rule assigns to the nearest centroid rather than opening an 11th — the
/// DP-means cap behavior.
pub const K_MAX: usize = 10;
/// Default DP-means split threshold (cosine). A positive engagement whose best
/// cosine to an existing centroid is `< τ` opens a new cluster (up to `K_MAX`).
///
/// This is the single most important constant and is a *starting point* for the
/// grid search, not a tuned shipped value (`multi-vector-preference.md` §1, Open
/// Questions): ~0.50.6 for OpenAI-1536-D embeddings. 0.55 sits in that band.
pub const DEFAULT_SPLIT_THRESHOLD: f32 = 0.55;
/// Default number of clusters queried at serve time (`M`).
///
/// `min(K_active, 3)` — 3 is `PinnerSage`'s serve-time count
/// (`multi-vector-preference.md` §3).
pub const DEFAULT_TOP_M: usize = 3;
/// Default importance half-life (seconds) for the forward-decay composition.
///
/// 30 days — a stale interest fades out of the top-`M` over weeks, not minutes;
/// independent of the per-signal-type decay so a slow-burning taste persists
/// even as individual view scores decay fast. (`multi-vector-preference.md` §3 /
/// Open Questions: tie-to-signal-vs-independent is a follow-up A/B.)
pub const DEFAULT_IMPORTANCE_HALF_LIFE_SECS: f64 = 30.0 * 24.0 * 3600.0;
/// The persistence format version byte written at the head of every new
/// multi-cluster row.
///
/// **Discrimination is structural, not byte-value-based.** A legacy single-vector
/// row is `[count:8 LE][dim:4 LE][f32*dim]`, so its first byte is the low byte of
/// the little-endian `update_count` — which can equal *any* value, including this
/// sentinel (e.g. a user with `update_count == 2`). The first byte is therefore
/// only a cheap hint: [`MultiPreferenceVectors::restore`] commits a row to the
/// warm tier **only if** the first byte matches AND [`decode_multi_value`] parses
/// it into ≥1 dimension-correct cluster; otherwise it falls through to the legacy
/// decoder. A real multi row always parses (dim-correct, ≥1 cluster); a colliding
/// legacy row fails the parse (its `dim` field, read from the count bytes, never
/// matches the store dim) and is correctly rescued as legacy. This keeps the
/// zero-migration upgrade path sound for *every* legacy `update_count`, not just
/// counts whose low byte avoids the sentinel.
const FORMAT_VERSION: u8 = 0x02;
/// Multi-cluster checkpoint header length: `[version:1][n_clusters:1][dim:4 LE]`.
/// Single-sourced so `encode_multi_value`/`decode_multi_value` cannot drift.
const MULTI_HEADER_LEN: usize = 6;
/// Per-cluster fixed prefix length before the centroid floats:
/// `[update_count:8 LE][importance_at_anchor:4 LE][anchor_ts:8 LE]`. The centroid
/// (`dim * 4` bytes) follows. Single-sourced for the same drift reason; when
/// Approach C widens the per-cluster record this constant moves in one place.
const CLUSTER_FIXED_LEN: usize = 8 + 4 + 8;
/// One interest cluster: a centroid plus the per-cluster adaptive-LR state and
/// the forward-decay importance anchor.
///
/// The layout is intentionally a superset of what Approach B needs: `anchor_ts`
/// and `importance_at_anchor` compose with the decay kernel today, and the same
/// fields carry the medoid-recluster bookkeeping Approach C will reset — so *this
/// per-cluster record* needs no format migration when C ships. (C's separate
/// interaction-embedding window is an additive [`Tag::PreferenceWindow`] row, not
/// a change to this layout — see the module-level Persistence note and
/// `multi-vector-preference.md` §5/§6.)
#[derive(Debug, Clone)]
struct Cluster {
/// Unit-normalized centroid (the cluster's query/representative vector).
centroid: Vec<f32>,
/// Per-cluster adaptive-LR update count. Preserved exactly as the
/// single-vector tier: `alpha = base / (1 + ln_1p(count))`. Approach C's
/// periodic recluster is the only thing that resets this.
update_count: u64,
/// Engagement mass at `anchor_ts`, before forward-decay.
importance_at_anchor: f32,
/// Nanosecond timestamp the importance anchor was last advanced.
anchor_ts_ns: u64,
}
impl Cluster {
/// Current (forward-decayed) importance at `now_ns`, via the canonical decay
/// kernel — decay the anchored mass forward, add zero new weight.
fn current_importance(&self, now_ns: u64, lambda: f64) -> f64 {
forward_decay_step(
f64::from(self.importance_at_anchor),
self.anchor_ts_ns,
now_ns,
lambda,
0.0,
)
.new_score
}
/// Fold a new engagement's mass into the importance anchor at `event_ns`.
///
/// Uses the same forward-decay kernel as the signal ledger: in-order events
/// decay the prior mass forward then add the new weight; out-of-order events
/// pre-decay the weight and fold it in without regressing the anchor.
fn add_importance(&mut self, weight: f64, event_ns: u64, lambda: f64) {
let step = forward_decay_step(
f64::from(self.importance_at_anchor),
self.anchor_ts_ns,
event_ns,
lambda,
weight,
);
#[allow(clippy::cast_possible_truncation)]
{
self.importance_at_anchor = step.new_score as f32;
}
if step.advance_timestamp {
self.anchor_ts_ns = event_ns;
}
}
}
/// Per-user multi-vector preference store.
///
/// Holds the warm-tier cluster sets and **owns** the single-vector cold-start
/// store, so it is a drop-in superset of [`PreferenceVectors`]: every method the
/// rest of the engine called on the single store has the same signature here and
/// routes to whichever tier the user is in, plus the new fan-out methods
/// ([`query_vectors`](Self::query_vectors), [`cosine_similarity`](Self::cosine_similarity)).
///
/// Thread-safe via `DashMap`; concurrent updates to different users never contend.
pub struct MultiPreferenceVectors {
/// `user_id` -> active clusters (warm tier). A user appears here only once
/// they have crossed [`COLD_START_N`]; cold-start users live entirely in
/// `cold_start`.
clusters: DashMap<u64, Vec<Cluster>>,
/// The single adaptive-LR vector for cold-start users — the existing,
/// unchanged K=1 tier. Also the source the first cluster is seeded from on
/// crossover.
cold_start: PreferenceVectors,
/// Per-user positive-interaction count, gating the cold-start → warm
/// transition. Distinct from a cluster's `update_count` (which is per-cluster
/// adaptive LR): this counts *total* engagements to decide the tier.
interaction_counts: DashMap<u64, u64>,
/// Embedding dimensionality. Every centroid must match.
dim: usize,
/// Base learning rate for the per-cluster adaptive EMA.
base_lr: f32,
/// DP-means split threshold τ (cosine).
split_threshold: f32,
/// Cluster cap.
k_max: usize,
/// Cold-start interaction threshold N.
cold_start_n: u64,
/// Forward-decay rate for cluster importance (`ln 2 / half_life_secs`).
importance_lambda: f64,
}
impl MultiPreferenceVectors {
/// Create a multi-vector store for `dim`-dimensional embeddings with default
/// tuning constants.
#[must_use]
pub fn new(dim: usize) -> Self {
Self::with_params(
dim,
0.1,
DEFAULT_SPLIT_THRESHOLD,
K_MAX,
COLD_START_N,
DEFAULT_IMPORTANCE_HALF_LIFE_SECS,
)
}
/// Create with a custom base learning rate, default everything else.
/// Kept signature-compatible with [`PreferenceVectors::with_learning_rate`].
#[must_use]
pub fn with_learning_rate(dim: usize, learning_rate: f32) -> Self {
Self::with_params(
dim,
learning_rate,
DEFAULT_SPLIT_THRESHOLD,
K_MAX,
COLD_START_N,
DEFAULT_IMPORTANCE_HALF_LIFE_SECS,
)
}
/// Full constructor for benchmarks / tuning sweeps.
///
/// TODO(tuning-config): every production construction site (`db/open.rs`,
/// `db/mod.rs`) currently calls [`new`](Self::new), so `split_threshold`,
/// `k_max`, `cold_start_n`, and the importance half-life ship as recompile-only
/// constants — this constructor is reachable only from tests/benches. The
/// research doc mandates a per-corpus grid search of these (esp. τ); thread
/// them through the schema/`TidalDbBuilder` + server config the way
/// `top_clusters` already is, then route the open-time sites here. Tracked in
/// `docs/planning/ROADMAP.md` (multi-vector follow-ups). The shipped defaults
/// sit inside the doc's stated 0.50.6 band for OpenAI-1536-D embeddings.
#[must_use]
pub fn with_params(
dim: usize,
base_lr: f32,
split_threshold: f32,
k_max: usize,
cold_start_n: u64,
importance_half_life_secs: f64,
) -> Self {
Self {
clusters: DashMap::new(),
cold_start: PreferenceVectors::with_learning_rate(dim, base_lr),
interaction_counts: DashMap::new(),
dim,
base_lr,
split_threshold,
k_max: k_max.max(1),
cold_start_n,
importance_lambda: std::f64::consts::LN_2 / importance_half_life_secs.max(f64::EPSILON),
}
}
// ── Single-vector-compatible surface (drop-in for PreferenceVectors) ──────
/// The user's primary query vector: the **top-importance** cluster centroid
/// for a warm user, else the cold-start single vector. `None` if neither tier
/// has a vector for this user.
///
/// This keeps callers that resolve exactly one query vector (e.g. the SEARCH
/// executor's single-vector boost, and the m12p2 `for_you` ANN resolution)
/// working unchanged while the richer fan-out is opt-in via
/// [`query_vectors`](Self::query_vectors).
#[must_use]
pub fn get(&self, user_id: u64) -> Option<Vec<f32>> {
if let Some(clusters) = self.clusters.get(&user_id) {
return self
.top_importance_cluster(&clusters, now_ns())
.map(|c| c.centroid.clone());
}
self.cold_start.get(user_id)
}
/// Set the cold-start vector directly (cold-start initialization). Returns
/// `false` on a dimension mismatch **or if the user is already on the warm
/// (clustered) tier**. Does not create clusters — a subsequent `update` past
/// the threshold seeds clusters from this vector.
///
/// `set` is cold-start initialization only. Writing `cold_start` behind a
/// warm user's live clusters would put them in both tiers (breaking the
/// disjointness invariant `len()`/`checkpoint` rely on) and a later
/// checkpoint would clobber the warm row; rejecting it keeps the tiers
/// disjoint, mirroring the dimension-mismatch contract (rejected, no
/// mutation).
///
/// The per-user `interaction_counts` guard is held across the warm check and
/// the cold-start write so `set` serializes against a concurrent `update_at`
/// crossover for the same user — otherwise the check could observe "not warm",
/// a crossover could complete and `cold_start.remove`, and the write would then
/// resurrect the cold row, landing the user in both tiers (the same race
/// `update_at` is structured to prevent).
#[must_use]
#[allow(clippy::significant_drop_tightening)]
pub fn set(&self, user_id: u64, vec: Vec<f32>) -> bool {
let _guard = self.interaction_counts.entry(user_id).or_insert(0);
if self.clusters.contains_key(&user_id) {
return false;
}
self.cold_start.set(user_id, vec)
}
/// Record a positive engagement: route to the cold-start tier below the
/// threshold, else assign to / split clusters in the warm tier.
///
/// Returns `false` on a dimension mismatch (no state mutated).
#[must_use]
pub fn update(&self, user_id: u64, interaction_embedding: &[f32]) -> bool {
self.update_at(user_id, interaction_embedding, now_ns())
}
/// [`update`](Self::update) at an explicit event timestamp (deterministic
/// tests + importance anchoring). `now_ns` anchors the cluster importance.
//
// Concurrency: the per-user `interaction_counts` entry guard is held across
// the ENTIRE tier transition (count bump → cold/warm decision → cluster seed
// → cold-start removal) so it is atomic for this user. Without it, two
// concurrent same-user updates straddling the `COLD_START_N` boundary could
// interleave such that a cold-path `cold_start.update` lands AFTER a
// warm-path `cold_start.remove`, leaving the user in BOTH tiers — which a
// later checkpoint would materialize as two rows under one key, clobbering
// the warm row and demoting the user to cold-start on restore. The lock
// order is always `interaction_counts → clusters → cold_start` (never
// inverted by any other method), so holding this outermost guard is
// deadlock-free. Inside, the `clusters` entry guard is still scoped to its
// block so no `clusters` guard is held across the `cold_start` accesses.
#[must_use]
#[allow(clippy::significant_drop_tightening)]
pub fn update_at(&self, user_id: u64, interaction_embedding: &[f32], now_ns: u64) -> bool {
if interaction_embedding.len() != self.dim {
return false;
}
// Bump the total interaction count and hold the guard for the whole
// transition (see the concurrency note above). `count_after` decides the
// tier; the guard makes the decision-and-mutation atomic per user.
let mut count_guard = self.interaction_counts.entry(user_id).or_insert(0);
let count_after = *count_guard + 1;
*count_guard = count_after;
if count_after < self.cold_start_n && !self.clusters.contains_key(&user_id) {
// Cold-start tier: the existing single adaptive-LR vector, unchanged.
// Held `count_guard` blocks a concurrent same-user warm crossover from
// removing the cold row out from under this insert.
return self.cold_start.update(user_id, interaction_embedding);
}
// Warm tier. On the very first crossover, migrate the cold-start vector
// into cluster 0 so we do not discard the taste learned cold. Read the
// cold-start seed BEFORE taking the cluster guard so we never hold the
// `clusters` entry guard across another DashMap access.
let mut normalized = interaction_embedding.to_vec();
l2_normalize(&mut normalized);
let seed = (!self.clusters.contains_key(&user_id))
.then(|| {
self.cold_start
.get(user_id)
.map(|v| (v, self.cold_start.update_count(user_id)))
})
.flatten();
// Scope the `clusters` entry guard so it is released before we touch the
// cold-start store (a different DashMap) below.
{
let mut entry = self.clusters.entry(user_id).or_default();
let clusters = entry.value_mut();
if clusters.is_empty()
&& let Some((centroid, update_count)) = seed
{
// Seed cluster 0 from the cold-start vector. Carry the cold-start
// update count so the adaptive LR resumes where it was rather than
// snapping back to the base rate on crossover, and seed its
// importance from the cold-start engagement MASS (≥1) so the
// dominant cold taste is NOT invisible to the top-M fan-out when a
// divergent crossover interest immediately opens another cluster.
clusters.push(Cluster {
centroid,
update_count,
importance_at_anchor: importance_seed(update_count),
anchor_ts_ns: now_ns,
});
}
self.assign_or_split(clusters, &normalized, now_ns);
}
// The cold-start vector is now superseded by the clusters; drop it so a
// checkpoint does not double-store a stale K=1 alongside the clusters.
// Still under `count_guard`, so no concurrent cold-path update for this
// user can re-insert it after this point.
self.cold_start.remove(user_id);
true
}
/// Blend `interaction_embedding` at an explicit learning rate, bypassing the
/// adaptive count. Routes to whichever tier the user is in (clusters: into
/// the nearest centroid; cold start: the single vector). Signature-compatible
/// with [`PreferenceVectors::update_with_custom_rate`].
///
/// Returns `false` on a dimension mismatch.
#[must_use]
pub fn update_with_custom_rate(
&self,
user_id: u64,
interaction_embedding: &[f32],
lr: f32,
) -> bool {
if interaction_embedding.len() != self.dim {
return false;
}
if let Some(mut clusters) = self.clusters.get_mut(&user_id) {
let mut normalized = interaction_embedding.to_vec();
l2_normalize(&mut normalized);
if let Some(idx) = nearest_cluster(clusters.value(), &normalized) {
blend_into(&mut clusters.value_mut()[idx].centroid, &normalized, lr);
}
return true;
}
self.cold_start
.update_with_custom_rate(user_id, interaction_embedding, lr)
}
/// Cosine similarity between the user's preference and a candidate embedding.
///
/// For a warm user this is the **max cosine over the user's clusters**
/// (`multi-vector-preference.md` §4: a candidate retrieved via the "cars"
/// cluster must not be scored against a "cooking"-dominated mean). For a
/// cold-start user it is the single-vector cosine. `None` if the user has no
/// preference or dimensions mismatch.
#[must_use]
pub fn cosine_similarity(&self, user_id: u64, candidate: &[f32]) -> Option<f32> {
if candidate.len() != self.dim {
return None;
}
if let Some(clusters) = self.clusters.get(&user_id) {
return nearest_cosine_in(&clusters, candidate);
}
self.cold_start.cosine_similarity(user_id, candidate)
}
// ── Multi-vector fan-out surface ──────────────────────────────────────────
/// The top-M cluster centroids by current (forward-decayed) importance, for
/// the query-time fan-out. `M = min(K_active, top_m)`.
///
/// Deterministic (top-M by importance, descending; ties broken by centroid
/// bytes) so the same user/time yields the same fan-out — tidalDB values
/// reproducible queries. For a cold-start user, returns the single vector as
/// a one-element fan-out so the caller has a uniform interface. Empty when the
/// user has no preference at all.
#[must_use]
pub fn query_vectors(&self, user_id: u64, now_ns: u64, top_m: usize) -> Vec<Vec<f32>> {
let m = top_m.max(1);
if let Some(clusters) = self.clusters.get(&user_id) {
let mut ranked: Vec<(f64, &Cluster)> = clusters
.iter()
.map(|c| (c.current_importance(now_ns, self.importance_lambda), c))
.collect();
// Descending importance; deterministic tie-break on centroid bytes.
ranked.sort_by(|a, b| {
b.0.partial_cmp(&a.0)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| centroid_cmp(&a.1.centroid, &b.1.centroid))
});
return ranked
.into_iter()
.take(m)
.map(|(_, c)| c.centroid.clone())
.collect();
}
self.cold_start
.get(user_id)
.map(|v| vec![v])
.unwrap_or_default()
}
/// Number of active clusters for a user (0 for a cold-start or unknown user).
#[must_use]
pub fn cluster_count(&self, user_id: u64) -> usize {
self.clusters.get(&user_id).map_or(0, |c| c.len())
}
/// Whether the user is on the warm (multi-cluster) tier.
#[must_use]
pub fn is_warm(&self, user_id: u64) -> bool {
self.clusters.contains_key(&user_id)
}
/// Whether the user has any stored preference (either tier), without
/// allocating. Equivalent to `get(user_id).is_some()` but a pair of O(1)
/// hash lookups — for serve-path existence guards that would otherwise clone
/// a full centroid just to discard it.
#[must_use]
pub fn contains(&self, user_id: u64) -> bool {
self.is_warm(user_id) || self.cold_start.contains(user_id)
}
/// Current decayed importance of each cluster, for tests / introspection.
#[must_use]
pub fn cluster_importances(&self, user_id: u64, now_ns: u64) -> Vec<f64> {
self.clusters.get(&user_id).map_or_else(Vec::new, |c| {
c.iter()
.map(|cl| cl.current_importance(now_ns, self.importance_lambda))
.collect()
})
}
// ── Bookkeeping passthroughs ──────────────────────────────────────────────
/// Remove all per-user state (both tiers) — returns the store to its
/// just-opened state so a rebuild does not fold on already-folded vectors.
pub fn clear(&self) {
self.clusters.clear();
self.interaction_counts.clear();
self.cold_start.clear();
}
/// Number of users with any stored preference (either tier).
#[must_use]
pub fn len(&self) -> usize {
// A warm user has no cold-start row (removed on crossover), so the two
// sets are disjoint and the sum is the true distinct-user count.
self.clusters.len() + self.cold_start.len()
}
/// Whether no preferences are stored at all.
#[must_use]
pub fn is_empty(&self) -> bool {
self.clusters.is_empty() && self.cold_start.is_empty()
}
/// Total positive interactions recorded for a user (the tier gate).
#[must_use]
pub fn interaction_count(&self, user_id: u64) -> u64 {
self.interaction_counts.get(&user_id).map_or(0, |c| *c)
}
// ── Internal clustering ───────────────────────────────────────────────────
/// DP-means assign-or-split for one engagement (already unit-normalized).
fn assign_or_split(&self, clusters: &mut Vec<Cluster>, embedding: &[f32], now_ns: u64) {
let nearest = nearest_cluster_cos(clusters, embedding);
let best_cos = nearest.map_or(f32::NEG_INFINITY, |(_, cos)| cos);
let open_new = best_cos < self.split_threshold && clusters.len() < self.k_max;
let Some((idx, _)) = nearest.filter(|_| !open_new) else {
// Empty cluster set, or a distinct interest under the cap: open a new
// cluster seeded at importance 1.0.
clusters.push(Cluster {
centroid: embedding.to_vec(),
update_count: 0,
importance_at_anchor: 0.0,
anchor_ts_ns: now_ns,
});
let last = clusters.len() - 1;
clusters[last].add_importance(1.0, now_ns, self.importance_lambda);
return;
};
// Assign to the nearest centroid (also the over-cap path): blend via the
// per-cluster adaptive LR, then bump importance.
let lr = adaptive_lr(self.base_lr, clusters[idx].update_count);
blend_into(&mut clusters[idx].centroid, embedding, lr);
clusters[idx].update_count += 1;
clusters[idx].add_importance(1.0, now_ns, self.importance_lambda);
}
/// The current top-importance cluster (does not clone).
///
/// MUST select the same cluster as `query_vectors(.., 1)[0]` so `get()` and the
/// fan-out's primary agree. `query_vectors` sorts importance-descending and
/// takes the FIRST element, breaking importance ties by ascending centroid
/// bytes (`centroid_cmp(a, b)`). `max_by` returns the LAST maximal element, so
/// to pick the same cluster on a tie the tie-break here is REVERSED
/// (`centroid_cmp(b, a)`) — last-of-equal under the reversed order is the
/// ascending-bytes winner. The `get()==query_vectors(..,1)[0]` invariant is
/// pinned by `get_matches_query_vectors_primary_on_ties`.
fn top_importance_cluster<'c>(
&self,
clusters: &'c [Cluster],
now_ns: u64,
) -> Option<&'c Cluster> {
clusters.iter().max_by(|a, b| {
a.current_importance(now_ns, self.importance_lambda)
.partial_cmp(&b.current_importance(now_ns, self.importance_lambda))
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| centroid_cmp(&b.centroid, &a.centroid))
})
}
}
// ── Checkpoint / Restore ──────────────────────────────────────────────────────
impl MultiPreferenceVectors {
/// Checkpoint every user's preference (both tiers) under `Tag::Preference`,
/// one row per user, atomically swapped in a single `WriteBatch`.
///
/// Warm users are written in the **version-tagged multi-cluster** layout;
/// cold-start users are written in the **legacy single-vector** layout so an
/// older binary (or the standalone single-vector path) can still read them.
/// Key suffix `[user: 8B BE]`, sentinel entity id 0 — identical to the
/// single-vector checkpoint so the single prefix-scan restore is preserved.
///
/// # Errors
///
/// Returns storage errors from the underlying engine.
pub fn checkpoint(&self, storage: &dyn crate::storage::StorageEngine) -> crate::Result<()> {
use crate::{
schema::EntityId,
storage::{Tag, WriteBatch, encode_key, entity_tag_prefix},
};
let prefix = entity_tag_prefix(EntityId::new(0), Tag::Preference);
let mut batch = WriteBatch::with_capacity(self.clusters.len() + self.cold_start.len() + 1);
// Stage deletion of every existing row so the post-swap snapshot holds
// exactly the current set (a crash mid-checkpoint never leaves a torn mix).
for item in storage.scan_prefix(&prefix) {
let (key, _) = item.map_err(crate::schema::TidalError::from)?;
batch.delete(key);
}
// Warm users: multi-cluster rows.
for entry in &self.clusters {
let user_id = *entry.key();
let value = encode_multi_value(entry.value());
let key = encode_key(EntityId::new(0), Tag::Preference, &user_id.to_be_bytes());
batch.put(key, value);
}
// Cold-start users: legacy single-vector rows (so the format stays a
// strict superset and a downgrade can still read them). The single store
// writes its own legacy rows via its checkpoint, but we cannot call that
// here without double-deleting; re-encode inline to one batch instead.
//
// Defense-in-depth: SKIP any user already present in `clusters`. The
// tiers are kept disjoint by `update_at`/`set`, but a legacy cold row and
// a warm cluster row share the same `Tag::Preference` key, so if a user
// ever appeared in both maps the later legacy `put` would clobber the warm
// row (last-write-wins) and silently demote them on restore. Skipping
// makes that impossible regardless of how the maps got there.
//
// Materialize the warm key-set FIRST, so the `skip` closure does a plain
// `HashSet` lookup rather than taking a `clusters` guard while
// `append_legacy_rows` holds a `cold_start` shard guard — which would
// invert the canonical `clusters → cold_start` lock order (`update_at`).
let warm_keys: std::collections::HashSet<u64> =
self.clusters.iter().map(|e| *e.key()).collect();
self.cold_start
.append_legacy_rows(&mut batch, |uid| warm_keys.contains(&uid));
storage
.write_batch(batch)
.map_err(crate::schema::TidalError::from)?;
Ok(())
}
/// Restore from a `Tag::Preference` checkpoint. Reads both the new
/// multi-cluster layout and legacy single-vector rows (loaded as one K=1
/// cluster / a cold-start vector respectively). Skips rows whose dimension
/// does not match, and drops torn clusters individually rather than failing
/// the user. Every restored centroid is NaN-neutralized + re-normalized at the
/// load boundary so the cosine invariant holds even for a tampered row.
///
/// # Errors
///
/// Returns storage errors from the underlying engine.
pub fn restore(&self, storage: &dyn crate::storage::StorageEngine) -> crate::Result<()> {
use crate::{
schema::EntityId,
storage::{Tag, entity_tag_prefix, parse_key},
};
let prefix = entity_tag_prefix(EntityId::new(0), Tag::Preference);
let mut warm = 0u64;
let mut cold = 0u64;
for entry in storage.scan_prefix(&prefix) {
let (key, value) = entry.map_err(crate::schema::TidalError::from)?;
let Some((_, Tag::Preference, suffix)) = parse_key(&key) else {
continue;
};
if suffix.len() < 8 {
continue;
}
let user_id = u64::from_be_bytes(suffix[0..8].try_into().unwrap_or([0u8; 8]));
// Multi-cluster row: the first byte is only a hint (it can collide
// with a legacy count's low byte — see `FORMAT_VERSION`), so commit to
// the warm tier ONLY if the value actually parses as ≥1 dimension-
// correct cluster. A colliding legacy row fails this parse and falls
// through to the legacy decoder below — never silently dropped.
if value.first() == Some(&FORMAT_VERSION)
&& let Some(clusters) = decode_multi_value(&value, self.dim)
&& !clusters.is_empty()
{
let total: u64 = clusters.iter().map(|c| c.update_count).sum();
// The interaction-count gate: a warm user is, by construction,
// already past the threshold. The exact original count is not
// recoverable from the checkpoint, so compute a conservative
// over-estimate (`sum(update_count) + cluster_count`) and floor it
// at `cold_start_n`. The FLOOR (not the arithmetic) is what
// guarantees a restored warm user is never demoted on the next
// update; the estimate is loose because cluster 0's update_count
// already absorbs the seeded cold-start blends.
let restored_interactions = total + clusters.len() as u64;
self.clusters.insert(user_id, clusters);
self.interaction_counts
.insert(user_id, restored_interactions.max(self.cold_start_n));
warm += 1;
continue;
}
// First byte matched `FORMAT_VERSION` but it did not load as a
// dim-correct multi row. It is EITHER a genuine multi row for a
// DIFFERENT store dim (a schema embedding-dim change → must be skipped,
// exactly as before the collision fix) OR a legacy row whose
// `update_count` low byte merely aliases the sentinel (must be rescued
// as legacy below). Distinguish them structurally: a genuine multi row's
// length matches its OWN embedded `[n, dim]` header exactly; a colliding
// legacy row never does. Skipping the former prevents loading a garbage
// cold-start vector from a multi row's bytes.
if is_well_formed_multi_header(&value) {
continue;
}
// Legacy single-vector row (also the rescue path for a legacy row
// whose `update_count` low byte happened to equal `FORMAT_VERSION`).
// Decoded via the single-sourced `decode_legacy_row` so the legacy
// format has exactly one encoder/decoder pair across both stores.
if let Some((update_count, vec)) =
crate::entities::preference::decode_legacy_row(&value, self.dim)
{
self.cold_start.insert_restored(user_id, vec, update_count);
// A legacy row is a cold-start (K=1) user; seed the interaction
// gate from its update count so it stays cold until it organically
// crosses the threshold (a clean migration, not an instant warm-up).
self.interaction_counts.insert(user_id, update_count);
cold += 1;
}
}
if warm + cold > 0 {
tracing::info!(
warm,
cold,
"multi-vector preferences restored from checkpoint"
);
}
Ok(())
}
}
// ── Free helpers ──────────────────────────────────────────────────────────────
/// Encode a multi-cluster value (`multi-vector-preference.md` §5):
/// `[version:1][n_clusters:1][dim:4 LE]` then per cluster
/// `[update_count:8 LE][importance_at_anchor:4 LE f32][anchor_ts:8 LE][f32*dim]`.
fn encode_multi_value(clusters: &[Cluster]) -> Vec<u8> {
let n = clusters.len().min(u8::MAX as usize);
let dim = clusters.first().map_or(0, |c| c.centroid.len());
let mut value = Vec::with_capacity(MULTI_HEADER_LEN + n * (CLUSTER_FIXED_LEN + dim * 4));
value.push(FORMAT_VERSION);
#[allow(clippy::cast_possible_truncation)]
value.push(n as u8);
#[allow(clippy::cast_possible_truncation)]
value.extend_from_slice(&(dim as u32).to_le_bytes());
for c in clusters.iter().take(n) {
value.extend_from_slice(&c.update_count.to_le_bytes());
value.extend_from_slice(&c.importance_at_anchor.to_le_bytes());
value.extend_from_slice(&c.anchor_ts_ns.to_le_bytes());
for v in &c.centroid {
value.extend_from_slice(&v.to_le_bytes());
}
}
value
}
/// Decode a multi-cluster value. Returns the clusters that parsed cleanly
/// (torn clusters dropped individually) or `None` if the header is malformed or
/// the dimension does not match this store. Each centroid is NaN-neutralized +
/// re-normalized at the boundary.
fn decode_multi_value(value: &[u8], expected_dim: usize) -> Option<Vec<Cluster>> {
// [version:1][n:1][dim:4]
if value.len() < MULTI_HEADER_LEN {
return None;
}
let n = value[1] as usize;
let dim = u32::from_le_bytes(value[2..6].try_into().ok()?) as usize;
if dim != expected_dim {
return None;
}
let cluster_bytes = CLUSTER_FIXED_LEN + dim * 4;
let mut clusters = Vec::with_capacity(n);
let mut off = MULTI_HEADER_LEN;
for _ in 0..n {
if off + cluster_bytes > value.len() {
// Torn tail: drop this and any remaining clusters, keep the prefix.
break;
}
let update_count = u64::from_le_bytes(value[off..off + 8].try_into().ok()?);
let importance = f32::from_le_bytes(value[off + 8..off + 12].try_into().ok()?);
let anchor_ts = u64::from_le_bytes(value[off + 12..off + 20].try_into().ok()?);
let mut centroid = Vec::with_capacity(dim);
for i in 0..dim {
let p = off + CLUSTER_FIXED_LEN + i * 4;
let f = f32::from_le_bytes(value[p..p + 4].try_into().ok()?);
// Neutralize NaN at the load boundary so a torn row cannot poison
// cosine scoring downstream (mirrors the single-vector restore).
centroid.push(if f.is_nan() { 0.0 } else { f });
}
l2_normalize(&mut centroid);
clusters.push(Cluster {
centroid,
update_count,
importance_at_anchor: if importance.is_finite() {
importance
} else {
0.0
},
anchor_ts_ns: anchor_ts,
});
off += cluster_bytes;
}
Some(clusters)
}
/// Whether `value` is a structurally well-formed multi-cluster row for its OWN
/// embedded `[n, dim]` header — i.e. its length matches `MULTI_HEADER_LEN +
/// n*(CLUSTER_FIXED_LEN + dim*4)` exactly.
///
/// Used by `restore` to disambiguate the two ways a `FORMAT_VERSION`-leading row
/// can fail `decode_multi_value(self.dim)`: a genuine multi row written for a
/// *different* store dim (well-formed for its own header → must be SKIPPED on a
/// schema dim change) vs a legacy row whose `update_count` low byte merely aliases
/// the sentinel (never length-consistent as a multi row → must be rescued as
/// legacy). An exact-length match cannot be satisfied by a legacy
/// `[count:8][dim:4][f32*legacy_dim]` row for the `[n, dim]` it would parse to.
fn is_well_formed_multi_header(value: &[u8]) -> bool {
if value.len() < MULTI_HEADER_LEN || value.first() != Some(&FORMAT_VERSION) {
return false;
}
let n = value[1] as usize;
let Some(dim) = value
.get(2..6)
.and_then(|b| b.try_into().ok())
.map(|b| u32::from_le_bytes(b) as usize)
else {
return false;
};
value.len() == MULTI_HEADER_LEN + n * (CLUSTER_FIXED_LEN + dim * 4)
}
/// Adaptive learning rate `base / (1 + ln_1p(count))` — identical to the
/// single-vector tier, applied per cluster.
#[allow(clippy::cast_precision_loss, clippy::cast_possible_truncation)]
fn adaptive_lr(base: f32, count: u64) -> f32 {
(f64::from(base) / (1.0 + (count as f64).ln_1p())) as f32
}
/// Importance mass to seed a crossover cluster with: the cold-start engagement
/// count it absorbed, floored at 1. Without this the seeded cold taste enters the
/// warm tier at importance 0.0 and is dropped from the top-M fan-out the moment a
/// divergent interest opens another cluster (which always outranks it at 1.0).
/// `.max(1)` guards a `set()`-seeded centroid (`update_count == 0` but a real
/// taste vector) from re-entering invisibly.
#[allow(clippy::cast_precision_loss)]
fn importance_seed(update_count: u64) -> f32 {
update_count.max(1) as f32
}
/// `pref = (1 - lr) * pref + lr * interaction`, then re-normalize.
fn blend_into(pref: &mut [f32], interaction: &[f32], lr: f32) {
for (p, &i) in pref.iter_mut().zip(interaction.iter()) {
*p = (1.0 - lr).mul_add(*p, lr * i);
}
l2_normalize(pref);
}
/// Index + cosine of the nearest centroid to `embedding` (vectors assumed
/// unit-length), or `None` for an empty cluster set. Ties keep the earliest
/// index, so the DP-means split test (`assign_or_split`, which needs the cosine)
/// and plain nearest-assignment ([`nearest_cluster`]) share one argmax and cannot
/// drift in their tie-break behavior.
fn nearest_cluster_cos(clusters: &[Cluster], embedding: &[f32]) -> Option<(usize, f32)> {
clusters
.iter()
.enumerate()
.map(|(i, c)| (i, dot(&c.centroid, embedding)))
.fold(None, |acc: Option<(usize, f32)>, (i, cos)| match acc {
Some((_, bc)) if bc >= cos => acc,
_ => Some((i, cos)),
})
}
/// Index of the nearest centroid by cosine (vectors assumed unit-length).
fn nearest_cluster(clusters: &[Cluster], embedding: &[f32]) -> Option<usize> {
nearest_cluster_cos(clusters, embedding).map(|(i, _)| i)
}
/// Max cosine of a (raw) candidate over the cluster centroids. Centroids are
/// unit-length; the candidate is normalized on the fly (matching the
/// single-vector `cosine_similarity` contract). `None` for an empty cluster set.
fn nearest_cosine_in(clusters: &[Cluster], candidate: &[f32]) -> Option<f32> {
if clusters.is_empty() {
return None;
}
let norm: f32 = candidate.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm < f32::EPSILON {
return Some(0.0);
}
let best = clusters
.iter()
.map(|c| dot(&c.centroid, candidate) / norm)
.fold(f32::NEG_INFINITY, f32::max);
Some(best)
}
/// Dot product of two equal-length slices (overlapping prefix if they differ).
fn dot(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
}
/// Total order over centroid byte patterns, for deterministic tie-breaks.
fn centroid_cmp(a: &[f32], b: &[f32]) -> std::cmp::Ordering {
a.iter()
.zip(b.iter())
.map(|(x, y)| x.to_bits().cmp(&y.to_bits()))
.find(|o| *o != std::cmp::Ordering::Equal)
.unwrap_or_else(|| a.len().cmp(&b.len()))
}
/// L2-normalize in place; an all-zero vector is left untouched.
fn l2_normalize(vec: &mut [f32]) {
let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > f32::EPSILON {
for v in vec.iter_mut() {
*v /= norm;
}
}
}
/// Wall-clock nanoseconds, via the engine's clock-anomaly-safe `Timestamp`.
fn now_ns() -> u64 {
crate::schema::Timestamp::now().as_nanos()
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::float_cmp)]
mod tests {
use super::*;
/// A unit vector along axis `axis` in `dim` dimensions.
fn axis_vec(dim: usize, axis: usize) -> Vec<f32> {
let mut v = vec![0.0f32; dim];
v[axis] = 1.0;
v
}
/// A blend of two axes (for an embedding near a centroid but not identical).
fn blend_vec(dim: usize, a: usize, b: usize, wa: f32, wb: f32) -> Vec<f32> {
let mut v = vec![0.0f32; dim];
v[a] = wa;
v[b] = wb;
l2_normalize(&mut v);
v
}
// ── Cold-start boundary ──────────────────────────────────────────────────
#[test]
fn cold_start_below_threshold_uses_single_vector() {
let pv = MultiPreferenceVectors::new(4);
// N-1 updates: still cold start, no clusters.
for _ in 0..(COLD_START_N - 1) {
assert!(pv.update_at(1, &axis_vec(4, 0), 1000));
}
assert!(!pv.is_warm(1), "must still be cold-start below N");
assert_eq!(pv.cluster_count(1), 0);
// get() returns the cold-start single vector.
let v = pv.get(1).unwrap();
assert!((v[0] - 1.0).abs() < 1e-5);
}
#[test]
fn crossing_threshold_seeds_first_cluster_from_cold_start() {
let pv = MultiPreferenceVectors::new(4);
// N-1 cold-start updates along axis 0.
for _ in 0..(COLD_START_N - 1) {
assert!(pv.update_at(1, &axis_vec(4, 0), 1000));
}
let cold = pv.get(1).unwrap();
// The Nth update (also axis 0, so it assigns to the seeded cluster) crosses.
assert!(pv.update_at(1, &axis_vec(4, 0), 2000));
assert!(pv.is_warm(1), "must be warm at/after N");
assert_eq!(
pv.cluster_count(1),
1,
"single coherent interest = one cluster"
);
// The cluster centroid is close to the cold-start vector it was seeded from.
let warm = pv.get(1).unwrap();
let cos = dot(&cold, &warm);
assert!(
cos > 0.99,
"seeded cluster must inherit the cold-start taste, cos={cos}"
);
// The cold-start row is dropped on crossover (no double-count alongside
// the cluster rows).
assert!(
pv.cold_start.get(1).is_none(),
"cold-start vector dropped on crossover"
);
}
#[test]
fn cold_start_dim_mismatch_rejected() {
let pv = MultiPreferenceVectors::new(4);
assert!(!pv.update_at(1, &[1.0, 0.0], 1000));
assert_eq!(
pv.interaction_count(1),
0,
"rejected update must not bump count"
);
}
// ── DP-means split / merge / cap ─────────────────────────────────────────
#[test]
fn distinct_interests_split_into_separate_clusters() {
// τ default 0.55; orthogonal axes have cosine 0 < τ ⇒ split.
let pv = MultiPreferenceVectors::new(8);
// Warm the user up on axis 0 past the threshold.
for t in 0..COLD_START_N {
assert!(pv.update_at(1, &axis_vec(8, 0), 1000 + t));
}
assert_eq!(pv.cluster_count(1), 1);
// Now engage with a genuinely distinct interest (axis 3).
assert!(pv.update_at(1, &axis_vec(8, 3), 5000));
assert_eq!(
pv.cluster_count(1),
2,
"orthogonal interest opens a new cluster"
);
// And a third (axis 6).
assert!(pv.update_at(1, &axis_vec(8, 6), 6000));
assert_eq!(pv.cluster_count(1), 3);
}
#[test]
fn similar_interests_merge_into_one_cluster() {
let pv = MultiPreferenceVectors::new(8);
for t in 0..COLD_START_N {
assert!(pv.update_at(1, &axis_vec(8, 0), 1000 + t));
}
assert_eq!(pv.cluster_count(1), 1);
// A vector with cosine > τ (0.55) to the axis-0 centroid: mostly axis 0.
// blend (0.9, 0.1) has cosine ≈ 0.9/sqrt(0.82) ≈ 0.994 to axis 0 ⇒ assign.
let near = blend_vec(8, 0, 1, 0.95, 0.05);
assert!(dot(&near, &axis_vec(8, 0)) > DEFAULT_SPLIT_THRESHOLD);
assert!(pv.update_at(1, &near, 5000));
assert_eq!(
pv.cluster_count(1),
1,
"a similar interest blends, not splits"
);
}
#[test]
fn k_max_cap_never_exceeded() {
// Use a low τ so every distinct axis splits, then push past K_MAX.
let pv = MultiPreferenceVectors::with_params(64, 0.1, 0.99, K_MAX, COLD_START_N, 1e9);
// Warm up.
for t in 0..COLD_START_N {
assert!(pv.update_at(1, &axis_vec(64, 0), 1000 + t));
}
// Engage with 30 distinct orthogonal interests — far more than K_MAX.
for axis in 1..31 {
assert!(pv.update_at(1, &axis_vec(64, axis), 2000 + axis as u64));
}
assert!(
pv.cluster_count(1) <= K_MAX,
"cluster count {} must never exceed K_MAX={K_MAX}",
pv.cluster_count(1)
);
assert_eq!(pv.cluster_count(1), K_MAX, "saturates exactly at the cap");
}
#[test]
fn over_cap_assigns_nearest_never_evicts() {
let pv = MultiPreferenceVectors::with_params(64, 0.1, 0.99, 3, COLD_START_N, 1e9);
for t in 0..COLD_START_N {
assert!(pv.update_at(1, &axis_vec(64, 0), 1000 + t));
}
// Open up to the cap (3) with distinct axes.
assert!(pv.update_at(1, &axis_vec(64, 1), 2000));
assert!(pv.update_at(1, &axis_vec(64, 2), 2001));
assert_eq!(pv.cluster_count(1), 3);
// A 4th distinct interest cannot open an 11th/4th cluster — assigned to
// nearest, count stays at the cap (no eviction = no lost interest).
assert!(pv.update_at(1, &axis_vec(64, 3), 3000));
assert_eq!(
pv.cluster_count(1),
3,
"over-cap assigns to nearest, never evicts"
);
}
// ── Determinism ──────────────────────────────────────────────────────────
#[test]
fn clustering_is_deterministic_for_same_input_order() {
let build = || {
let pv = MultiPreferenceVectors::new(8);
let seq = [0usize, 0, 3, 0, 3, 6, 3, 6, 0, 6];
for (t, &axis) in seq.iter().enumerate() {
assert!(pv.update_at(7, &axis_vec(8, axis), 1000 + t as u64));
}
pv
};
let a = build();
let b = build();
assert_eq!(a.cluster_count(7), b.cluster_count(7));
let va = a.query_vectors(7, 100_000, 3);
let vb = b.query_vectors(7, 100_000, 3);
assert_eq!(va.len(), vb.len());
for (ca, cb) in va.iter().zip(vb.iter()) {
assert_eq!(ca, cb, "same input order must yield byte-identical fan-out");
}
}
// ── Per-cluster adaptive learning rate ───────────────────────────────────
#[test]
fn per_cluster_adaptive_lr_matches_single_vector_formula() {
// The adaptive LR helper must equal the documented per-cluster formula.
assert!((adaptive_lr(0.1, 0) - 0.1).abs() < 1e-6);
let lr_100 = adaptive_lr(0.1, 100);
let lr_1000 = adaptive_lr(0.1, 1000);
assert!(lr_100 < 0.1 && lr_1000 < lr_100, "LR decays with count");
}
#[test]
fn cluster_update_count_advances_on_assignment() {
let pv = MultiPreferenceVectors::new(8);
for t in 0..COLD_START_N {
assert!(pv.update_at(1, &axis_vec(8, 0), 1000 + t));
}
// Many more assignments to the same cluster stabilize it (count grows).
for t in 0..50 {
assert!(pv.update_at(1, &blend_vec(8, 0, 1, 0.97, 0.03), 2000 + t));
}
assert_eq!(pv.cluster_count(1), 1);
// The centroid stays near axis 0 (a settled cluster resists drift).
let v = pv.get(1).unwrap();
assert!(dot(&v, &axis_vec(8, 0)) > 0.9);
}
// ── Forward-decayed importance + top-M selection ─────────────────────────
#[test]
fn stale_interest_falls_out_of_top_m_via_decay() {
// half-life 1s so importance decays fast in the test window.
let pv = MultiPreferenceVectors::with_params(8, 0.1, 0.55, K_MAX, COLD_START_N, 1.0);
let sec = 1_000_000_000u64;
// Warm up on axis 0 (the "old" interest).
for t in 0..COLD_START_N {
assert!(pv.update_at(1, &axis_vec(8, 0), t * sec));
}
// Much later, engage heavily with axes 3 and 6 (the "fresh" interests).
for t in 0..5 {
assert!(pv.update_at(1, &axis_vec(8, 3), (100 + t) * sec));
assert!(pv.update_at(1, &axis_vec(8, 6), (100 + t) * sec));
}
assert_eq!(pv.cluster_count(1), 3);
// At t = 105s, the axis-0 cluster (last touched ~104s ago) has decayed far
// below the two fresh ones; top-2 must be the fresh interests.
let now = 105 * sec;
let top2 = pv.query_vectors(1, now, 2);
assert_eq!(top2.len(), 2);
let axis0 = axis_vec(8, 0);
for v in &top2 {
assert!(
dot(v, &axis0) < 0.5,
"the stale axis-0 interest must not be in the top-2"
);
}
}
#[test]
fn query_vectors_top_m_respects_min_k_active() {
let pv = MultiPreferenceVectors::new(8);
for t in 0..COLD_START_N {
assert!(pv.update_at(1, &axis_vec(8, 0), 1000 + t));
}
assert!(pv.update_at(1, &axis_vec(8, 3), 5000));
// Only 2 clusters; asking for top-5 returns 2.
assert_eq!(pv.query_vectors(1, 100_000, 5).len(), 2);
}
#[test]
fn query_vectors_cold_start_returns_single_vector() {
let pv = MultiPreferenceVectors::new(4);
assert!(pv.update_at(1, &axis_vec(4, 0), 1000));
let vs = pv.query_vectors(1, 2000, 3);
assert_eq!(vs.len(), 1, "cold-start fan-out is one vector");
}
// ── Nearest-cosine (Stage-3 boost) ───────────────────────────────────────
#[test]
fn cosine_similarity_is_max_over_clusters() {
let pv = MultiPreferenceVectors::new(8);
for t in 0..COLD_START_N {
assert!(pv.update_at(1, &axis_vec(8, 0), 1000 + t));
}
assert!(pv.update_at(1, &axis_vec(8, 3), 5000));
// A candidate aligned with the axis-3 cluster scores ~1.0 even though the
// axis-0 cluster would score ~0 — the single-vector mean would dilute it.
let cand = axis_vec(8, 3);
let cos = pv.cosine_similarity(1, &cand).unwrap();
assert!(cos > 0.99, "max-over-clusters cosine, got {cos}");
}
// ── Checkpoint / restore ─────────────────────────────────────────────────
fn warm_user(pv: &MultiPreferenceVectors, user: u64) {
for t in 0..COLD_START_N {
assert!(pv.update_at(user, &axis_vec(8, 0), 1000 + t));
}
assert!(pv.update_at(user, &axis_vec(8, 3), 5000));
assert!(pv.update_at(user, &axis_vec(8, 6), 6000));
}
#[test]
fn checkpoint_restore_roundtrip_multi_cluster() {
use crate::storage::InMemoryBackend;
let storage = InMemoryBackend::new();
let src = MultiPreferenceVectors::new(8);
warm_user(&src, 1); // 3 clusters
// A cold-start user too.
assert!(src.update_at(2, &axis_vec(8, 1), 1000));
src.checkpoint(&storage).unwrap();
let dst = MultiPreferenceVectors::new(8);
dst.restore(&storage).unwrap();
assert_eq!(dst.cluster_count(1), 3, "warm user clusters round-trip");
assert!(dst.is_warm(1));
assert!(!dst.is_warm(2), "cold-start user restores as cold-start");
// The fan-out is preserved.
let now = 7000;
let a = src.query_vectors(1, now, 3);
let b = dst.query_vectors(1, now, 3);
assert_eq!(a.len(), b.len());
for (ca, cb) in a.iter().zip(b.iter()) {
for (x, y) in ca.iter().zip(cb.iter()) {
assert!(
(x - y).abs() < 1e-6,
"centroid component must survive round-trip"
);
}
}
}
#[test]
fn restore_reads_legacy_single_vector_rows_as_cold_start() {
use crate::{
schema::EntityId,
storage::{InMemoryBackend, StorageEngine, Tag, encode_key},
};
let storage = InMemoryBackend::new();
// Hand-write a LEGACY row: [count:8 LE][dim:4 LE][f32*dim] (no version byte).
let dim = 8usize;
let mut value = Vec::new();
value.extend_from_slice(&3u64.to_le_bytes());
value.extend_from_slice(&(dim as u32).to_le_bytes());
let mut vec = axis_vec(dim, 2);
for v in &vec {
value.extend_from_slice(&v.to_le_bytes());
}
l2_normalize(&mut vec);
let key = encode_key(EntityId::new(0), Tag::Preference, &42u64.to_be_bytes());
storage.put(&key, &value).unwrap();
let pv = MultiPreferenceVectors::new(dim);
pv.restore(&storage).unwrap();
assert!(!pv.is_warm(42), "legacy row is a K=1 cold-start user");
let restored = pv.get(42).unwrap();
assert!(
dot(&restored, &axis_vec(dim, 2)) > 0.99,
"legacy taste preserved"
);
// The interaction gate is seeded from the legacy count, so it stays cold
// until it organically crosses the threshold.
assert_eq!(pv.interaction_count(42), 3);
}
#[test]
fn restore_renormalizes_torn_cluster_to_unit_length() {
use crate::{
schema::EntityId,
storage::{InMemoryBackend, StorageEngine, Tag, encode_key},
};
let storage = InMemoryBackend::new();
let dim = 3usize;
// Hand-write a NEW-format row with one cluster whose centroid is NaN-laced
// and non-unit: [3, NaN, 4] zeroes to [3,0,4] (norm 5) ⇒ must normalize.
let mut value = Vec::new();
value.push(FORMAT_VERSION);
value.push(1u8); // n_clusters
value.extend_from_slice(&(dim as u32).to_le_bytes());
value.extend_from_slice(&5u64.to_le_bytes()); // update_count
value.extend_from_slice(&2.0f32.to_le_bytes()); // importance_at_anchor
value.extend_from_slice(&1234u64.to_le_bytes()); // anchor_ts
for f in [3.0f32, f32::NAN, 4.0] {
value.extend_from_slice(&f.to_le_bytes());
}
let key = encode_key(EntityId::new(0), Tag::Preference, &9u64.to_be_bytes());
storage.put(&key, &value).unwrap();
let pv = MultiPreferenceVectors::new(dim);
pv.restore(&storage).unwrap();
assert!(pv.is_warm(9));
let v = pv.get(9).unwrap();
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!(
(norm - 1.0).abs() < 1e-5,
"torn centroid re-normalized, norm={norm}"
);
assert!((v[0] - 0.6).abs() < 1e-5 && (v[2] - 0.8).abs() < 1e-5);
}
#[test]
fn restore_drops_torn_tail_cluster_keeps_prefix() {
use crate::{
schema::EntityId,
storage::{InMemoryBackend, StorageEngine, Tag, encode_key},
};
let storage = InMemoryBackend::new();
let dim = 3usize;
// Header claims 2 clusters but the value only holds 1 full cluster + a
// truncated second. The first must survive; the torn tail is dropped.
let mut value = Vec::new();
value.push(FORMAT_VERSION);
value.push(2u8); // claims 2
value.extend_from_slice(&(dim as u32).to_le_bytes());
// cluster 0 (complete)
value.extend_from_slice(&1u64.to_le_bytes());
value.extend_from_slice(&1.0f32.to_le_bytes());
value.extend_from_slice(&100u64.to_le_bytes());
for f in axis_vec(dim, 0) {
value.extend_from_slice(&f.to_le_bytes());
}
// cluster 1 (TORN: only a few bytes)
value.extend_from_slice(&[0u8, 1u8, 2u8]);
let key = encode_key(EntityId::new(0), Tag::Preference, &11u64.to_be_bytes());
storage.put(&key, &value).unwrap();
let pv = MultiPreferenceVectors::new(dim);
pv.restore(&storage).unwrap();
assert_eq!(pv.cluster_count(11), 1, "torn tail dropped, prefix kept");
}
#[test]
fn restore_skips_dimension_mismatch() {
use crate::storage::InMemoryBackend;
let storage = InMemoryBackend::new();
let src = MultiPreferenceVectors::new(8);
warm_user(&src, 1);
src.checkpoint(&storage).unwrap();
// Restore into a store with a DIFFERENT dim — must skip, not crash.
let dst = MultiPreferenceVectors::new(16);
dst.restore(&storage).unwrap();
assert_eq!(dst.cluster_count(1), 0);
assert!(dst.is_empty());
}
#[test]
fn warm_user_survives_restart_no_demotion_to_cold_start() {
use crate::storage::InMemoryBackend;
let storage = InMemoryBackend::new();
let src = MultiPreferenceVectors::new(8);
warm_user(&src, 1);
src.checkpoint(&storage).unwrap();
let dst = MultiPreferenceVectors::new(8);
dst.restore(&storage).unwrap();
// A post-restore engagement must keep the user warm (not reset to cold).
assert!(dst.update_at(1, &axis_vec(8, 0), 9000));
assert!(
dst.is_warm(1),
"restored warm user must not demote on next update"
);
}
// ── Importance composition with the canonical decay kernel ───────────────
#[test]
fn importance_uses_forward_decay_kernel() {
let half_life_secs = 3600.0;
let pv =
MultiPreferenceVectors::with_params(4, 0.1, 0.55, K_MAX, COLD_START_N, half_life_secs);
let sec = 1_000_000_000u64;
for t in 0..COLD_START_N {
assert!(pv.update_at(1, &axis_vec(4, 0), t * sec));
}
// Importance right after the last engagement.
let last_ts = (COLD_START_N - 1) * sec;
let imp_now = pv.cluster_importances(1, last_ts);
// One half-life later it must be ~half (within 1%).
let imp_later = pv.cluster_importances(1, last_ts + 3600 * sec);
assert!(!imp_now.is_empty());
let ratio = imp_later[0] / imp_now[0];
assert!(
(ratio - 0.5).abs() < 0.02,
"importance must decay by half over one half-life, ratio={ratio}"
);
}
// ── Regression: format-version collision (BLOCKER) ────────────────────────
#[test]
fn restore_rescues_legacy_row_whose_count_low_byte_equals_format_version() {
use crate::{
schema::EntityId,
storage::{InMemoryBackend, StorageEngine, Tag, encode_key},
};
// A legacy single-vector row's first byte is the LOW byte of its LE
// update_count. For counts ≡ FORMAT_VERSION (mod 256) — e.g. 2, 258, 514 —
// it aliases the multi-cluster sentinel and the row MUST still restore as
// cold-start, never be silently dropped.
let dim = 8usize;
for count in [2u64, 258, 514] {
let storage = InMemoryBackend::new();
let mut vec = axis_vec(dim, 3);
l2_normalize(&mut vec);
let value = crate::entities::preference::encode_legacy_row(count, &vec);
assert_eq!(
value[0], FORMAT_VERSION,
"count {count} must alias the sentinel"
);
let key = encode_key(EntityId::new(0), Tag::Preference, &7u64.to_be_bytes());
storage.put(&key, &value).unwrap();
let pv = MultiPreferenceVectors::new(dim);
pv.restore(&storage).unwrap();
assert!(!pv.is_warm(7), "legacy row stays cold-start, count={count}");
let restored = pv
.get(7)
.unwrap_or_else(|| panic!("legacy row silently dropped for count={count}"));
assert!(
dot(&restored, &axis_vec(dim, 3)) > 0.99,
"taste preserved, count={count}"
);
assert_eq!(pv.interaction_count(7), count);
}
}
#[test]
fn cold_start_user_at_count_two_survives_full_checkpoint_restore() {
use crate::storage::InMemoryBackend;
// End-to-end: a real user with exactly 2 interactions (the modal cold-start
// count, whose legacy row leads with 0x02) round-trips checkpoint+restore.
let storage = InMemoryBackend::new();
let src = MultiPreferenceVectors::new(8);
assert!(src.update_at(42, &axis_vec(8, 1), 1000));
assert!(src.update_at(42, &axis_vec(8, 1), 2000));
assert!(!src.is_warm(42));
assert_eq!(src.interaction_count(42), 2);
src.checkpoint(&storage).unwrap();
let dst = MultiPreferenceVectors::new(8);
dst.restore(&storage).unwrap();
let restored = dst
.get(42)
.expect("count-2 cold-start user must survive restore");
assert!(dot(&restored, &axis_vec(8, 1)) > 0.99);
}
// ── Regression: seeded cold-start cluster importance (CRITICAL) ────────────
#[test]
fn crossover_diverging_from_cold_seed_keeps_seed_visible_in_top_m() {
// 4 cold-start updates on axis 0, then a crossover on an ORTHOGONAL axis 3
// that opens a new cluster. The dominant cold taste must NOT enter the warm
// tier invisible: its importance is seeded from the cold engagement mass.
let pv = MultiPreferenceVectors::new(8);
for t in 0..(COLD_START_N - 1) {
assert!(pv.update_at(1, &axis_vec(8, 0), 1000 + t));
}
assert!(pv.update_at(1, &axis_vec(8, 3), 5000)); // crossover, opens cluster 1
assert_eq!(pv.cluster_count(1), 2, "orthogonal crossover splits");
let imps = pv.cluster_importances(1, 5000);
assert!(
imps.iter().all(|&i| i > 0.0),
"no cluster may enter the warm tier at importance 0.0, got {imps:?}"
);
// The 4-interaction seed outranks the 1-interaction fresh cluster, so it is
// the primary vector AND present in a top-1 fan-out.
let primary = pv.get(1).unwrap();
assert!(
dot(&primary, &axis_vec(8, 0)) > 0.99,
"the dominant cold taste must be primary, not the 1-shot interest"
);
let top1 = pv.query_vectors(1, 5000, 1);
assert_eq!(top1.len(), 1);
assert!(dot(&top1[0], &axis_vec(8, 0)) > 0.99);
}
// ── Regression: tier disjointness (CRITICAL race + set guard) ──────────────
#[test]
fn checkpoint_skips_cold_row_for_user_also_in_clusters_no_demotion() {
use crate::storage::InMemoryBackend;
// The disjointness-violating state cannot arise via the public API after the
// tier-transition + set() fixes, but checkpoint must be robust to it anyway:
// force a stale cold-start row alongside a warm user's clusters and confirm
// the warm row survives (the legacy row is skipped, not last-write clobbered).
let pv = MultiPreferenceVectors::new(8);
warm_user(&pv, 1); // 3 clusters, cold_start row removed on crossover
assert!(pv.cold_start.set(1, axis_vec(8, 5)));
assert!(
pv.is_warm(1) && pv.cold_start.contains(1),
"user forced into BOTH tiers"
);
let storage = InMemoryBackend::new();
pv.checkpoint(&storage).unwrap();
let dst = MultiPreferenceVectors::new(8);
dst.restore(&storage).unwrap();
assert!(
dst.is_warm(1),
"warm user must NOT be demoted by a stale cold row"
);
assert_eq!(dst.cluster_count(1), 3);
}
#[test]
fn set_on_warm_user_is_rejected_and_keeps_tiers_disjoint() {
let pv = MultiPreferenceVectors::new(8);
warm_user(&pv, 1);
assert!(pv.is_warm(1));
assert!(
!pv.set(1, axis_vec(8, 5)),
"set() must be rejected for a warm user"
);
assert!(!pv.cold_start.contains(1), "no stale cold row created");
assert!(pv.is_warm(1));
}
#[test]
fn concurrent_crossover_keeps_user_in_exactly_one_tier() {
use std::sync::Arc;
// Hammer the cold→warm boundary from two threads on the same fresh user.
// The per-user atomic tier transition must never leave the user in BOTH
// tiers (which a later checkpoint would collapse, demoting them).
for round in 0..200u64 {
let pv = Arc::new(MultiPreferenceVectors::new(8));
for t in 0..(COLD_START_N - 1) {
assert!(pv.update_at(1, &axis_vec(8, 0), round * 1000 + t));
}
let a = {
let pv = Arc::clone(&pv);
std::thread::spawn(move || pv.update_at(1, &axis_vec(8, 0), 100_000 + round * 10))
};
let b = {
let pv = Arc::clone(&pv);
std::thread::spawn(move || pv.update_at(1, &axis_vec(8, 3), 100_001 + round * 10))
};
assert!(a.join().unwrap());
assert!(b.join().unwrap());
assert!(
!(pv.is_warm(1) && pv.cold_start.contains(1)),
"round {round}: user must be in exactly one tier, never both"
);
assert!(pv.is_warm(1), "round {round}: user crossed to warm");
}
}
#[test]
fn concurrent_set_and_crossover_keep_tiers_disjoint() {
use std::sync::Arc;
// `set()` racing a crossover `update_at` on the same user must not land the
// user in both tiers: the shared per-user `interaction_counts` guard
// serializes them, so either set runs before the crossover (its cold vector
// is then seeded + removed) or after (it observes warm and is rejected).
for round in 0..200u64 {
let pv = Arc::new(MultiPreferenceVectors::new(8));
for t in 0..(COLD_START_N - 1) {
assert!(pv.update_at(1, &axis_vec(8, 0), round * 1000 + t));
}
let a = {
let pv = Arc::clone(&pv);
std::thread::spawn(move || pv.update_at(1, &axis_vec(8, 3), 50_000 + round))
};
let b = {
let pv = Arc::clone(&pv);
std::thread::spawn(move || pv.set(1, axis_vec(8, 5)))
};
let _ = a.join().unwrap();
let _ = b.join().unwrap();
assert!(
!(pv.is_warm(1) && pv.cold_start.contains(1)),
"round {round}: set() racing a crossover left the user in both tiers"
);
}
}
#[test]
fn restore_skips_dim_mismatched_multi_row_not_loaded_as_garbage() {
use crate::storage::InMemoryBackend;
// A genuine multi-cluster row written for dim 8, with cluster 0's
// update_count crafted so the row bytes a legacy decoder would read as
// `dim` (value[8..12]) equal the NEW store dim (4). Restoring into a dim-4
// store must SKIP it (a schema embedding-dim change), NOT mis-decode it as a
// legacy cold-start vector of garbage bytes (the pre-fix fall-through bug).
let src = MultiPreferenceVectors::new(8);
{
let mut clusters = src.clusters.entry(1).or_default();
clusters.push(Cluster {
centroid: axis_vec(8, 0),
// update_count = 4 << 16 ⇒ row bytes value[8..12] == [4,0,0,0] == dim 4.
update_count: 4u64 << 16,
importance_at_anchor: 1.0,
anchor_ts_ns: 100,
});
}
let storage = InMemoryBackend::new();
src.checkpoint(&storage).unwrap();
let dst = MultiPreferenceVectors::new(4); // DIFFERENT dim
dst.restore(&storage).unwrap();
assert!(
dst.is_empty() && dst.get(1).is_none(),
"a dim-mismatched multi row must be skipped, never loaded as a legacy garbage vector"
);
}
#[test]
fn get_matches_query_vectors_primary_on_ties() {
let pv = MultiPreferenceVectors::new(8);
let now = 10_000u64;
// Normal warm user: get() agrees with the top of the fan-out.
warm_user(&pv, 1);
assert_eq!(
pv.get(1).as_deref(),
pv.query_vectors(1, now, 1).first().map(Vec::as_slice)
);
// Tie case: two clusters with identical importance + anchor and byte-distinct
// centroids. `get()` (max_by, last-of-equal under a REVERSED centroid
// tie-break) must pick the SAME cluster as `query_vectors[0]` (sort-first,
// ascending centroid bytes) — see `top_importance_cluster`.
let tie = |a, b| {
vec![
Cluster {
centroid: axis_vec(8, a),
update_count: 1,
importance_at_anchor: 3.0,
anchor_ts_ns: 0,
},
Cluster {
centroid: axis_vec(8, b),
update_count: 1,
importance_at_anchor: 3.0,
anchor_ts_ns: 0,
},
]
};
pv.clusters.insert(2, tie(0, 1));
pv.clusters.insert(3, tie(1, 0)); // reversed insertion order
assert_eq!(
pv.get(2).unwrap(),
pv.query_vectors(2, now, 1)[0],
"get() must equal query_vectors(..,1)[0] on a tie"
);
assert_eq!(pv.get(3).unwrap(), pv.query_vectors(3, now, 1)[0]);
assert_eq!(
pv.get(2).unwrap(),
pv.get(3).unwrap(),
"tie winner must be insertion-order-independent"
);
}
mod proptests {
use proptest::prelude::*;
use super::*;
proptest! {
/// Every cluster centroid stays unit-length (or zero) after any
/// sequence of warm-tier updates.
#[test]
fn centroids_stay_unit_length(
axes in proptest::collection::vec(0usize..8, 6..40),
) {
let pv = MultiPreferenceVectors::new(8);
for (t, &axis) in axes.iter().enumerate() {
let _ = pv.update_at(1, &axis_vec(8, axis), 1000 + t as u64);
}
if let Some(clusters) = pv.clusters.get(&1) {
for c in clusters.iter() {
let norm: f32 = c.centroid.iter().map(|x| x * x).sum::<f32>().sqrt();
prop_assert!(
(norm - 1.0).abs() < 1e-4 || norm < f32::EPSILON,
"centroid norm {norm} not unit"
);
}
}
}
/// Cluster count is always in [1, K_MAX] for a warm user, and never
/// exceeds the cap regardless of how adversarial the interest order is.
#[test]
fn cluster_count_bounded(
axes in proptest::collection::vec(0usize..32, 10..80),
) {
// dim 32, tight τ so distinct axes split aggressively.
let pv = MultiPreferenceVectors::with_params(32, 0.1, 0.9, K_MAX, COLD_START_N, 1e9);
for (t, &axis) in axes.iter().enumerate() {
let _ = pv.update_at(1, &axis_vec(32, axis), 1000 + t as u64);
}
let k = pv.cluster_count(1);
if pv.is_warm(1) {
prop_assert!((1..=K_MAX).contains(&k), "cluster count {k} out of [1, {K_MAX}]");
}
}
/// Checkpoint→restore preserves the warm cluster count for any warm user.
#[test]
fn checkpoint_restore_preserves_cluster_count(
axes in proptest::collection::vec(0usize..8, 8..30),
) {
use crate::storage::InMemoryBackend;
let src = MultiPreferenceVectors::new(8);
for (t, &axis) in axes.iter().enumerate() {
let _ = src.update_at(5, &axis_vec(8, axis), 1000 + t as u64);
}
let storage = InMemoryBackend::new();
src.checkpoint(&storage).unwrap();
let dst = MultiPreferenceVectors::new(8);
dst.restore(&storage).unwrap();
prop_assert_eq!(src.cluster_count(5), dst.cluster_count(5));
}
}
}
}