use std::collections::HashMap;
use std::sync::Arc;
use async_trait::async_trait;
use bytes::Bytes;
use futures::stream::{self, BoxStream, StreamExt};
use lunaris_core::storage::keyword::{KeywordHit, KeywordPort};
use lunaris_core::storage::types::{
CypherQuery, Filter, GraphResult, Lsn, QueueMsg, Row, VectorHit, WriteOp,
};
use lunaris_core::{
BiTemporal, Embedder, Hlc, HlcClock, LunarisError, StorageCapabilities, StorageError,
StoragePort, StubEmbedder,
};
use lunaris_rerank::{NoopReranker, RerankCandidate, Reranker};
use lunaris_retrieve::{Query, QueryContext, RawHit, Retriever, SourceOp, Vector};
use parking_lot::Mutex;
use serde_json::json;
#[derive(Default)]
struct RecordingStorage {
vector_hits: Mutex<Vec<VectorHit>>,
chunks_by_key: Mutex<HashMap<Vec<u8>, Vec<u8>>>,
}
impl RecordingStorage {
fn new() -> Self {
Self::default()
}
fn set_vector_hits(&self, hits: Vec<VectorHit>) {
*self.vector_hits.lock() = hits;
}
}
#[async_trait]
impl StoragePort for RecordingStorage {
async fn atomic_write(
&self,
_scope: &lunaris_core::Scope,
_ops: &[WriteOp],
) -> Result<Lsn, StorageError> {
Ok(Lsn::ZERO)
}
async fn vector_search(
&self,
_scope: &lunaris_core::Scope,
_index: &str,
_query: &[f32],
_k: usize,
_filter: Option<&Filter>,
_as_of: Option<Hlc>,
_rerank: bool,
) -> Result<Vec<VectorHit>, StorageError> {
Ok(self.vector_hits.lock().clone())
}
async fn graph_traverse(
&self,
_scope: &lunaris_core::Scope,
_q: &CypherQuery,
_as_of: Option<Hlc>,
) -> Result<GraphResult, StorageError> {
Err(StorageError::NotSupported("RecordingStorage::graph_traverse"))
}
async fn scan_range(
&self,
_scope: &lunaris_core::Scope,
_prefix: &[u8],
_as_of: Option<Hlc>,
) -> Result<BoxStream<'_, Result<(Bytes, Bytes), StorageError>>, StorageError> {
Ok(stream::iter(Vec::<Result<(Bytes, Bytes), StorageError>>::new()).boxed())
}
async fn read_as_of(
&self,
_scope: &lunaris_core::Scope,
key: &[u8],
_as_of: Hlc,
) -> Result<Option<Row<Bytes>>, StorageError> {
if let Some(v) = self.chunks_by_key.lock().get(key).cloned() {
return Ok(Some(Row {
key: key.to_vec(),
value: Bytes::from(v),
bt: BiTemporal::at(Hlc::ZERO, Hlc::ZERO),
}));
}
Ok(None)
}
async fn publish(
&self,
_scope: &lunaris_core::Scope,
_topic: &str,
_partition: u16,
_payload: Bytes,
) -> Result<u64, StorageError> {
Err(StorageError::NotSupported("RecordingStorage::publish"))
}
async fn subscribe(
&self,
_scope: &lunaris_core::Scope,
_group: &str,
_topic: &str,
_partition: u16,
) -> Result<BoxStream<'static, Result<QueueMsg, StorageError>>, StorageError> {
Err(StorageError::NotSupported("RecordingStorage::subscribe"))
}
fn capabilities(&self) -> StorageCapabilities {
StorageCapabilities {
bi_temporal_native: false,
graph_native: false,
rerank_native: false,
queue_native: false,
max_vector_dim: 768,
native_rrf: false,
max_scopes_recommended: 0,
cypher_dialect: lunaris_core::CypherDialect::Legacy,
graph_decay_native: false,
graph_navigate_native: false,
}
}
}
#[async_trait]
impl KeywordPort for RecordingStorage {
async fn keyword_search(
&self,
_scope: &lunaris_core::Scope,
_index: &str,
_query: &str,
_k: usize,
_filter: Option<&Filter>,
_as_of: Option<Hlc>,
) -> Result<Vec<KeywordHit>, StorageError> {
Ok(Vec::new())
}
}
fn vh(id: &[u8], score: f32) -> VectorHit {
VectorHit { id: id.to_vec(), score, rerank_applied: false, metadata: json!({}) }
}
fn build_ctx(
rec: Arc<RecordingStorage>,
) -> (Arc<dyn StoragePort>, Arc<dyn KeywordPort>, Arc<dyn Embedder>) {
let storage: Arc<dyn StoragePort> = rec.clone();
let keyword: Arc<dyn KeywordPort> = rec.clone();
let embedder: Arc<dyn Embedder> = Arc::new(StubEmbedder::new(768));
(storage, keyword, embedder)
}
fn seed_chunk(rec: &RecordingStorage, text: &str) -> Vec<u8> {
use lunaris_core::primitives::Chunk;
let clock = HlcClock::new(0);
let episode_id = ulid::Ulid::new();
let chunk = Chunk::new(
lunaris_core::Scope::dev(),
episode_id,
text,
4,
0,
vec!["Notes".to_string()],
&clock,
);
let id_bytes = chunk.id.to_bytes().to_vec();
let key = lunaris_core::keyspace::chunk_key(&lunaris_core::Scope::dev(), chunk.id);
rec.chunks_by_key.lock().insert(key, serde_json::to_vec(&chunk).unwrap());
id_bytes
}
struct LexicographicReranker;
#[async_trait]
impl Reranker for LexicographicReranker {
async fn rerank(
&self,
_query: &str,
mut docs: Vec<RerankCandidate>,
) -> Result<Vec<RerankCandidate>, LunarisError> {
docs.sort_by(|a, b| b.id.cmp(&a.id));
let n = docs.len();
for (i, d) in docs.iter_mut().enumerate() {
d.score = (n - i) as f32;
}
Ok(docs)
}
fn applies(&self) -> bool {
true
}
}
struct RecordingReranker {
received: Mutex<Vec<RerankCandidate>>,
}
impl RecordingReranker {
fn new() -> Self {
Self { received: Mutex::new(Vec::new()) }
}
}
#[async_trait]
impl Reranker for RecordingReranker {
async fn rerank(
&self,
_query: &str,
docs: Vec<RerankCandidate>,
) -> Result<Vec<RerankCandidate>, LunarisError> {
*self.received.lock() = docs.clone();
Ok(docs)
}
fn applies(&self) -> bool {
true
}
}
#[tokio::test]
async fn rerank_with_noop_preserves_order() {
let rec = Arc::new(RecordingStorage::new());
let id_a = seed_chunk(&rec, "alpha document text");
let id_b = seed_chunk(&rec, "beta document text");
let id_c = seed_chunk(&rec, "gamma document text");
rec.set_vector_hits(vec![vh(&id_a, 0.9), vh(&id_b, 0.7), vh(&id_c, 0.5)]);
let (storage, keyword, embedder) = build_ctx(rec.clone());
let ctx = QueryContext::new(
Query::text("anything"),
lunaris_core::Scope::dev(),
embedder,
storage,
keyword,
);
let root = Vector::new("chunks", 30).rerank(Arc::new(NoopReranker));
let raw = root.retrieve(&ctx).await.unwrap();
assert_eq!(raw.len(), 3);
assert_eq!(raw[0].id, id_a);
assert_eq!(raw[1].id, id_b);
assert_eq!(raw[2].id, id_c);
for h in &raw {
assert!(!h.rerank_applied, "NoopReranker MUST set rerank_applied=false");
assert_eq!(h.source_op, SourceOp::Reranked);
}
}
#[tokio::test]
async fn rerank_with_mock_inverts_order() {
let rec = Arc::new(RecordingStorage::new());
let id_a = seed_chunk(&rec, "doc a");
let id_b = seed_chunk(&rec, "doc b");
let id_c = seed_chunk(&rec, "doc c");
rec.set_vector_hits(vec![vh(&id_a, 0.9), vh(&id_b, 0.7), vh(&id_c, 0.5)]);
let (storage, keyword, embedder) = build_ctx(rec.clone());
let ctx =
QueryContext::new(Query::text("q"), lunaris_core::Scope::dev(), embedder, storage, keyword);
let root = Vector::new("chunks", 30).rerank(Arc::new(LexicographicReranker));
let raw = root.retrieve(&ctx).await.unwrap();
assert_eq!(raw.len(), 3);
let mut expected = vec![id_a.clone(), id_b.clone(), id_c.clone()];
expected.sort_by(|a, b| b.cmp(a));
let actual: Vec<Vec<u8>> = raw.iter().map(|h| h.id.clone()).collect();
assert_eq!(actual, expected, "LexicographicReranker MUST sort hits by id desc");
for h in &raw {
assert!(h.rerank_applied, "real reranker MUST set rerank_applied=true");
assert_eq!(h.source_op, SourceOp::Reranked);
}
}
#[tokio::test]
async fn rerank_truncates_to_k_in() {
let rec = Arc::new(RecordingStorage::new());
let mut hits = Vec::with_capacity(50);
for i in 0..50 {
let id = seed_chunk(&rec, &format!("doc {i}"));
hits.push(vh(&id, 1.0 - (i as f32) * 0.01));
}
rec.set_vector_hits(hits);
let (storage, keyword, embedder) = build_ctx(rec.clone());
let ctx =
QueryContext::new(Query::text("q"), lunaris_core::Scope::dev(), embedder, storage, keyword);
let recorder = Arc::new(RecordingReranker::new());
let root = Vector::new("chunks", 50).rerank(recorder.clone() as Arc<dyn Reranker>);
let raw = root.retrieve(&ctx).await.unwrap();
assert_eq!(raw.len(), 30, "rerank must truncate to k_in=30 before calling reranker");
let received_count = recorder.received.lock().len();
assert_eq!(received_count, 30, "reranker MUST receive exactly 30 candidates");
}
#[tokio::test]
async fn rerank_with_top_in_widens_the_truncation_window() {
let rec = Arc::new(RecordingStorage::new());
let mut hits = Vec::with_capacity(50);
for i in 0..50 {
let id = seed_chunk(&rec, &format!("doc {i}"));
hits.push(vh(&id, 1.0 - (i as f32) * 0.01));
}
rec.set_vector_hits(hits);
let (storage, keyword, embedder) = build_ctx(rec.clone());
let ctx =
QueryContext::new(Query::text("q"), lunaris_core::Scope::dev(), embedder, storage, keyword);
let recorder = Arc::new(RecordingReranker::new());
let upstream: Box<dyn Retriever> = Box::new(Vector::new("chunks", 50));
let root = lunaris_retrieve::RerankRetriever::with_top_in(
upstream,
recorder.clone() as Arc<dyn Reranker>,
45,
);
let raw = root.retrieve(&ctx).await.unwrap();
assert_eq!(
raw.len(),
45,
"with_top_in(45) must widen the truncation window past DEFAULT_RERANK_TOP_IN=30"
);
assert_eq!(
recorder.received.lock().len(),
45,
"reranker MUST receive exactly 45 candidates when k_in=45"
);
}
#[tokio::test]
async fn rerank_partial_hydrates_text() {
let rec = Arc::new(RecordingStorage::new());
let id_x = seed_chunk(&rec, "the quick brown fox");
let id_y = seed_chunk(&rec, "jumps over the lazy dog");
rec.set_vector_hits(vec![vh(&id_x, 0.9), vh(&id_y, 0.8)]);
let (storage, keyword, embedder) = build_ctx(rec.clone());
let ctx =
QueryContext::new(Query::text("q"), lunaris_core::Scope::dev(), embedder, storage, keyword);
let recorder = Arc::new(RecordingReranker::new());
let root = Vector::new("chunks", 30).rerank(recorder.clone() as Arc<dyn Reranker>);
let _ = root.retrieve(&ctx).await.unwrap();
let received = recorder.received.lock().clone();
assert_eq!(received.len(), 2);
let by_id: HashMap<Vec<u8>, String> = received.into_iter().map(|c| (c.id, c.text)).collect();
assert_eq!(by_id.get(&id_x).unwrap(), "the quick brown fox");
assert_eq!(by_id.get(&id_y).unwrap(), "jumps over the lazy dog");
}
#[tokio::test]
async fn rerank_validates_doc_count() {
struct DroppingReranker;
#[async_trait]
impl Reranker for DroppingReranker {
async fn rerank(
&self,
_q: &str,
mut docs: Vec<RerankCandidate>,
) -> Result<Vec<RerankCandidate>, LunarisError> {
docs.pop(); Ok(docs)
}
fn applies(&self) -> bool {
true
}
}
let rec = Arc::new(RecordingStorage::new());
let id_a = seed_chunk(&rec, "alpha");
let id_b = seed_chunk(&rec, "beta");
rec.set_vector_hits(vec![vh(&id_a, 0.9), vh(&id_b, 0.7)]);
let (storage, keyword, embedder) = build_ctx(rec.clone());
let ctx =
QueryContext::new(Query::text("q"), lunaris_core::Scope::dev(), embedder, storage, keyword);
let root = Vector::new("chunks", 30).rerank(Arc::new(DroppingReranker));
let res = root.retrieve(&ctx).await;
let err = res.expect_err("dropping reranker MUST surface OperatorFailed");
let msg = format!("{err}");
assert!(msg.contains("reranker returned"), "must mention doc-count mismatch; got: {msg}");
}
#[tokio::test]
async fn rerank_preserves_degraded_flag_through_rerank_pass() {
struct DegradedSource(Vec<u8>);
#[async_trait]
impl Retriever for DegradedSource {
async fn retrieve(&self, _ctx: &QueryContext) -> Result<Vec<RawHit>, LunarisError> {
Ok(vec![RawHit {
id: self.0.clone(),
score: 0.5,
rerank_applied: false,
degraded: true,
metadata: json!({}),
source_op: SourceOp::Vector,
}])
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
}
let rec = Arc::new(RecordingStorage::new());
let id_a = seed_chunk(&rec, "alpha");
let (storage, keyword, embedder) = build_ctx(rec.clone());
let ctx =
QueryContext::new(Query::text("q"), lunaris_core::Scope::dev(), embedder, storage, keyword);
let upstream: Box<dyn Retriever> = Box::new(DegradedSource(id_a.clone()));
let root = lunaris_retrieve::rerank(upstream, Arc::new(NoopReranker));
let raw = root.retrieve(&ctx).await.unwrap();
assert_eq!(raw.len(), 1);
assert!(raw[0].degraded, "rerank MUST preserve upstream degraded=true");
}
struct FixedScoreReranker(f32);
#[async_trait]
impl Reranker for FixedScoreReranker {
async fn rerank(
&self,
_query: &str,
mut docs: Vec<RerankCandidate>,
) -> Result<Vec<RerankCandidate>, LunarisError> {
for d in &mut docs {
d.score = self.0;
}
Ok(docs)
}
fn applies(&self) -> bool {
true
}
}
#[tokio::test]
async fn rerank_min_score_gate_drops_all_hits_below_threshold() {
let rec = Arc::new(RecordingStorage::new());
let id_a = seed_chunk(&rec, "alpha document text");
let id_b = seed_chunk(&rec, "beta document text");
rec.set_vector_hits(vec![vh(&id_a, 0.9), vh(&id_b, 0.7)]);
let (storage, keyword, embedder) = build_ctx(rec.clone());
let ctx =
QueryContext::new(Query::text("q"), lunaris_core::Scope::dev(), embedder, storage, keyword);
let root =
Vector::new("chunks", 30).rerank(Arc::new(FixedScoreReranker(0.2))).with_min_score(0.5);
let raw = root.retrieve(&ctx).await.unwrap();
assert!(
raw.is_empty(),
"hits scoring 0.2 must be dropped by a 0.5 threshold (abstention), got {} hits",
raw.len()
);
}
#[tokio::test]
async fn rerank_min_score_gate_keeps_hits_at_or_above_threshold() {
let rec = Arc::new(RecordingStorage::new());
let id_a = seed_chunk(&rec, "alpha document text");
let id_b = seed_chunk(&rec, "beta document text");
rec.set_vector_hits(vec![vh(&id_a, 0.9), vh(&id_b, 0.7)]);
let (storage, keyword, embedder) = build_ctx(rec.clone());
let ctx =
QueryContext::new(Query::text("q"), lunaris_core::Scope::dev(), embedder, storage, keyword);
let root =
Vector::new("chunks", 30).rerank(Arc::new(FixedScoreReranker(0.5))).with_min_score(0.5);
let raw = root.retrieve(&ctx).await.unwrap();
assert_eq!(raw.len(), 2, "hits scoring exactly at the threshold must be KEPT (>=, inclusive)");
}
#[tokio::test]
async fn rerank_min_score_gate_default_none_preserves_behavior() {
let rec = Arc::new(RecordingStorage::new());
let id_a = seed_chunk(&rec, "alpha document text");
rec.set_vector_hits(vec![vh(&id_a, 0.9)]);
let (storage, keyword, embedder) = build_ctx(rec.clone());
let ctx =
QueryContext::new(Query::text("q"), lunaris_core::Scope::dev(), embedder, storage, keyword);
let root = Vector::new("chunks", 30).rerank(Arc::new(FixedScoreReranker(0.0)));
let raw = root.retrieve(&ctx).await.unwrap();
assert_eq!(raw.len(), 1, "min_score=None (default) must never drop hits");
}
#[tokio::test]
async fn rerank_min_score_gate_skipped_when_reranker_does_not_apply() {
let rec = Arc::new(RecordingStorage::new());
let id_a = seed_chunk(&rec, "alpha document text");
let id_b = seed_chunk(&rec, "beta document text");
rec.set_vector_hits(vec![vh(&id_a, 0.9), vh(&id_b, 0.7)]);
let (storage, keyword, embedder) = build_ctx(rec.clone());
let ctx =
QueryContext::new(Query::text("q"), lunaris_core::Scope::dev(), embedder, storage, keyword);
let root = Vector::new("chunks", 30).rerank(Arc::new(NoopReranker)).with_min_score(0.99);
let raw = root.retrieve(&ctx).await.unwrap();
assert_eq!(
raw.len(),
2,
"min_score gate MUST be skipped on the NoopReranker (applies()==false) fallback path"
);
}
#[tokio::test]
async fn rerank_min_score_gate_preserves_count_validation_contract() {
struct DroppingReranker;
#[async_trait]
impl Reranker for DroppingReranker {
async fn rerank(
&self,
_q: &str,
mut docs: Vec<RerankCandidate>,
) -> Result<Vec<RerankCandidate>, LunarisError> {
docs.pop();
for d in &mut docs {
d.score = 0.9; }
Ok(docs)
}
fn applies(&self) -> bool {
true
}
}
let rec = Arc::new(RecordingStorage::new());
let id_a = seed_chunk(&rec, "alpha");
let id_b = seed_chunk(&rec, "beta");
rec.set_vector_hits(vec![vh(&id_a, 0.9), vh(&id_b, 0.7)]);
let (storage, keyword, embedder) = build_ctx(rec.clone());
let ctx =
QueryContext::new(Query::text("q"), lunaris_core::Scope::dev(), embedder, storage, keyword);
let root = Vector::new("chunks", 30).rerank(Arc::new(DroppingReranker)).with_min_score(0.1);
let res = root.retrieve(&ctx).await;
let err = res.expect_err("count mismatch must surface even with a gate configured");
assert!(format!("{err}").contains("reranker returned"));
}