use std::any::Any;
use std::collections::HashMap;
use std::sync::{Arc, Mutex, RwLock};
use crate::serve::kv_persist::spiller::KvCacheSpill;
use crate::serve::kv_persist::EngineBindable;
use crate::serve::quant_select::QuantType;
type FamilyKey = (String, &'static str);
pub trait FamilyHookFactory: Send + Sync {
fn try_construct(
&self,
engine_dyn: Arc<dyn Any + Send + Sync>,
) -> Option<(Arc<Mutex<dyn KvCacheSpill>>, Arc<dyn EngineBindable>)>;
}
pub struct KvPersistRegistry {
hooks: RwLock<HashMap<FamilyKey, Arc<dyn EngineBindable>>>,
factories: RwLock<HashMap<FamilyKey, Arc<dyn FamilyHookFactory>>>,
}
impl Default for KvPersistRegistry {
fn default() -> Self {
Self::new()
}
}
impl KvPersistRegistry {
pub fn new() -> Self {
Self {
hooks: RwLock::new(HashMap::new()),
factories: RwLock::new(HashMap::new()),
}
}
pub fn register(&self, repo: String, quant: QuantType, hook: Arc<dyn EngineBindable>) {
let mut g = self
.hooks
.write()
.expect("KvPersistRegistry::hooks RwLock poisoned");
g.insert((repo, quant.as_str()), hook);
}
pub fn unregister(&self, repo: &str, quant: QuantType) -> bool {
let mut g = self
.hooks
.write()
.expect("KvPersistRegistry::hooks RwLock poisoned");
let key = (repo.to_string(), quant.as_str());
g.remove(&key).is_some()
}
pub fn registered_count(&self) -> usize {
self.hooks
.read()
.expect("KvPersistRegistry::hooks RwLock poisoned")
.len()
}
pub fn bind_for(&self, repo: &str, quant: QuantType, engine_dyn: Arc<dyn Any + Send + Sync>) {
let g = self
.hooks
.read()
.expect("KvPersistRegistry::hooks RwLock poisoned");
let key = (repo.to_string(), quant.as_str());
if let Some(hook) = g.get(&key) {
hook.bind_engine(engine_dyn);
}
}
pub fn unbind_for(&self, repo: &str, quant: QuantType) {
let g = self
.hooks
.read()
.expect("KvPersistRegistry::hooks RwLock poisoned");
let key = (repo.to_string(), quant.as_str());
if let Some(hook) = g.get(&key) {
hook.unbind_engine();
}
}
pub fn register_factory(
&self,
repo: String,
quant: QuantType,
factory: Arc<dyn FamilyHookFactory>,
) {
let mut g = self
.factories
.write()
.expect("KvPersistRegistry::factories RwLock poisoned");
g.insert((repo, quant.as_str()), factory);
}
pub fn factory_count(&self) -> usize {
self.factories
.read()
.expect("KvPersistRegistry::factories RwLock poisoned")
.len()
}
pub fn contains_factory(&self, repo: &str, quant: QuantType) -> bool {
let g = self
.factories
.read()
.expect("KvPersistRegistry::factories RwLock poisoned");
let key = (repo.to_string(), quant.as_str());
g.contains_key(&key)
}
pub fn try_substitute_on_load(
&self,
repo: &str,
quant: QuantType,
engine_dyn: Arc<dyn Any + Send + Sync>,
) -> Option<Arc<Mutex<dyn KvCacheSpill>>> {
let key = (repo.to_string(), quant.as_str());
let factory = {
let g = self
.factories
.read()
.expect("KvPersistRegistry::factories RwLock poisoned");
g.get(&key).cloned()
};
let factory = factory?;
let (kv_hook, bindable_hook) = factory.try_construct(engine_dyn)?;
{
let mut g = self
.hooks
.write()
.expect("KvPersistRegistry::hooks RwLock poisoned");
g.insert(key, bindable_hook);
}
Some(kv_hook)
}
pub fn contains(&self, repo: &str, quant: QuantType) -> bool {
let g = self
.hooks
.read()
.expect("KvPersistRegistry::hooks RwLock poisoned");
let key = (repo.to_string(), quant.as_str());
g.contains_key(&key)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicU32, Ordering};
struct MockBindable {
binds: AtomicU32,
unbinds: AtomicU32,
last_bound_type: std::sync::Mutex<Option<std::any::TypeId>>,
}
impl MockBindable {
fn new() -> Arc<Self> {
Arc::new(Self {
binds: AtomicU32::new(0),
unbinds: AtomicU32::new(0),
last_bound_type: std::sync::Mutex::new(None),
})
}
fn bind_count(&self) -> u32 {
self.binds.load(Ordering::SeqCst)
}
fn unbind_count(&self) -> u32 {
self.unbinds.load(Ordering::SeqCst)
}
fn last_bound_type(&self) -> Option<std::any::TypeId> {
*self.last_bound_type.lock().unwrap()
}
}
impl EngineBindable for MockBindable {
fn bind_engine(&self, engine_dyn: Arc<dyn Any + Send + Sync>) {
self.binds.fetch_add(1, Ordering::SeqCst);
let id: std::any::TypeId = (*engine_dyn).type_id();
*self.last_bound_type.lock().unwrap() = Some(id);
}
fn unbind_engine(&self) {
self.unbinds.fetch_add(1, Ordering::SeqCst);
}
}
#[derive(Debug)]
struct DummyEnginePayload {
_marker: u32,
}
#[test]
fn new_registry_has_zero_hooks() {
let reg = KvPersistRegistry::new();
assert_eq!(reg.registered_count(), 0);
assert!(!reg.contains("acme/m1", QuantType::Q4_K_M));
}
#[test]
fn registry_register_then_bind_round_trip() {
let reg = KvPersistRegistry::new();
let mock = MockBindable::new();
reg.register(
"acme/m1".to_string(),
QuantType::Q4_K_M,
mock.clone() as Arc<dyn EngineBindable>,
);
assert_eq!(reg.registered_count(), 1);
assert!(reg.contains("acme/m1", QuantType::Q4_K_M));
let payload = Arc::new(DummyEnginePayload { _marker: 0xABCD });
let payload_dyn: Arc<dyn Any + Send + Sync> = payload.clone();
reg.bind_for("acme/m1", QuantType::Q4_K_M, payload_dyn);
assert_eq!(mock.bind_count(), 1, "bind_engine fired exactly once");
assert_eq!(mock.unbind_count(), 0);
assert_eq!(
mock.last_bound_type(),
Some(std::any::TypeId::of::<DummyEnginePayload>()),
"type-id round-trip survives Arc<dyn Any> erasure"
);
}
#[test]
fn registry_unbind_clears_engine_handle() {
let reg = KvPersistRegistry::new();
let mock = MockBindable::new();
reg.register(
"acme/m1".to_string(),
QuantType::Q4_K_M,
mock.clone() as Arc<dyn EngineBindable>,
);
let payload = Arc::new(DummyEnginePayload { _marker: 1 });
reg.bind_for(
"acme/m1",
QuantType::Q4_K_M,
payload as Arc<dyn Any + Send + Sync>,
);
assert_eq!(mock.bind_count(), 1);
reg.unbind_for("acme/m1", QuantType::Q4_K_M);
assert_eq!(mock.unbind_count(), 1);
}
#[test]
fn registry_bind_for_unknown_repo_quant_is_noop() {
let reg = KvPersistRegistry::new();
let mock = MockBindable::new();
reg.register(
"acme/known".to_string(),
QuantType::Q4_K_M,
mock.clone() as Arc<dyn EngineBindable>,
);
let payload = Arc::new(DummyEnginePayload { _marker: 2 });
reg.bind_for(
"acme/unknown",
QuantType::Q4_K_M,
payload.clone() as Arc<dyn Any + Send + Sync>,
);
reg.bind_for(
"acme/known",
QuantType::Q8_0,
payload as Arc<dyn Any + Send + Sync>,
);
assert_eq!(mock.bind_count(), 0, "no hook fired on unknown key");
reg.unbind_for("nobody/here", QuantType::Q4_K_M);
assert_eq!(mock.unbind_count(), 0);
}
#[test]
fn registry_re_register_overwrites_prior_hook() {
let reg = KvPersistRegistry::new();
let first = MockBindable::new();
let second = MockBindable::new();
reg.register(
"acme/m1".to_string(),
QuantType::Q4_K_M,
first.clone() as Arc<dyn EngineBindable>,
);
reg.register(
"acme/m1".to_string(),
QuantType::Q4_K_M,
second.clone() as Arc<dyn EngineBindable>,
);
assert_eq!(reg.registered_count(), 1, "no count growth on overwrite");
let payload = Arc::new(DummyEnginePayload { _marker: 3 });
reg.bind_for(
"acme/m1",
QuantType::Q4_K_M,
payload as Arc<dyn Any + Send + Sync>,
);
assert_eq!(first.bind_count(), 0, "first registration was overwritten");
assert_eq!(second.bind_count(), 1, "second registration is live");
}
#[test]
fn registry_unregister_idempotent() {
let reg = KvPersistRegistry::new();
let mock = MockBindable::new();
reg.register(
"acme/m1".to_string(),
QuantType::Q4_K_M,
mock as Arc<dyn EngineBindable>,
);
assert!(reg.unregister("acme/m1", QuantType::Q4_K_M));
assert_eq!(reg.registered_count(), 0);
assert!(!reg.unregister("acme/m1", QuantType::Q4_K_M));
assert!(!reg.unregister("nobody/here", QuantType::Q4_K_M));
}
#[test]
fn registry_distinct_quant_grows_table() {
let reg = KvPersistRegistry::new();
let m1 = MockBindable::new();
let m2 = MockBindable::new();
reg.register(
"acme/m1".to_string(),
QuantType::Q4_K_M,
m1.clone() as Arc<dyn EngineBindable>,
);
reg.register(
"acme/m1".to_string(),
QuantType::Q8_0,
m2.clone() as Arc<dyn EngineBindable>,
);
assert_eq!(reg.registered_count(), 2);
let payload = Arc::new(DummyEnginePayload { _marker: 4 });
reg.bind_for(
"acme/m1",
QuantType::Q8_0,
payload as Arc<dyn Any + Send + Sync>,
);
assert_eq!(m1.bind_count(), 0);
assert_eq!(m2.bind_count(), 1);
}
struct MockKvSpillForFactory;
impl crate::serve::kv_persist::spiller::KvCacheSpill for MockKvSpillForFactory {
fn block_alignment(&self) -> u32 {
crate::serve::kv_persist::format::BLOCK_TOKENS
}
fn snapshot_block(
&self,
_layer_rank: usize,
_range: std::ops::Range<u32>,
) -> Option<Vec<u8>> {
None
}
fn restore_block(
&mut self,
_layer_rank: usize,
_range: std::ops::Range<u32>,
_payload: &[u8],
) -> std::result::Result<(), crate::serve::multi_model::SpillErrorKind> {
Ok(())
}
}
#[derive(Debug)]
struct ExpectedEngineHandle {
marker: u32,
}
struct MockFactoryExpecting;
impl FamilyHookFactory for MockFactoryExpecting {
fn try_construct(
&self,
engine_dyn: Arc<dyn Any + Send + Sync>,
) -> Option<(Arc<Mutex<dyn KvCacheSpill>>, Arc<dyn EngineBindable>)> {
match engine_dyn.downcast::<ExpectedEngineHandle>() {
Ok(_handle) => {
let kv: Arc<Mutex<dyn KvCacheSpill>> =
Arc::new(Mutex::new(MockKvSpillForFactory));
let bindable: Arc<dyn EngineBindable> = MockBindable::new();
Some((kv, bindable))
}
Err(_) => None,
}
}
}
struct MockFactoryAlwaysNone;
impl FamilyHookFactory for MockFactoryAlwaysNone {
fn try_construct(
&self,
_engine_dyn: Arc<dyn Any + Send + Sync>,
) -> Option<(Arc<Mutex<dyn KvCacheSpill>>, Arc<dyn EngineBindable>)> {
None
}
}
#[test]
fn registry_register_factory_round_trip() {
let reg = KvPersistRegistry::new();
assert_eq!(reg.factory_count(), 0);
assert!(!reg.contains_factory("acme/m1", QuantType::Q4_K_M));
reg.register_factory(
"acme/m1".to_string(),
QuantType::Q4_K_M,
Arc::new(MockFactoryExpecting),
);
assert_eq!(reg.factory_count(), 1);
assert!(reg.contains_factory("acme/m1", QuantType::Q4_K_M));
assert_eq!(reg.registered_count(), 0);
}
#[test]
fn registry_try_substitute_on_load_succeeds_on_type_match() {
let reg = KvPersistRegistry::new();
let stub = MockBindable::new();
reg.register(
"acme/m1".to_string(),
QuantType::Q4_K_M,
stub.clone() as Arc<dyn EngineBindable>,
);
assert_eq!(reg.registered_count(), 1);
reg.register_factory(
"acme/m1".to_string(),
QuantType::Q4_K_M,
Arc::new(MockFactoryExpecting),
);
let payload = Arc::new(ExpectedEngineHandle { marker: 0xBEEF });
let payload_dyn: Arc<dyn Any + Send + Sync> = payload;
let result = reg.try_substitute_on_load("acme/m1", QuantType::Q4_K_M, payload_dyn);
assert!(result.is_some(), "factory yields a kv_hook on type match");
let next_payload = Arc::new(ExpectedEngineHandle { marker: 0x1234 });
reg.bind_for(
"acme/m1",
QuantType::Q4_K_M,
next_payload as Arc<dyn Any + Send + Sync>,
);
assert_eq!(
stub.bind_count(),
0,
"after substitution, the original stub no longer fires"
);
assert_eq!(reg.registered_count(), 1);
assert_eq!(reg.factory_count(), 1);
}
#[test]
fn registry_try_substitute_on_load_returns_none_on_type_mismatch() {
let reg = KvPersistRegistry::new();
let stub = MockBindable::new();
reg.register(
"acme/m1".to_string(),
QuantType::Q4_K_M,
stub.clone() as Arc<dyn EngineBindable>,
);
reg.register_factory(
"acme/m1".to_string(),
QuantType::Q4_K_M,
Arc::new(MockFactoryExpecting),
);
let wrong_payload = Arc::new(DummyEnginePayload { _marker: 0xFEED });
let result = reg.try_substitute_on_load(
"acme/m1",
QuantType::Q4_K_M,
wrong_payload as Arc<dyn Any + Send + Sync>,
);
assert!(result.is_none(), "type mismatch ⇒ no substitution");
let next_payload = Arc::new(DummyEnginePayload { _marker: 0xBABE });
reg.bind_for(
"acme/m1",
QuantType::Q4_K_M,
next_payload as Arc<dyn Any + Send + Sync>,
);
assert_eq!(
stub.bind_count(),
1,
"stub remains registered after non-matching factory"
);
}
#[test]
fn registry_try_substitute_on_load_no_factory_is_noop() {
let reg = KvPersistRegistry::new();
let stub = MockBindable::new();
reg.register(
"acme/m1".to_string(),
QuantType::Q4_K_M,
stub.clone() as Arc<dyn EngineBindable>,
);
let payload = Arc::new(ExpectedEngineHandle { marker: 0xC0DE });
let result = reg.try_substitute_on_load(
"acme/m1",
QuantType::Q4_K_M,
payload as Arc<dyn Any + Send + Sync>,
);
assert!(result.is_none(), "no factory ⇒ no substitution");
assert_eq!(reg.factory_count(), 0);
let payload2 = Arc::new(ExpectedEngineHandle { marker: 1 });
reg.bind_for(
"acme/m1",
QuantType::Q4_K_M,
payload2 as Arc<dyn Any + Send + Sync>,
);
assert_eq!(stub.bind_count(), 1);
}
#[test]
fn registry_factory_always_none_does_not_substitute() {
let reg = KvPersistRegistry::new();
let stub = MockBindable::new();
reg.register(
"acme/m1".to_string(),
QuantType::Q4_K_M,
stub.clone() as Arc<dyn EngineBindable>,
);
reg.register_factory(
"acme/m1".to_string(),
QuantType::Q4_K_M,
Arc::new(MockFactoryAlwaysNone),
);
let payload = Arc::new(ExpectedEngineHandle { marker: 5 });
let result = reg.try_substitute_on_load(
"acme/m1",
QuantType::Q4_K_M,
payload as Arc<dyn Any + Send + Sync>,
);
assert!(result.is_none());
let payload2 = Arc::new(ExpectedEngineHandle { marker: 6 });
reg.bind_for(
"acme/m1",
QuantType::Q4_K_M,
payload2 as Arc<dyn Any + Send + Sync>,
);
assert_eq!(stub.bind_count(), 1);
}
#[test]
fn registry_factory_re_register_overwrites_prior() {
let reg = KvPersistRegistry::new();
reg.register_factory(
"acme/m1".to_string(),
QuantType::Q4_K_M,
Arc::new(MockFactoryAlwaysNone),
);
reg.register_factory(
"acme/m1".to_string(),
QuantType::Q4_K_M,
Arc::new(MockFactoryExpecting),
);
assert_eq!(reg.factory_count(), 1, "no count growth on overwrite");
let stub = MockBindable::new();
reg.register(
"acme/m1".to_string(),
QuantType::Q4_K_M,
stub as Arc<dyn EngineBindable>,
);
let payload = Arc::new(ExpectedEngineHandle { marker: 7 });
let result = reg.try_substitute_on_load(
"acme/m1",
QuantType::Q4_K_M,
payload as Arc<dyn Any + Send + Sync>,
);
assert!(result.is_some(), "second-registered factory is live");
}
#[test]
fn registry_multi_threaded_bind_for_is_safe() {
use std::thread;
const N_THREADS: u32 = 8;
const N_ITERS: u32 = 32;
let reg = Arc::new(KvPersistRegistry::new());
let mock = MockBindable::new();
reg.register(
"acme/m1".to_string(),
QuantType::Q4_K_M,
mock.clone() as Arc<dyn EngineBindable>,
);
let handles: Vec<_> = (0..N_THREADS)
.map(|_| {
let reg = Arc::clone(®);
thread::spawn(move || {
for i in 0..N_ITERS {
let payload = Arc::new(DummyEnginePayload { _marker: i });
reg.bind_for(
"acme/m1",
QuantType::Q4_K_M,
payload as Arc<dyn Any + Send + Sync>,
);
}
})
})
.collect();
for h in handles {
h.join().expect("worker did not panic");
}
assert_eq!(mock.bind_count(), N_THREADS * N_ITERS);
}
}