use async_trait::async_trait;
use nodedb_types::Surrogate;
use crate::types::{DatabaseId, TenantId, VShardId};
pub struct VectorInsertParams {
pub collection: String,
pub vector: Vec<f32>,
pub dim: usize,
pub field_name: String,
pub surrogate: Surrogate,
pub pk_bytes: Option<Vec<u8>>,
}
#[async_trait]
pub trait VectorDispatcher: Send + Sync {
async fn dispatch_insert(
&self,
tenant_id: TenantId,
vshard: VShardId,
params: VectorInsertParams,
provenance: nodedb_types::sync::wire::SyncProvenance,
) -> crate::Result<Vec<u8>>;
async fn dispatch_delete(
&self,
tenant_id: TenantId,
vshard: VShardId,
collection: String,
surrogate: Surrogate,
field_name: String,
provenance: nodedb_types::sync::wire::SyncProvenance,
) -> crate::Result<Vec<u8>>;
fn assign_surrogate(
&self,
database_id: DatabaseId,
tenant_id: TenantId,
collection: &str,
doc_id: &str,
) -> crate::Result<Surrogate>;
}
pub struct SharedStateVectorDispatcher<'a> {
pub shared: &'a crate::control::state::SharedState,
}
#[async_trait]
impl<'a> VectorDispatcher for SharedStateVectorDispatcher<'a> {
async fn dispatch_insert(
&self,
tenant_id: TenantId,
vshard: VShardId,
params: VectorInsertParams,
provenance: nodedb_types::sync::wire::SyncProvenance,
) -> crate::Result<Vec<u8>> {
use crate::bridge::envelope::PhysicalPlan;
use crate::control::server::wal_dispatch::{VectorPutWalArgs, wal_append_vector_put};
use nodedb_physical::physical_plan::VectorOp;
let prov = provenance;
wal_append_vector_put(
&self.shared.wal,
tenant_id,
vshard,
DatabaseId::DEFAULT,
VectorPutWalArgs {
collection: ¶ms.collection,
vector: ¶ms.vector,
dim: params.dim,
field_name: ¶ms.field_name,
surrogate: params.surrogate,
provenance: Some(&prov),
},
)?;
let plan = PhysicalPlan::Vector(VectorOp::Insert {
collection: params.collection,
vector: params.vector,
dim: params.dim,
field_name: params.field_name,
surrogate: params.surrogate,
pk_bytes: params.pk_bytes,
provenance: Some(prov),
});
super::raft_dispatch::dispatch_sync_payload(self.shared, tenant_id, vshard, plan).await
}
async fn dispatch_delete(
&self,
tenant_id: TenantId,
vshard: VShardId,
collection: String,
surrogate: Surrogate,
field_name: String,
provenance: nodedb_types::sync::wire::SyncProvenance,
) -> crate::Result<Vec<u8>> {
use crate::bridge::envelope::PhysicalPlan;
use crate::control::server::wal_dispatch::{
VectorDeleteWalArgs, wal_append_vector_delete_by_surrogate,
};
use nodedb_physical::physical_plan::VectorOp;
let prov = provenance;
wal_append_vector_delete_by_surrogate(
&self.shared.wal,
tenant_id,
vshard,
DatabaseId::DEFAULT,
VectorDeleteWalArgs {
collection: &collection,
surrogate,
field_name: &field_name,
provenance: Some(&prov),
},
)?;
let plan = PhysicalPlan::Vector(VectorOp::DeleteBySurrogate {
collection,
surrogate,
field_name,
provenance: Some(prov),
});
super::raft_dispatch::dispatch_sync_payload(self.shared, tenant_id, vshard, plan).await
}
fn assign_surrogate(
&self,
database_id: DatabaseId,
tenant_id: TenantId,
collection: &str,
doc_id: &str,
) -> crate::Result<Surrogate> {
self.shared
.surrogate_assigner
.assign(database_id, tenant_id, collection, doc_id.as_bytes())
}
}
pub struct NoOpVectorDispatcher;
#[async_trait]
impl VectorDispatcher for NoOpVectorDispatcher {
async fn dispatch_insert(
&self,
_tenant_id: TenantId,
_vshard: VShardId,
_params: VectorInsertParams,
_provenance: nodedb_types::sync::wire::SyncProvenance,
) -> crate::Result<Vec<u8>> {
Err(super::raft_dispatch::noop_dispatch_error("vector insert"))
}
async fn dispatch_delete(
&self,
_tenant_id: TenantId,
_vshard: VShardId,
_collection: String,
_surrogate: Surrogate,
_field_name: String,
_provenance: nodedb_types::sync::wire::SyncProvenance,
) -> crate::Result<Vec<u8>> {
Err(super::raft_dispatch::noop_dispatch_error("vector delete"))
}
fn assign_surrogate(
&self,
_database_id: DatabaseId,
_tenant_id: TenantId,
_collection: &str,
_doc_id: &str,
) -> crate::Result<Surrogate> {
Ok(Surrogate::ZERO)
}
}
#[cfg(test)]
mod tests {
use std::sync::{Arc, Mutex};
use super::super::session::SyncSession;
use super::super::wire::*;
use super::*;
type MockCallLog = Arc<Mutex<Vec<(TenantId, String, String)>>>;
struct MockDispatcher {
insert_calls: MockCallLog,
delete_calls: MockCallLog,
result: crate::Result<()>,
}
impl MockDispatcher {
fn ok() -> (Self, MockCallLog, MockCallLog) {
let inserts = Arc::new(Mutex::new(Vec::new()));
let deletes = Arc::new(Mutex::new(Vec::new()));
(
Self {
insert_calls: inserts.clone(),
delete_calls: deletes.clone(),
result: Ok(()),
},
inserts,
deletes,
)
}
fn err() -> Self {
Self {
insert_calls: Arc::new(Mutex::new(Vec::new())),
delete_calls: Arc::new(Mutex::new(Vec::new())),
result: Err(crate::Error::Internal {
detail: "mock failure".to_string(),
}),
}
}
}
#[async_trait]
impl VectorDispatcher for MockDispatcher {
async fn dispatch_insert(
&self,
tenant_id: TenantId,
_vshard: VShardId,
params: VectorInsertParams,
provenance: nodedb_types::sync::wire::SyncProvenance,
) -> crate::Result<Vec<u8>> {
let seq = provenance.seq;
self.insert_calls
.lock()
.unwrap()
.push((tenant_id, params.collection, String::new()));
super::super::test_support::mock_applied_ack(&self.result, seq)
}
async fn dispatch_delete(
&self,
tenant_id: TenantId,
_vshard: VShardId,
collection: String,
_surrogate: Surrogate,
_field_name: String,
provenance: nodedb_types::sync::wire::SyncProvenance,
) -> crate::Result<Vec<u8>> {
let seq = provenance.seq;
self.delete_calls
.lock()
.unwrap()
.push((tenant_id, collection, String::new()));
super::super::test_support::mock_applied_ack(&self.result, seq)
}
fn assign_surrogate(
&self,
_database_id: DatabaseId,
_tenant_id: TenantId,
_collection: &str,
_doc_id: &str,
) -> crate::Result<Surrogate> {
Ok(Surrogate::ZERO)
}
}
fn make_session() -> SyncSession {
SyncSession::new("test-vector-session".to_string())
}
fn make_insert_msg(collection: &str, id: &str, vector: Vec<f32>) -> VectorInsertMsg {
let dim = vector.len();
VectorInsertMsg {
lite_id: "lite-test".to_string(),
collection: collection.to_string(),
id: id.to_string(),
vector,
dim,
field_name: String::new(),
batch_id: 1,
producer_id: 0,
epoch: 0,
seq: 0,
}
}
fn make_delete_msg(collection: &str, id: &str) -> VectorDeleteMsg {
VectorDeleteMsg {
lite_id: "lite-test".to_string(),
collection: collection.to_string(),
id: id.to_string(),
field_name: String::new(),
batch_id: 2,
producer_id: 0,
epoch: 0,
seq: 0,
}
}
#[tokio::test]
async fn unauthenticated_insert_returns_rejection() {
let mut session = make_session();
let (mock, inserts, _) = MockDispatcher::ok();
let msg = make_insert_msg("vecs", "v1", vec![1.0, 0.0, 0.0]);
let frame = session.handle_vector_insert(&msg, &mock).await;
assert!(frame.is_some());
let ack: VectorInsertAckMsg = frame.unwrap().decode_body().unwrap();
assert!(!ack.accepted);
assert!(inserts.lock().unwrap().is_empty());
}
#[tokio::test]
async fn authenticated_insert_dispatches_and_acks() {
let mut session = make_session();
session.authenticated = true;
let (mock, inserts, _) = MockDispatcher::ok();
let msg = make_insert_msg("vecs", "v1", vec![1.0, 0.0, 0.0]);
let frame = session.handle_vector_insert(&msg, &mock).await;
assert!(frame.is_some());
let ack: VectorInsertAckMsg = frame.unwrap().decode_body().unwrap();
assert!(ack.accepted);
assert_eq!(ack.id, "v1");
assert_eq!(inserts.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn insert_dimension_mismatch_rejects() {
let mut session = make_session();
session.authenticated = true;
let (mock, _, _) = MockDispatcher::ok();
let mut msg = make_insert_msg("vecs", "v1", vec![1.0, 0.0, 0.0]);
msg.dim = 5;
let frame = session.handle_vector_insert(&msg, &mock).await;
let ack: VectorInsertAckMsg = frame.unwrap().decode_body().unwrap();
assert!(!ack.accepted);
assert!(ack.reject_reason.unwrap().contains("dimension mismatch"));
}
#[tokio::test]
async fn insert_dispatch_failure_rejects() {
let mut session = make_session();
session.authenticated = true;
let mock = MockDispatcher::err();
let msg = make_insert_msg("vecs", "v1", vec![1.0, 0.0]);
let frame = session.handle_vector_insert(&msg, &mock).await;
let ack: VectorInsertAckMsg = frame.unwrap().decode_body().unwrap();
assert!(!ack.accepted);
assert!(ack.reject_reason.is_some());
}
#[tokio::test]
async fn authenticated_delete_dispatches_and_acks() {
let mut session = make_session();
session.authenticated = true;
let (mock, _, deletes) = MockDispatcher::ok();
let msg = make_delete_msg("vecs", "v1");
let frame = session.handle_vector_delete(&msg, &mock).await;
let ack: VectorDeleteAckMsg = frame.unwrap().decode_body().unwrap();
assert!(ack.accepted);
assert_eq!(ack.id, "v1");
assert_eq!(deletes.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn unauthenticated_delete_returns_rejection() {
let mut session = make_session();
let (mock, _, deletes) = MockDispatcher::ok();
let msg = make_delete_msg("vecs", "v1");
let frame = session.handle_vector_delete(&msg, &mock).await;
let ack: VectorDeleteAckMsg = frame.unwrap().decode_body().unwrap();
assert!(!ack.accepted);
assert!(deletes.lock().unwrap().is_empty());
}
}