use lattice_fann::{FannError, Network};
pub type AdapterId = String;
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum RouterError {
#[error("gate network error: {0}")]
Gate(#[from] FannError),
#[error("k={k} exceeds available adapter count {available}")]
KTooLarge {
k: usize,
available: usize,
},
#[error("k must be >= 1, got {k}")]
InvalidK {
k: usize,
},
#[error("context vector length {got} does not match gate input size {expected}")]
InputSizeMismatch {
expected: usize,
got: usize,
},
#[error(
"router: k={k} exceeds usable adapter count {usable} \
(gate produced fewer scores than adapters)"
)]
GateTooNarrow {
k: usize,
usable: usize,
},
#[error("duplicate adapter id in available set: {id}")]
DuplicateAdapterId {
id: String,
},
}
pub struct AdapterRouter {
gate: Network,
}
impl AdapterRouter {
pub fn new(gate: Network) -> Self {
Self { gate }
}
pub fn route(
&mut self,
context_vector: &[f32],
available: &[AdapterId],
k: usize,
) -> Result<Vec<(AdapterId, f32)>, RouterError> {
if k == 0 {
return Err(RouterError::InvalidK { k });
}
if k > available.len() {
return Err(RouterError::KTooLarge {
k,
available: available.len(),
});
}
let mut seen = std::collections::HashSet::new();
for id in available {
if !seen.insert(id.as_str()) {
return Err(RouterError::DuplicateAdapterId { id: id.clone() });
}
}
let expected_input = self.gate.num_inputs();
if context_vector.len() != expected_input {
return Err(RouterError::InputSizeMismatch {
expected: expected_input,
got: context_vector.len(),
});
}
let scores = self.gate.forward(context_vector)?;
let n = available.len().min(scores.len());
if k > n {
return Err(RouterError::GateTooNarrow { k, usable: n });
}
let mut indexed: Vec<(usize, f32)> = scores[..n].iter().copied().enumerate().collect();
indexed.select_nth_unstable_by(k - 1, |a, b| {
b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal)
});
indexed[..k].sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
let weight = 1.0 / k as f32;
let selected: Vec<(AdapterId, f32)> = indexed[..k]
.iter()
.map(|(idx, _)| (available[*idx].clone(), weight))
.collect();
Ok(selected)
}
pub fn input_size(&self) -> usize {
self.gate.num_inputs()
}
pub fn output_size(&self) -> usize {
self.gate.num_outputs()
}
}
#[cfg(test)]
mod tests {
use super::*;
use lattice_fann::{Activation, NetworkBuilder};
fn make_router(inputs: usize, outputs: usize) -> AdapterRouter {
let net = NetworkBuilder::new()
.input(inputs)
.output(outputs, Activation::Linear)
.build()
.unwrap();
AdapterRouter::new(net)
}
#[test]
fn route_returns_k_entries() {
let mut router = make_router(4, 6);
let available: Vec<AdapterId> = (0..6).map(|i| format!("adapter-{i}")).collect();
let ctx = vec![1.0f32; 4];
let result = router.route(&ctx, &available, 3).unwrap();
assert_eq!(result.len(), 3);
}
#[test]
fn route_weights_sum_to_one() {
let mut router = make_router(4, 4);
let available: Vec<AdapterId> = (0..4).map(|i| format!("a{i}")).collect();
let ctx = vec![1.0f32; 4];
let result = router.route(&ctx, &available, 4).unwrap();
let weight_sum: f32 = result.iter().map(|(_, w)| w).sum();
assert!((weight_sum - 1.0).abs() < 1e-6, "weights must sum to 1.0");
}
#[test]
fn route_uniform_weight() {
let k = 3usize;
let mut router = make_router(2, 5);
let available: Vec<AdapterId> = (0..5).map(|i| format!("a{i}")).collect();
let ctx = vec![0.5f32; 2];
let result = router.route(&ctx, &available, k).unwrap();
let expected_w = 1.0 / k as f32;
for (_, w) in &result {
assert!(
(w - expected_w).abs() < 1e-6,
"each weight must be 1/k={expected_w}"
);
}
}
#[test]
fn route_k_zero_errors() {
let mut router = make_router(2, 3);
let available: Vec<AdapterId> = vec!["a".into(), "b".into(), "c".into()];
assert!(router.route(&[1.0, 2.0], &available, 0).is_err());
}
#[test]
fn route_k_exceeds_available_errors() {
let mut router = make_router(2, 2);
let available: Vec<AdapterId> = vec!["a".into()];
assert!(router.route(&[1.0, 2.0], &available, 2).is_err());
}
#[test]
fn route_wrong_input_size_errors() {
let mut router = make_router(4, 2);
let available: Vec<AdapterId> = vec!["a".into(), "b".into()];
assert!(router.route(&[1.0, 2.0, 3.0], &available, 1).is_err());
}
#[test]
fn route_narrow_gate_returns_err() {
let mut router = make_router(2, 3);
let available: Vec<AdapterId> = (0..5).map(|i| format!("a{i}")).collect();
let ctx = vec![1.0f32; 2];
let result = router.route(&ctx, &available, 4);
assert!(
matches!(result, Err(RouterError::GateTooNarrow { k: 4, usable: 3 })),
"expected GateTooNarrow {{k:4, usable:3}}, got {result:?}"
);
}
#[test]
fn route_duplicate_adapter_ids_returns_err() {
let mut router = make_router(2, 3);
let available: Vec<AdapterId> = vec!["same".into(), "same".into(), "other".into()];
let ctx = vec![1.0f32; 2];
let result = router.route(&ctx, &available, 2);
assert!(
matches!(result, Err(RouterError::DuplicateAdapterId { .. })),
"duplicate adapter id should return DuplicateAdapterId error, got {result:?}"
);
}
}