use std::collections::HashMap;
use std::fmt;
use std::hash::Hash;
use std::marker::PhantomData;
use std::pin::Pin;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use std::time::Instant;
use dashmap::DashMap;
use futures_core::Stream;
use rand::rngs::SmallRng;
use rand::{Rng, RngExt};
use tower::Service;
use tower::discover::{Change, Discover};
use tower::limit::ConcurrencyLimit;
use tower::ready_cache::ReadyCache;
use super::types::{LlmRequest, LlmRequestKind, LlmResponse};
use crate::client::BoxFuture;
use crate::error::{LiterLlmError, Result};
#[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct Weight(u32);
impl Weight {
pub const ZERO: Weight = Weight(0);
pub const ONE: Weight = Weight(1);
pub const MAX: Weight = Weight(u32::MAX);
#[must_use]
pub fn from_f64(f: f64) -> Self {
if f.is_nan() || f < 0.0 {
Self::ZERO
} else if f.is_infinite() {
Self::MAX
} else {
#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
let w = f.round().min(f64::from(u32::MAX)) as u32;
Self(w)
}
}
#[must_use]
pub fn as_u32(self) -> u32 {
self.0
}
}
impl Default for Weight {
fn default() -> Self {
Self::ONE
}
}
impl fmt::Display for Weight {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.0)
}
}
#[derive(Clone)]
#[cfg_attr(alef, alef(skip))]
pub enum RoutingStrategy {
RoundRobin,
Fallback,
LatencyBased,
CostBased,
WeightedRandom {
weights: Vec<Weight>,
},
Semantic(Arc<dyn super::route_classify::RouteClassifier>),
}
impl std::fmt::Debug for RoutingStrategy {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::RoundRobin => write!(f, "RoundRobin"),
Self::Fallback => write!(f, "Fallback"),
Self::LatencyBased => write!(f, "LatencyBased"),
Self::CostBased => write!(f, "CostBased"),
Self::WeightedRandom { weights } => f.debug_struct("WeightedRandom").field("weights", weights).finish(),
Self::Semantic(_) => write!(f, "Semantic(…)"),
}
}
}
#[derive(Debug)]
struct DeploymentMetrics {
latency_ema: f64,
request_count: u64,
}
impl Default for DeploymentMetrics {
fn default() -> Self {
Self {
latency_ema: 0.0,
request_count: 0,
}
}
}
impl DeploymentMetrics {
fn record_latency(&mut self, latency_secs: f64) {
const ALPHA: f64 = 0.3;
if self.request_count == 0 {
self.latency_ema = latency_secs;
} else {
self.latency_ema = ALPHA * latency_secs + (1.0 - ALPHA) * self.latency_ema;
}
self.request_count += 1;
}
}
#[cfg_attr(alef, alef(skip))]
pub struct RouterState {
metrics: Arc<DashMap<usize, DeploymentMetrics>>,
}
impl RouterState {
fn new() -> Self {
Self {
metrics: Arc::new(DashMap::new()),
}
}
}
impl Clone for RouterState {
fn clone(&self) -> Self {
Self {
metrics: Arc::clone(&self.metrics),
}
}
}
#[cfg_attr(alef, alef(skip))]
pub struct Router<S> {
deployments: Vec<S>,
strategy: RoutingStrategy,
counter: Arc<AtomicUsize>,
state: RouterState,
deployment_models: Vec<String>,
weighted_random_rng: Arc<Mutex<SmallRng>>,
}
impl<S> Router<S> {
pub fn new(deployments: Vec<S>, strategy: RoutingStrategy) -> Result<Self> {
if deployments.is_empty() {
return Err(LiterLlmError::BadRequest {
message: "Router requires at least one deployment".into(),
status: 400,
});
}
if let RoutingStrategy::WeightedRandom { ref weights } = strategy {
if weights.len() != deployments.len() {
return Err(LiterLlmError::BadRequest {
message: format!(
"WeightedRandom: weights length ({}) must match deployments length ({})",
weights.len(),
deployments.len()
),
status: 400,
});
}
let total: u64 = weights.iter().map(|w| u64::from(w.as_u32())).sum();
if total == 0 {
return Err(LiterLlmError::BadRequest {
message: "WeightedRandom: total weight must be positive".into(),
status: 400,
});
}
}
let deployment_models = (0..deployments.len()).map(|i| i.to_string()).collect();
Ok(Self {
deployments,
strategy,
counter: Arc::new(AtomicUsize::new(0)),
state: RouterState::new(),
deployment_models,
weighted_random_rng: Arc::new(Mutex::new(rand::make_rng::<SmallRng>())),
})
}
#[must_use]
pub fn with_deployment_models(mut self, models: Vec<String>) -> Self {
let expected = self.deployments.len();
if models.len() != expected {
tracing::warn!(
expected,
got = models.len(),
"Router::with_deployment_models: length mismatch; missing entries keep the positional-index fallback"
);
}
for (i, model) in models.into_iter().take(expected).enumerate() {
self.deployment_models[i] = model;
}
self
}
}
impl<S: Clone> Clone for Router<S> {
fn clone(&self) -> Self {
Self {
deployments: self.deployments.clone(),
strategy: self.strategy.clone(),
counter: Arc::clone(&self.counter),
state: self.state.clone(),
deployment_models: self.deployment_models.clone(),
weighted_random_rng: Arc::clone(&self.weighted_random_rng),
}
}
}
impl<S> Service<LlmRequest> for Router<S>
where
S: Service<LlmRequest, Response = LlmResponse, Error = LiterLlmError> + Clone + Send + 'static,
S::Future: Send + 'static,
{
type Response = LlmResponse;
type Error = LiterLlmError;
type Future = BoxFuture<'static, Result<LlmResponse>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<()>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, req: LlmRequest) -> Self::Future {
match &self.strategy {
RoutingStrategy::RoundRobin => self.call_round_robin(req),
RoutingStrategy::Fallback => self.call_fallback(req),
RoutingStrategy::LatencyBased => self.call_latency_based(req),
RoutingStrategy::CostBased => self.call_cost_based(req),
RoutingStrategy::WeightedRandom { weights } => self.call_weighted_random(weights, req),
RoutingStrategy::Semantic(classifier) => self.call_semantic(classifier, req),
}
}
}
impl<S> Router<S>
where
S: Service<LlmRequest, Response = LlmResponse, Error = LiterLlmError> + Clone + Send + 'static,
S::Future: Send + 'static,
{
fn call_round_robin(&self, req: LlmRequest) -> BoxFuture<'static, Result<LlmResponse>> {
let idx = self.counter.fetch_add(1, Ordering::Relaxed) % self.deployments.len();
let mut svc = self.deployments[idx].clone();
Box::pin(async move { svc.call(req).await })
}
fn call_fallback(&self, req: LlmRequest) -> BoxFuture<'static, Result<LlmResponse>> {
let deployments = self.deployments.clone();
Box::pin(async move {
let mut last_err: Option<LiterLlmError> = None;
for mut svc in deployments {
match svc.call(req.clone()).await {
Ok(resp) => return Ok(resp),
Err(e) if e.is_transient() => {
tracing::warn!(
error = %e,
"deployment failed with transient error; trying next deployment"
);
last_err = Some(e);
}
Err(e) => return Err(e),
}
}
Err(last_err.unwrap_or(LiterLlmError::ServerError {
message: "all deployments failed".into(),
status: 500,
}))
})
}
fn call_latency_based(&self, req: LlmRequest) -> BoxFuture<'static, Result<LlmResponse>> {
let state = self.state.clone();
let n = self.deployments.len();
let mut best_idx = 0;
let mut best_ema = f64::MAX;
for i in 0..n {
let ema = state.metrics.get(&i).map_or(0.0, |m| m.latency_ema);
if ema < best_ema {
best_ema = ema;
best_idx = i;
}
}
let mut svc = self.deployments[best_idx].clone();
let idx = best_idx;
Box::pin(async move {
let start = Instant::now();
let result = svc.call(req).await;
let latency = start.elapsed().as_secs_f64();
state.metrics.entry(idx).or_default().record_latency(latency);
result
})
}
fn call_cost_based(&self, req: LlmRequest) -> BoxFuture<'static, Result<LlmResponse>> {
let model = req.model().map(ToOwned::to_owned);
let deployments = self.deployments.clone();
Box::pin(async move {
let mut last_err: Option<LiterLlmError> = None;
for mut svc in deployments {
match svc.call(req.clone()).await {
Ok(resp) => {
log_estimated_cost(model.as_deref(), &resp);
return Ok(resp);
}
Err(e) if e.is_transient() => {
last_err = Some(e);
}
Err(e) => return Err(e),
}
}
Err(last_err.unwrap_or(LiterLlmError::ServerError {
message: "all deployments failed".into(),
status: 500,
}))
})
}
fn call_weighted_random(&self, weights: &[Weight], req: LlmRequest) -> BoxFuture<'static, Result<LlmResponse>> {
let idx = {
let mut rng = self
.weighted_random_rng
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
weighted_random_select(weights, &mut *rng)
};
let mut svc = self.deployments[idx].clone();
Box::pin(async move { svc.call(req).await })
}
fn call_semantic(
&self,
classifier: &Arc<dyn super::route_classify::RouteClassifier>,
req: LlmRequest,
) -> BoxFuture<'static, Result<LlmResponse>> {
use super::route_classify::ClassifyContext;
let classifier = Arc::clone(classifier);
let deployments = self.deployments.clone();
let counter = Arc::clone(&self.counter);
let deployment_models = self.deployment_models.clone();
let (prompt, system_prompt) = extract_semantic_prompt(&req);
Box::pin(async move {
let meta: HashMap<String, String> = HashMap::new();
let ctx = ClassifyContext {
prompt: &prompt,
system_prompt: system_prompt.as_deref(),
metadata: &meta,
available_models: &deployment_models,
};
let verdict = classifier.classify(&ctx).await;
let idx = resolve_semantic_index(verdict.as_deref(), &deployment_models, &counter);
deployments[idx].clone().call(req).await
})
}
}
fn log_estimated_cost(model: Option<&str>, resp: &LlmResponse) {
if let (Some(model_name), Some(usage)) = (model, resp.usage())
&& let Some(cost) = crate::cost::completion_cost(model_name, usage.prompt_tokens, usage.completion_tokens)
{
tracing::debug!(model = %model_name, cost_usd = cost, "cost-based routing: estimated cost");
}
}
fn extract_semantic_prompt(req: &LlmRequest) -> (String, Option<String>) {
let LlmRequestKind::Chat(r) = &req.kind else {
return (String::new(), None);
};
let prompt = r
.messages
.iter()
.rev()
.find_map(|m| {
if let crate::types::Message::User(u) = m {
match &u.content {
crate::types::UserContent::Text(t) => Some(t.clone()),
crate::types::UserContent::Parts(_) => None,
}
} else {
None
}
})
.unwrap_or_default();
let system = r.messages.iter().find_map(|m| {
if let crate::types::Message::System(s) = m {
s.content.as_text()
} else {
None
}
});
(prompt, system)
}
fn resolve_semantic_index(verdict: Option<&str>, deployment_models: &[String], counter: &AtomicUsize) -> usize {
verdict
.and_then(|model_id| deployment_models.iter().position(|m| m == model_id))
.unwrap_or_else(|| counter.fetch_add(1, Ordering::Relaxed) % deployment_models.len())
}
fn weighted_random_select(weights: &[Weight], rng: &mut impl Rng) -> usize {
let total: u64 = weights.iter().map(|w| u64::from(w.as_u32())).sum();
if total == 0 {
return 0;
}
let threshold = rng.random_range(0..total);
let mut cumulative: u64 = 0;
for (i, w) in weights.iter().enumerate() {
cumulative += u64::from(w.as_u32());
if threshold < cumulative {
return i;
}
}
weights.len() - 1
}
pub trait UpstreamDiscover: Discover<Key = String> + Unpin + Send {}
impl<D> UpstreamDiscover for D where D: Discover<Key = String> + Unpin + Send {}
#[derive(Debug, thiserror::Error)]
#[cfg_attr(alef, alef(skip))]
pub enum RouterError {
#[error("discovery error (code 2001): {source}")]
Discover {
source: tower::BoxError,
code: u32,
},
#[error("no ready upstream available (code 2002)")]
NoReadyUpstream {
code: u32,
},
}
impl RouterError {
#[must_use]
pub fn code(&self) -> u32 {
match self {
Self::Discover { code, .. } | Self::NoReadyUpstream { code } => *code,
}
}
}
impl From<RouterError> for LiterLlmError {
fn from(e: RouterError) -> Self {
LiterLlmError::ServerError {
message: e.to_string(),
status: 503,
}
}
}
#[cfg_attr(alef, alef(skip))]
pub struct StaticDiscover<S> {
keys: std::collections::VecDeque<String>,
services: std::collections::VecDeque<S>,
}
impl<S> StaticDiscover<S> {
pub fn new(services: impl IntoIterator<Item = (String, S)>) -> Self {
let (keys, services): (std::collections::VecDeque<_>, std::collections::VecDeque<_>) =
services.into_iter().unzip();
Self { keys, services }
}
}
impl<S: Unpin> Unpin for StaticDiscover<S> {}
impl<S: Unpin> Stream for StaticDiscover<S> {
type Item = std::result::Result<Change<String, S>, std::convert::Infallible>;
fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
match (self.keys.pop_front(), self.services.pop_front()) {
(Some(key), Some(svc)) => Poll::Ready(Some(Ok(Change::Insert(key, svc)))),
_ => Poll::Ready(None),
}
}
}
pub const DEFAULT_CONCURRENCY_LIMIT: usize = 256;
#[derive(Debug, Clone)]
#[cfg_attr(alef, alef(skip))]
pub struct ProviderConfig {
pub concurrency_limit: usize,
}
impl Default for ProviderConfig {
fn default() -> Self {
Self {
concurrency_limit: DEFAULT_CONCURRENCY_LIMIT,
}
}
}
#[cfg_attr(alef, alef(skip))]
pub struct DynamicRouter<D>
where
D: Discover<Key = String>,
D::Service: Service<LlmRequest, Response = LlmResponse, Error = LiterLlmError>,
{
discover: D,
services: ReadyCache<String, ConcurrencyLimit<D::Service>, LlmRequest>,
provider_configs: HashMap<String, ProviderConfig>,
counter: AtomicUsize,
_marker: PhantomData<LlmRequest>,
}
impl<D> fmt::Debug for DynamicRouter<D>
where
D: Discover<Key = String> + fmt::Debug,
D::Service: Service<LlmRequest, Response = LlmResponse, Error = LiterLlmError> + fmt::Debug,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("DynamicRouter")
.field("discover", &self.discover)
.finish_non_exhaustive()
}
}
impl<D> DynamicRouter<D>
where
D: Discover<Key = String> + Unpin,
D::Service: Service<LlmRequest, Response = LlmResponse, Error = LiterLlmError> + Send + Unpin + 'static,
<D::Service as Service<LlmRequest>>::Future: Send + 'static,
D::Error: Into<tower::BoxError>,
{
pub fn new(discover: D) -> Self {
Self {
discover,
services: ReadyCache::default(),
provider_configs: HashMap::new(),
counter: AtomicUsize::new(0),
_marker: PhantomData,
}
}
pub fn with_provider_config(mut self, key: impl Into<String>, config: ProviderConfig) -> Self {
self.provider_configs.insert(key.into(), config);
self
}
#[must_use]
pub fn len(&self) -> usize {
self.services.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.services.is_empty()
}
fn update_from_discover(&mut self, cx: &mut Context<'_>) -> std::result::Result<(), RouterError> {
loop {
match Pin::new(&mut self.discover).poll_discover(cx) {
Poll::Pending => return Ok(()),
Poll::Ready(None) => return Ok(()),
Poll::Ready(Some(Err(e))) => {
return Err(RouterError::Discover {
source: e.into(),
code: 2001,
});
}
Poll::Ready(Some(Ok(Change::Insert(key, svc)))) => {
let limit = self
.provider_configs
.get(&key)
.map_or(DEFAULT_CONCURRENCY_LIMIT, |c| c.concurrency_limit);
tracing::debug!(provider = %key, concurrency_limit = limit, "discovered new upstream");
self.services.push(key, ConcurrencyLimit::new(svc, limit));
}
Poll::Ready(Some(Ok(Change::Remove(key)))) => {
tracing::debug!(provider = %key, "upstream removed from discovery");
self.services.evict(&key);
}
}
}
}
}
fn next_ready_index(counter: &AtomicUsize, ready_len: usize) -> usize {
counter.fetch_add(1, Ordering::Relaxed) % ready_len
}
impl<D> Service<LlmRequest> for DynamicRouter<D>
where
D: Discover<Key = String> + Unpin + Send,
D::Service: Service<LlmRequest, Response = LlmResponse, Error = LiterLlmError> + Send + Unpin + 'static,
<D::Service as Service<LlmRequest>>::Future: Send + 'static,
D::Error: Into<tower::BoxError>,
{
type Response = LlmResponse;
type Error = LiterLlmError;
type Future = BoxFuture<'static, Result<LlmResponse>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<()>> {
if let Err(e) = self.update_from_discover(cx) {
return Poll::Ready(Err(e.into()));
}
let _ = self.services.poll_pending(cx);
if self.services.ready_len() > 0 {
Poll::Ready(Ok(()))
} else {
Poll::Pending
}
}
fn call(&mut self, req: LlmRequest) -> Self::Future {
let ready_len = self.services.ready_len();
if ready_len == 0 {
return Box::pin(async { Err(RouterError::NoReadyUpstream { code: 2002 }.into()) });
}
let index = next_ready_index(&self.counter, ready_len);
let fut = self.services.call_ready_index(index, req);
Box::pin(fut)
}
}
#[cfg(test)]
mod tests {
use std::collections::VecDeque;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
use futures_core::Stream;
use rand::SeedableRng;
use super::*;
use crate::tower::service::LlmService;
use crate::tower::tests_common::{MockClient, chat_req};
use crate::tower::types::LlmRequest;
#[test]
fn weight_clamps_nan_to_zero() {
assert_eq!(Weight::from_f64(f64::NAN).as_u32(), 0);
}
#[test]
fn weight_clamps_negative_to_zero() {
assert_eq!(Weight::from_f64(-1.0).as_u32(), 0);
assert_eq!(Weight::from_f64(-f64::INFINITY).as_u32(), 0);
}
#[test]
fn weight_clamps_inf_to_max() {
assert_eq!(Weight::from_f64(f64::INFINITY).as_u32(), u32::MAX);
}
#[test]
fn weight_rounds_normal_values() {
assert_eq!(Weight::from_f64(1.0).as_u32(), 1);
assert_eq!(Weight::from_f64(1.4).as_u32(), 1);
assert_eq!(Weight::from_f64(1.5).as_u32(), 2);
assert_eq!(Weight::from_f64(100.0).as_u32(), 100);
}
#[test]
fn weight_default_is_one() {
assert_eq!(Weight::default().as_u32(), 1);
}
#[test]
fn weighted_random_selects_proportionally() {
let weights = vec![Weight(1), Weight(2), Weight(3)];
let mut rng = SmallRng::seed_from_u64(0xF00D);
let mut counts = [0usize; 3];
for _ in 0..600u64 {
let idx = weighted_random_select(&weights, &mut rng);
assert!(idx < 3, "index {idx} out of range");
counts[idx] += 1;
}
for (i, &count) in counts.iter().enumerate() {
assert!(count > 0, "index {i} was never selected (counts: {counts:?})");
}
}
#[test]
fn weighted_random_rapid_calls_do_not_all_collapse_to_one_index() {
let weights = vec![Weight(1), Weight(2), Weight(3)];
let mut rng = SmallRng::seed_from_u64(0xC0FFEE);
let mut seen = std::collections::HashSet::new();
for _ in 0..500u64 {
seen.insert(weighted_random_select(&weights, &mut rng));
}
assert!(
seen.len() > 1,
"500 rapid back-to-back calls all returned the same index — entropy source is not advancing per call"
);
}
#[tokio::test]
async fn latency_based_routes_to_fastest() {
let deployments: Vec<LlmService<MockClient>> =
vec![LlmService::new(MockClient::ok()), LlmService::new(MockClient::ok())];
let mut router = Router::new(deployments, RoutingStrategy::LatencyBased).expect("non-empty deployments");
let resp = router.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(resp.is_ok());
let resp = router.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(resp.is_ok());
}
#[tokio::test]
async fn cost_based_falls_through_on_transient_error() {
let deployments: Vec<LlmService<MockClient>> = vec![
LlmService::new(MockClient::failing_service_unavailable()),
LlmService::new(MockClient::ok()),
];
let mut router = Router::new(deployments, RoutingStrategy::CostBased).expect("non-empty deployments");
let resp = router.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(resp.is_ok(), "should fall through to second deployment");
}
#[tokio::test]
async fn weighted_random_selects_valid_deployment() {
let deployments: Vec<LlmService<MockClient>> = vec![
LlmService::new(MockClient::ok()),
LlmService::new(MockClient::ok()),
LlmService::new(MockClient::ok()),
];
let mut router = Router::new(
deployments,
RoutingStrategy::WeightedRandom {
weights: vec![Weight(1), Weight(2), Weight(3)],
},
)
.expect("non-empty deployments");
for _ in 0..20 {
let resp = router.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(resp.is_ok());
}
}
#[tokio::test]
async fn weighted_random_rejects_mismatched_weights() {
let deployments: Vec<LlmService<MockClient>> =
vec![LlmService::new(MockClient::ok()), LlmService::new(MockClient::ok())];
let result = Router::new(
deployments,
RoutingStrategy::WeightedRandom {
weights: vec![Weight(1)],
},
);
assert!(result.is_err());
}
#[tokio::test]
async fn weighted_random_rejects_zero_total_weight() {
let deployments: Vec<LlmService<MockClient>> = vec![LlmService::new(MockClient::ok())];
let result = Router::new(
deployments,
RoutingStrategy::WeightedRandom {
weights: vec![Weight::ZERO],
},
);
assert!(result.is_err());
}
#[test]
fn weighted_random_select_returns_valid_index() {
let weights = vec![Weight(1), Weight(2), Weight(3)];
let mut rng = SmallRng::seed_from_u64(0xABCD);
for _ in 0..100 {
let idx = weighted_random_select(&weights, &mut rng);
assert!(idx < weights.len());
}
}
#[test]
fn deployment_metrics_ema_updates() {
let mut m = DeploymentMetrics::default();
m.record_latency(1.0);
assert!(
(m.latency_ema - 1.0).abs() < 1e-9,
"first sample should set EMA directly"
);
m.record_latency(0.0);
assert!(
(m.latency_ema - 0.7).abs() < 1e-9,
"EMA should be 0.7 after second sample"
);
}
struct VecDiscover {
items: VecDeque<std::result::Result<Change<String, LlmService<MockClient>>, std::convert::Infallible>>,
}
impl VecDiscover {
fn new(services: Vec<(String, LlmService<MockClient>)>) -> Self {
Self {
items: services.into_iter().map(|(k, v)| Ok(Change::Insert(k, v))).collect(),
}
}
}
impl Stream for VecDiscover {
type Item = std::result::Result<Change<String, LlmService<MockClient>>, std::convert::Infallible>;
fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Poll::Ready(self.items.pop_front())
}
}
impl Unpin for VecDiscover {}
#[tokio::test]
async fn dynamic_router_warms_ready_cache() {
let discover = VecDiscover::new(vec![
("openai".into(), LlmService::new(MockClient::ok())),
("anthropic".into(), LlmService::new(MockClient::ok())),
]);
let mut router = DynamicRouter::new(discover);
futures_util::future::poll_fn(|cx| match router.poll_ready(cx) {
Poll::Ready(Ok(())) => Poll::Ready(()),
Poll::Ready(Err(e)) => panic!("unexpected error: {e}"),
Poll::Pending => Poll::Pending,
})
.await;
assert!(!router.is_empty(), "at least one upstream should be ready");
}
#[tokio::test]
async fn dynamic_router_evicts_stale() {
struct InsertThenRemoveDiscover {
step: usize,
}
impl Stream for InsertThenRemoveDiscover {
type Item = std::result::Result<Change<String, LlmService<MockClient>>, std::convert::Infallible>;
fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let step = self.step;
self.step += 1;
match step {
0 => Poll::Ready(Some(Ok(Change::Insert(
"openai".into(),
LlmService::new(MockClient::ok()),
)))),
1 => Poll::Ready(Some(Ok(Change::Remove("openai".into())))),
_ => Poll::Ready(None),
}
}
}
impl Unpin for InsertThenRemoveDiscover {}
let discover = InsertThenRemoveDiscover { step: 0 };
let mut router = DynamicRouter::new(discover);
let mut noop_cx = std::task::Context::from_waker(futures_util::task::noop_waker_ref());
let _ = router.poll_ready(&mut noop_cx);
assert_eq!(router.len(), 0, "evicted service should be removed");
}
#[tokio::test]
async fn concurrency_limit_rejects_at_max() {
#[derive(Clone)]
struct BlockingService {
call_count: Arc<AtomicUsize>,
}
impl Service<LlmRequest> for BlockingService {
type Response = LlmResponse;
type Error = LiterLlmError;
type Future = BoxFuture<'static, Result<LlmResponse>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<()>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, _req: LlmRequest) -> Self::Future {
self.call_count.fetch_add(1, AtomicOrdering::SeqCst);
Box::pin(std::future::pending())
}
}
let counter = Arc::new(AtomicUsize::new(0));
let inner = BlockingService {
call_count: Arc::clone(&counter),
};
let mut limited = ConcurrencyLimit::new(inner, 1);
assert!(
futures_util::future::poll_fn(|cx| limited.poll_ready(cx)).await.is_ok(),
"first poll_ready should be ok"
);
let _held_fut = limited.call(LlmRequest::ListModels());
let mut noop_cx = std::task::Context::from_waker(futures_util::task::noop_waker_ref());
let poll = limited.poll_ready(&mut noop_cx);
assert!(
poll.is_pending(),
"second poll_ready should be Pending when limit=1 and one request is in-flight"
);
}
#[tokio::test]
async fn static_discover_yields_all_services() {
let mut discover = StaticDiscover::new(vec![
("a".to_owned(), LlmService::new(MockClient::ok())),
("b".to_owned(), LlmService::new(MockClient::ok())),
]);
let mut noop_cx = std::task::Context::from_waker(futures_util::task::noop_waker_ref());
let first = Pin::new(&mut discover).poll_next(&mut noop_cx);
assert!(matches!(first, Poll::Ready(Some(Ok(Change::Insert(ref k, _)))) if k == "a"));
let second = Pin::new(&mut discover).poll_next(&mut noop_cx);
assert!(matches!(second, Poll::Ready(Some(Ok(Change::Insert(ref k, _)))) if k == "b"));
let third = Pin::new(&mut discover).poll_next(&mut noop_cx);
assert!(matches!(third, Poll::Ready(None)));
}
struct FixedIndexClassifier {
target: String,
}
impl crate::tower::route_classify::RouteClassifier for FixedIndexClassifier {
fn classify<'a>(
&'a self,
_ctx: &'a crate::tower::route_classify::ClassifyContext<'a>,
) -> Pin<Box<dyn std::future::Future<Output = Option<String>> + Send + 'a>> {
let target = self.target.clone();
Box::pin(async move { Some(target) })
}
}
struct DeferringClassifier;
impl crate::tower::route_classify::RouteClassifier for DeferringClassifier {
fn classify<'a>(
&'a self,
_ctx: &'a crate::tower::route_classify::ClassifyContext<'a>,
) -> Pin<Box<dyn std::future::Future<Output = Option<String>> + Send + 'a>> {
Box::pin(async move { None })
}
}
#[tokio::test]
async fn router_semantic_strategy_uses_classifier() {
let deployments: Vec<LlmService<MockClient>> = vec![
LlmService::new(MockClient::failing_rate_limited()),
LlmService::new(MockClient::ok()),
];
let classifier = Arc::new(FixedIndexClassifier { target: "1".into() });
let mut router = Router::new(deployments, RoutingStrategy::Semantic(classifier)).expect("valid router");
let resp = router.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(resp.is_ok(), "classifier should have routed to the ok deployment");
}
#[tokio::test]
async fn router_semantic_strategy_fallback_to_round_robin_when_classifier_defers() {
let deployments: Vec<LlmService<MockClient>> =
vec![LlmService::new(MockClient::ok()), LlmService::new(MockClient::ok())];
let classifier = Arc::new(DeferringClassifier);
let mut router = Router::new(deployments, RoutingStrategy::Semantic(classifier)).expect("valid router");
let resp = router.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(resp.is_ok(), "fallback round-robin should handle the request");
}
struct RecordingClassifier {
seen: Arc<std::sync::Mutex<Vec<String>>>,
}
impl crate::tower::route_classify::RouteClassifier for RecordingClassifier {
fn classify<'a>(
&'a self,
ctx: &'a crate::tower::route_classify::ClassifyContext<'a>,
) -> Pin<Box<dyn std::future::Future<Output = Option<String>> + Send + 'a>> {
*self.seen.lock().expect("mutex not poisoned") = ctx.available_models.to_vec();
Box::pin(async move { None })
}
}
#[tokio::test]
async fn router_semantic_strategy_passes_real_model_ids_to_classifier() {
let deployments: Vec<LlmService<MockClient>> =
vec![LlmService::new(MockClient::ok()), LlmService::new(MockClient::ok())];
let seen = Arc::new(std::sync::Mutex::new(Vec::new()));
let classifier = Arc::new(RecordingClassifier {
seen: Arc::clone(&seen),
});
let mut router = Router::new(deployments, RoutingStrategy::Semantic(classifier))
.expect("valid router")
.with_deployment_models(vec!["gpt-4o".into(), "claude-3-5-sonnet".into()]);
let resp = router.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(resp.is_ok());
assert_eq!(
*seen.lock().expect("mutex not poisoned"),
vec!["gpt-4o".to_string(), "claude-3-5-sonnet".to_string()],
"classifier must see real model IDs, not positional index placeholders"
);
}
#[tokio::test]
async fn router_semantic_strategy_routes_by_real_model_id() {
let deployments: Vec<LlmService<MockClient>> = vec![
LlmService::new(MockClient::failing_rate_limited()),
LlmService::new(MockClient::ok()),
];
let classifier = Arc::new(FixedIndexClassifier {
target: "claude-3-5-sonnet".into(),
});
let mut router = Router::new(deployments, RoutingStrategy::Semantic(classifier))
.expect("valid router")
.with_deployment_models(vec!["gpt-4o".into(), "claude-3-5-sonnet".into()]);
let resp = router.call(LlmRequest::Chat(chat_req("gpt-4"))).await;
assert!(
resp.is_ok(),
"verdict 'claude-3-5-sonnet' should resolve to deployment 1 by name"
);
}
#[test]
fn next_ready_index_rotates_through_ready_set() {
let counter = AtomicUsize::new(0);
assert_eq!(next_ready_index(&counter, 3), 0);
assert_eq!(
next_ready_index(&counter, 3),
1,
"second call must not be pinned to index 0"
);
assert_eq!(next_ready_index(&counter, 3), 2);
assert_eq!(
next_ready_index(&counter, 3),
0,
"wraps back around after a full rotation"
);
}
}