use nodedb_types::DatabaseId;
use nodedb_types::Surrogate;
use nodedb_types::TenantId;
use serde::{Deserialize, Serialize};
use crate::bridge::scan_filter::ScanFilter;
use crate::control::state::SharedState;
#[derive(Serialize, Deserialize, zerompk::ToMessagePack, zerompk::FromMessagePack)]
#[msgpack(map)]
struct Hit {
id: u32,
distance: f32,
#[serde(skip_serializing_if = "Option::is_none")]
doc_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
body: Option<Vec<u8>>,
}
fn apply_rls_filter(hits: &mut Vec<Hit>, rls_filters: &[u8], top_k: usize) {
if rls_filters.is_empty() {
return;
}
let filters: Vec<ScanFilter> = match zerompk::from_msgpack(rls_filters) {
Ok(f) => f,
Err(_) => {
tracing::warn!("RLS filter decode failed at CP boundary — denying all hits");
hits.clear();
return;
}
};
hits.retain(|h| match h.body.as_deref() {
Some(body) => filters.iter().all(|f| f.matches_binary(body)),
None => false,
});
if hits.len() > top_k {
hits.truncate(top_k);
}
for h in hits.iter_mut() {
h.body = None;
}
}
pub(super) fn resolve_surrogate_pk(
state: &SharedState,
database_id: DatabaseId,
tenant_id: TenantId,
collection: &str,
surrogate: Surrogate,
) -> Option<String> {
let catalog = state.credentials.catalog();
let pk_bytes = catalog
.get_pk_for_surrogate(database_id, tenant_id, collection, surrogate)
.ok()??;
String::from_utf8(pk_bytes).ok()
}
pub fn translate_vector_search_payload(
payload: &[u8],
state: &SharedState,
database_id: DatabaseId,
tenant_id: TenantId,
collection: &str,
rls_filters: &[u8],
top_k: usize,
) -> Vec<u8> {
if payload.is_empty() {
return payload.to_vec();
}
let first = payload[0];
if first == b'[' || first == b'{' || first == b'"' {
return payload.to_vec();
}
let mut hits: Vec<Hit> = match zerompk::from_msgpack(payload) {
Ok(h) => h,
Err(_) => return payload.to_vec(),
};
apply_rls_filter(&mut hits, rls_filters, top_k);
for hit in &mut hits {
if hit.doc_id.is_some() {
continue;
}
if let Some(pk) = resolve_surrogate_pk(
state,
database_id,
tenant_id,
collection,
Surrogate::new(hit.id),
) {
hit.doc_id = Some(pk);
}
}
use std::collections::BTreeMap;
let flattened: Vec<BTreeMap<String, serde_json::Value>> = hits
.iter()
.map(|h| {
let mut obj: BTreeMap<String, serde_json::Value> = BTreeMap::new();
obj.insert("distance".into(), serde_json::json!(h.distance));
if let Some(ref body) = h.body
&& let Ok(map) = zerompk::from_msgpack::<
std::collections::HashMap<String, nodedb_types::Value>,
>(body)
{
for (k, v) in map {
if let Ok(j) = serde_json::to_value(&v) {
obj.insert(k, j);
}
}
}
if !obj.contains_key("id") {
if let Some(ref doc) = h.doc_id {
obj.insert("id".into(), serde_json::json!(doc));
} else {
obj.insert("id".into(), serde_json::json!(h.id));
}
}
obj.insert("_surrogate".into(), serde_json::json!(h.id));
obj
})
.collect();
if let Ok(s) = sonic_rs::to_string(&flattened) {
return s.into_bytes();
}
zerompk::to_msgpack_vec(&hits).unwrap_or_else(|_| payload.to_vec())
}