use std::collections::HashMap;
use std::sync::{Arc, RwLock};
use std::time::{Duration, Instant};
use tokio::task::JoinError;
use uuid::Uuid;
use khive_runtime::{BackendId, KhiveRuntime, SearchHit};
use khive_score::DeterministicScore;
use khive_types::namespace::Namespace;
#[derive(Clone)]
pub struct BackendEntry {
pub id: BackendId,
pub runtime: Arc<KhiveRuntime>,
}
#[derive(Default)]
pub struct BackendRegistry {
backends: HashMap<String, BackendEntry>,
primary: Option<String>,
}
impl BackendRegistry {
pub fn new() -> Self {
Self::default()
}
pub fn register(&mut self, id: BackendId, runtime: Arc<KhiveRuntime>) -> bool {
let key = id.as_str().to_string();
if self.backends.contains_key(&key) {
return false;
}
if self.primary.is_none() {
self.primary = Some(key.clone());
}
self.backends.insert(key, BackendEntry { id, runtime });
true
}
pub fn get(&self, id: &BackendId) -> Option<&BackendEntry> {
self.backends.get(id.as_str())
}
pub fn primary(&self) -> Option<&BackendEntry> {
self.primary.as_deref().and_then(|k| self.backends.get(k))
}
pub fn iter(&self) -> impl Iterator<Item = &BackendEntry> {
self.backends.values()
}
pub fn len(&self) -> usize {
self.backends.len()
}
pub fn is_empty(&self) -> bool {
self.backends.is_empty()
}
pub fn ids(&self) -> Vec<BackendId> {
self.backends.keys().map(BackendId::new).collect()
}
}
const DEFAULT_LOCATOR_TTL: Duration = Duration::from_secs(300);
struct LocatorEntry {
backend_id: BackendId,
inserted_at: Instant,
}
pub struct LocatorCache {
entries: RwLock<HashMap<Uuid, LocatorEntry>>,
ttl: Duration,
}
impl LocatorCache {
pub fn with_ttl(ttl: Duration) -> Self {
Self {
entries: RwLock::new(HashMap::new()),
ttl,
}
}
pub fn new() -> Self {
Self::with_ttl(DEFAULT_LOCATOR_TTL)
}
pub fn get(&self, id: Uuid) -> Option<BackendId> {
let now = Instant::now();
{
let guard = self.entries.read().unwrap_or_else(|e| e.into_inner());
if let Some(entry) = guard.get(&id) {
if now.duration_since(entry.inserted_at) < self.ttl {
return Some(entry.backend_id.clone());
}
} else {
return None;
}
}
let mut guard = self.entries.write().unwrap_or_else(|e| e.into_inner());
if let Some(entry) = guard.get(&id) {
if now.duration_since(entry.inserted_at) < self.ttl {
return Some(entry.backend_id.clone());
}
}
guard.remove(&id);
None
}
pub fn remove(&self, id: Uuid) {
let mut guard = self.entries.write().unwrap_or_else(|e| e.into_inner());
guard.remove(&id);
}
pub fn insert(&self, id: Uuid, backend_id: BackendId) {
let mut guard = self.entries.write().unwrap_or_else(|e| e.into_inner());
guard.insert(
id,
LocatorEntry {
backend_id,
inserted_at: Instant::now(),
},
);
}
pub fn purge_expired(&self) {
let now = Instant::now();
let mut guard = self.entries.write().unwrap_or_else(|e| e.into_inner());
guard.retain(|_, entry| now.duration_since(entry.inserted_at) < self.ttl);
}
pub fn len(&self) -> usize {
let guard = self.entries.read().unwrap_or_else(|e| e.into_inner());
guard.len()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
impl Default for LocatorCache {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug)]
pub struct BackendSearchResult {
pub backend_id: BackendId,
pub hits: Vec<SearchHit>,
pub error: Option<String>,
}
pub struct SubstrateCoordinator {
registry: BackendRegistry,
locator: Arc<LocatorCache>,
#[cfg(test)]
fail_backend_id: Option<String>,
}
impl SubstrateCoordinator {
pub fn new(registry: BackendRegistry) -> Self {
Self {
registry,
locator: Arc::new(LocatorCache::new()),
#[cfg(test)]
fail_backend_id: None,
}
}
pub fn with_locator_ttl(registry: BackendRegistry, ttl: Duration) -> Self {
Self {
registry,
locator: Arc::new(LocatorCache::with_ttl(ttl)),
#[cfg(test)]
fail_backend_id: None,
}
}
pub fn single(runtime: Arc<KhiveRuntime>) -> Self {
let mut registry = BackendRegistry::new();
registry.register(BackendId::main(), runtime);
Self {
registry,
locator: Arc::new(LocatorCache::new()),
#[cfg(test)]
fail_backend_id: None,
}
}
#[cfg(test)]
pub fn with_failing_backend(mut self, backend_id: &str) -> Self {
self.fail_backend_id = Some(backend_id.to_string());
self
}
pub fn registry(&self) -> &BackendRegistry {
&self.registry
}
pub fn locator_cache(&self) -> &Arc<LocatorCache> {
&self.locator
}
pub fn primary_runtime(&self) -> Option<Arc<KhiveRuntime>> {
self.registry.primary().map(|e| Arc::clone(&e.runtime))
}
pub fn backend_ids(&self) -> Vec<BackendId> {
self.registry.ids()
}
pub fn backend_count(&self) -> usize {
self.registry.len()
}
pub fn is_single_backend(&self) -> bool {
self.registry.len() <= 1
}
pub async fn locate(&self, id: Uuid, namespace: &Namespace) -> Option<BackendId> {
if let Some(backend_id) = self.locator.get(id) {
return Some(backend_id);
}
let entries: Vec<(BackendId, Arc<KhiveRuntime>)> = self
.registry
.iter()
.map(|e| (e.id.clone(), Arc::clone(&e.runtime)))
.collect();
if entries.is_empty() {
return None;
}
if entries.len() == 1 {
let (backend_id, runtime) = &entries[0];
let token = match runtime.authorize(namespace.clone()) {
Ok(t) => t,
Err(e) => {
tracing::warn!(error = %e, "locate: authorization denied for namespace");
return None;
}
};
let ns_str = namespace.as_str().to_string();
let entity_ns = ns_str.clone();
let entity_owned = match runtime.entities(&token) {
Ok(store) => store
.get_entity(id)
.await
.ok()
.flatten()
.map(|e| e.namespace == entity_ns)
.unwrap_or(false),
Err(_) => false,
};
if entity_owned {
self.locator.insert(id, backend_id.clone());
return Some(backend_id.clone());
}
let note_owned = match runtime.notes(&token) {
Ok(store) => store
.get_note(id)
.await
.ok()
.flatten()
.map(|n| n.namespace == ns_str)
.unwrap_or(false),
Err(_) => false,
};
if note_owned {
self.locator.insert(id, backend_id.clone());
return Some(backend_id.clone());
}
return None;
}
let ns_clone = namespace.clone();
let locator = Arc::clone(&self.locator);
let mut handles = Vec::with_capacity(entries.len());
for (backend_id, runtime) in entries {
let ns = ns_clone.clone();
let locator = Arc::clone(&locator);
let handle = tokio::spawn(async move {
let token = match runtime.authorize(ns.clone()) {
Ok(t) => t,
Err(e) => {
tracing::warn!(error = %e, "locate: authorization denied for namespace");
return None;
}
};
let ns_str = ns.as_str().to_string();
if let Ok(store) = runtime.entities(&token) {
if let Ok(Some(entity)) = store.get_entity(id).await {
if entity.namespace == ns_str {
locator.insert(id, backend_id.clone());
return Some(backend_id);
}
}
}
if let Ok(store) = runtime.notes(&token) {
if let Ok(Some(note)) = store.get_note(id).await {
if note.namespace == ns_str {
locator.insert(id, backend_id.clone());
return Some(backend_id);
}
}
}
None
});
handles.push(handle);
}
let results: Vec<Result<Option<BackendId>, JoinError>> =
futures_util::future::join_all(handles).await;
for result in results {
if let Ok(Some(backend_id)) = result {
return Some(backend_id);
}
}
None
}
pub fn invalidate(&self, id: Uuid) {
self.locator.remove(id);
}
pub async fn fan_out_search(
&self,
query: &str,
namespace: &Namespace,
limit: u32,
) -> (Vec<SearchHit>, Vec<BackendSearchResult>) {
let entries: Vec<(BackendId, Arc<KhiveRuntime>)> = self
.registry
.iter()
.map(|e| (e.id.clone(), Arc::clone(&e.runtime)))
.collect();
if entries.is_empty() {
return (vec![], vec![]);
}
if entries.len() == 1 {
let (backend_id, runtime) = &entries[0];
let token = match runtime.authorize(namespace.clone()) {
Ok(t) => t,
Err(e) => {
tracing::warn!(error = %e, "fan_out_search: authorization denied for namespace");
let backend_result = BackendSearchResult {
backend_id: backend_id.clone(),
hits: vec![],
error: Some(e.to_string()),
};
return (vec![], vec![backend_result]);
}
};
match runtime
.hybrid_search(&token, query, None, limit, None, None)
.await
{
Ok(hits) => {
let backend_result = BackendSearchResult {
backend_id: backend_id.clone(),
hits: hits.clone(),
error: None,
};
return (hits, vec![backend_result]);
}
Err(e) => {
let backend_result = BackendSearchResult {
backend_id: backend_id.clone(),
hits: vec![],
error: Some(e.to_string()),
};
return (vec![], vec![backend_result]);
}
}
}
let query = query.to_string();
let ns = namespace.clone();
#[cfg(test)]
let fail_id: Option<String> = self.fail_backend_id.clone();
#[cfg(not(test))]
let fail_id: Option<String> = None;
let mut handles = Vec::with_capacity(entries.len());
for (backend_id, runtime) in entries {
let q = query.clone();
let ns = ns.clone();
let should_fail = fail_id
.as_deref()
.map(|id| id == backend_id.as_str())
.unwrap_or(false);
let handle = tokio::spawn(async move {
if should_fail {
return (
backend_id,
Err(khive_runtime::RuntimeError::Internal(
"injected failure".to_string(),
)),
);
}
let token = match runtime.authorize(ns) {
Ok(t) => t,
Err(e) => {
tracing::warn!(error = %e, "fan_out_search: authorization denied for namespace");
return (backend_id, Err(e));
}
};
let result = runtime
.hybrid_search(&token, &q, None, limit, None, None)
.await;
(backend_id, result)
});
handles.push(handle);
}
type BackendSearchOutcome = (
BackendId,
Result<Vec<SearchHit>, khive_runtime::RuntimeError>,
);
let join_results: Vec<Result<BackendSearchOutcome, JoinError>> =
futures_util::future::join_all(handles).await;
let mut per_backend: Vec<BackendSearchResult> = Vec::new();
let mut ranked_lists: Vec<Vec<SearchHit>> = Vec::new();
for join_result in join_results {
match join_result {
Ok((backend_id, Ok(hits))) => {
ranked_lists.push(hits.clone());
per_backend.push(BackendSearchResult {
backend_id,
hits,
error: None,
});
}
Ok((backend_id, Err(e))) => {
per_backend.push(BackendSearchResult {
backend_id,
hits: vec![],
error: Some(e.to_string()),
});
}
Err(join_err) => {
tracing::warn!(error = %join_err, "backend search task failed");
}
}
}
let merged = rrf_merge_hits(ranked_lists, limit as usize);
(merged, per_backend)
}
}
fn rrf_merge_hits(lists: Vec<Vec<SearchHit>>, limit: usize) -> Vec<SearchHit> {
const K: f64 = 60.0;
let mut scores: HashMap<Uuid, (f64, Option<String>, Option<String>)> = HashMap::new();
for list in &lists {
for (i, hit) in list.iter().enumerate() {
let rank = (i + 1) as f64;
let rrf = 1.0 / (K + rank);
let entry = scores.entry(hit.entity_id).or_insert((0.0, None, None));
entry.0 += rrf;
if entry.1.is_none() {
entry.1 = hit.title.clone();
}
if entry.2.is_none() {
entry.2 = hit.snippet.clone();
}
}
}
let mut merged: Vec<SearchHit> = scores
.into_iter()
.map(|(id, (score, title, snippet))| {
let det_score = DeterministicScore::from_f64(score);
SearchHit {
entity_id: id,
score: det_score,
source: khive_runtime::SearchSource::Both,
title,
snippet,
}
})
.collect();
merged.sort_by(|a, b| b.score.cmp(&a.score).then(a.entity_id.cmp(&b.entity_id)));
merged.truncate(limit);
merged
}
mod futures_util {
pub mod future {
pub async fn join_all<F: std::future::Future>(
futs: Vec<F>,
) -> Vec<<F as std::future::Future>::Output> {
let mut results = Vec::with_capacity(futs.len());
for fut in futs {
results.push(fut.await);
}
results
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use khive_runtime::KhiveRuntime;
fn memory_runtime() -> Arc<KhiveRuntime> {
Arc::new(KhiveRuntime::memory().expect("memory runtime"))
}
#[test]
fn single_coordinator_is_single_backend() {
let coord = SubstrateCoordinator::single(memory_runtime());
assert!(coord.is_single_backend());
assert_eq!(coord.backend_count(), 1);
assert_eq!(coord.backend_ids().len(), 1);
assert_eq!(coord.backend_ids()[0].as_str(), "main");
}
#[test]
fn registry_register_dedup() {
let mut reg = BackendRegistry::new();
let rt = memory_runtime();
assert!(reg.register(BackendId::new("main"), Arc::clone(&rt)));
assert!(!reg.register(BackendId::new("main"), Arc::clone(&rt)));
assert_eq!(reg.len(), 1);
}
#[test]
fn registry_primary_is_first_registered() {
let mut reg = BackendRegistry::new();
let rt1 = memory_runtime();
let rt2 = memory_runtime();
reg.register(BackendId::new("main"), rt1);
reg.register(BackendId::new("lore"), rt2);
assert_eq!(reg.primary().unwrap().id.as_str(), "main");
}
#[test]
fn multi_backend_coordinator_not_single() {
let mut registry = BackendRegistry::new();
registry.register(BackendId::new("main"), memory_runtime());
registry.register(BackendId::new("lore"), memory_runtime());
let coord = SubstrateCoordinator::new(registry);
assert!(!coord.is_single_backend());
assert_eq!(coord.backend_count(), 2);
}
#[test]
fn backend_id_display() {
let id = BackendId::new("archive");
assert_eq!(id.to_string(), "archive");
assert_eq!(id.as_str(), "archive");
}
#[test]
fn backend_id_main_constant() {
assert_eq!(BackendId::main().as_str(), BackendId::MAIN);
}
#[test]
fn locator_cache_miss_returns_none() {
let cache = LocatorCache::new();
let id = Uuid::new_v4();
assert!(cache.get(id).is_none());
}
#[test]
fn locator_cache_insert_then_get_returns_backend() {
let cache = LocatorCache::new();
let id = Uuid::new_v4();
cache.insert(id, BackendId::new("main"));
let result = cache.get(id);
assert!(result.is_some());
assert_eq!(result.unwrap().as_str(), "main");
}
#[test]
fn locator_cache_expired_entry_returns_none() {
let cache = LocatorCache::with_ttl(Duration::from_nanos(1));
let id = Uuid::new_v4();
cache.insert(id, BackendId::new("main"));
std::thread::sleep(Duration::from_micros(1));
assert!(cache.get(id).is_none());
}
#[test]
fn locator_cache_purge_removes_expired() {
let cache = LocatorCache::with_ttl(Duration::from_nanos(1));
for _ in 0..5 {
cache.insert(Uuid::new_v4(), BackendId::new("main"));
}
std::thread::sleep(Duration::from_micros(1));
cache.purge_expired();
assert_eq!(cache.len(), 0);
}
#[tokio::test]
async fn locator_cache_miss_then_hit() {
let coord = SubstrateCoordinator::single(memory_runtime());
let ns = Namespace::local();
let runtime = coord.primary_runtime().unwrap();
let token = runtime.authorize(ns.clone()).unwrap();
let entity = runtime
.create_entity(&token, "concept", None, "LoRA", None, None, vec![])
.await
.expect("create entity");
let first = coord.locate(entity.id, &ns).await;
assert!(
first.is_some(),
"locate should find the entity on first call"
);
assert_eq!(first.unwrap().as_str(), BackendId::MAIN);
assert_eq!(coord.locator_cache().len(), 1, "cache should be populated");
let second = coord.locate(entity.id, &ns).await;
assert!(second.is_some(), "second locate should hit cache");
}
#[tokio::test]
async fn locator_cache_returns_none_for_unknown_uuid() {
let coord = SubstrateCoordinator::single(memory_runtime());
let ns = Namespace::local();
let unknown = Uuid::new_v4();
let result = coord.locate(unknown, &ns).await;
assert!(result.is_none(), "unknown UUID should resolve to None");
}
#[tokio::test]
async fn fan_out_search_single_backend_returns_hits() {
let coord = SubstrateCoordinator::single(memory_runtime());
let ns = Namespace::local();
let runtime = coord.primary_runtime().unwrap();
let token = runtime.authorize(ns.clone()).unwrap();
runtime
.create_entity(
&token,
"concept",
None,
"FlashAttention",
Some("IO-aware exact attention"),
None,
vec![],
)
.await
.expect("create entity");
let (hits, per_backend) = coord.fan_out_search("FlashAttention", &ns, 10).await;
assert!(!hits.is_empty(), "should find the entity");
assert_eq!(per_backend.len(), 1, "single backend report");
assert!(per_backend[0].error.is_none(), "no error");
}
#[tokio::test]
async fn fan_out_search_two_backends_merged() {
let mut registry = BackendRegistry::new();
let rt_main = memory_runtime();
let rt_lore = memory_runtime();
registry.register(BackendId::new("main"), Arc::clone(&rt_main));
registry.register(BackendId::new("lore"), Arc::clone(&rt_lore));
let coord = SubstrateCoordinator::new(registry);
let ns = Namespace::local();
let tok_main = rt_main.authorize(ns.clone()).unwrap();
rt_main
.create_entity(
&tok_main,
"concept",
None,
"LoRA",
Some("Low-rank adaptation"),
None,
vec![],
)
.await
.expect("create on main");
let tok_lore = rt_lore.authorize(ns.clone()).unwrap();
rt_lore
.create_entity(
&tok_lore,
"concept",
None,
"QLoRA",
Some("Quantised LoRA"),
None,
vec![],
)
.await
.expect("create on lore");
let (merged_hits, per_backend) = coord.fan_out_search("LoRA", &ns, 10).await;
assert_eq!(per_backend.len(), 2, "both backends in report");
assert!(
!merged_hits.is_empty(),
"merged results should not be empty"
);
}
#[tokio::test]
async fn fan_out_search_empty_registry_returns_empty() {
let coord = SubstrateCoordinator::new(BackendRegistry::new());
let ns = Namespace::local();
let (hits, per_backend) = coord.fan_out_search("anything", &ns, 10).await;
assert!(hits.is_empty());
assert!(per_backend.is_empty());
}
#[tokio::test]
async fn fan_out_partial_failure_preserves_working_backend_hits() {
let rt_main = memory_runtime();
let rt_lore = memory_runtime();
let ns = Namespace::local();
let tok_lore = rt_lore.authorize(ns.clone()).unwrap();
rt_lore
.create_entity(
&tok_lore,
"concept",
None,
"PartialFailureProbe",
Some("probe entity for partial-failure test"),
None,
vec![],
)
.await
.expect("create entity on lore");
let mut registry = BackendRegistry::new();
registry.register(BackendId::new("main"), Arc::clone(&rt_main));
registry.register(BackendId::new("lore"), Arc::clone(&rt_lore));
let coord = SubstrateCoordinator::new(registry).with_failing_backend("main");
let (merged_hits, per_backend) = coord.fan_out_search("PartialFailureProbe", &ns, 10).await;
assert_eq!(
per_backend.len(),
2,
"both backends should appear in the report"
);
let main_result = per_backend
.iter()
.find(|r| r.backend_id.as_str() == "main")
.expect("main backend result must be present");
assert!(
main_result.error.is_some(),
"main backend should report an error"
);
assert!(
main_result.hits.is_empty(),
"main backend should have no hits"
);
let lore_result = per_backend
.iter()
.find(|r| r.backend_id.as_str() == "lore")
.expect("lore backend result must be present");
assert!(
lore_result.error.is_none(),
"lore backend should have no error"
);
assert!(
!merged_hits.is_empty(),
"merged hits must include results from the working backend"
);
}
#[tokio::test]
async fn locate_finds_note_uuid() {
let coord = SubstrateCoordinator::single(memory_runtime());
let ns = Namespace::local();
let runtime = coord.primary_runtime().unwrap();
let token = runtime.authorize(ns.clone()).unwrap();
let note = runtime
.create_note(
&token,
"observation",
Some("locate-note-regression"),
"content for locate regression test",
None,
None,
vec![],
)
.await
.expect("create note");
let backend = coord.locate(note.id, &ns).await;
assert!(backend.is_some(), "locate should find the note's backend");
assert_eq!(backend.unwrap().as_str(), BackendId::MAIN);
assert_eq!(
coord.locator_cache().len(),
1,
"cache should be populated for the note"
);
}
#[test]
fn locator_cache_get_evicts_expired_entry() {
let cache = LocatorCache::with_ttl(Duration::from_nanos(1));
let id = Uuid::new_v4();
cache.insert(id, BackendId::new("main"));
assert_eq!(cache.len(), 1, "entry inserted");
std::thread::sleep(Duration::from_micros(1));
assert!(cache.get(id).is_none(), "expired entry returns None");
assert_eq!(cache.len(), 0, "expired entry must be evicted from the map");
}
#[test]
fn locator_cache_remove_evicts_live_entry() {
let cache = LocatorCache::new();
let id = Uuid::new_v4();
cache.insert(id, BackendId::new("main"));
assert!(cache.get(id).is_some(), "entry live before remove");
cache.remove(id);
assert!(cache.get(id).is_none(), "entry gone after remove");
assert_eq!(cache.len(), 0, "map must be empty after remove");
}
#[tokio::test]
async fn invalidate_clears_locate_cache() {
let coord = SubstrateCoordinator::single(memory_runtime());
let ns = Namespace::local();
let runtime = coord.primary_runtime().unwrap();
let token = runtime.authorize(ns.clone()).unwrap();
let entity = runtime
.create_entity(
&token,
"concept",
None,
"InvalidateTest",
None,
None,
vec![],
)
.await
.expect("create entity");
coord.locate(entity.id, &ns).await;
assert_eq!(coord.locator_cache().len(), 1, "cache populated");
coord.invalidate(entity.id);
assert_eq!(
coord.locator_cache().len(),
0,
"cache cleared after invalidate"
);
let found_again = coord.locate(entity.id, &ns).await;
assert!(found_again.is_some(), "locate re-finds after cache clear");
}
}