repositories / smith
smith
There are many coding harnesses - but this one is fast
owned by admin
smith-ai/src/mux.rs
Raw//! Endpoint routing with bounded retry and failover (`SMH-SPEC-SPEC0001`,
//! Providers).
//!
//! Resolution produces an ordered plan; routing walks it. Rate limits and
//! overloaded endpoints fail over immediately, transient faults get a bounded
//! retry budget per endpoint, invalid requests and authentication faults fail
//! without trying anything else, and exhausting the plan surfaces the last
//! fault instead of looping.
use crate::models::Catalog;
use futures::StreamExt;
use smith::error::{ProviderFault, Result, SmithError};
use smith::provider::{ProviderRequest, ProviderStream, StreamFn};
use smith::tool::CancelHandle;
use std::collections::HashMap;
use std::sync::Arc;
/// One endpoint in a failover plan: a model id and its stream.
#[derive(Clone)]
pub struct Endpoint {
/// Model id used in requests to this endpoint.
pub model: String,
/// The endpoint's stream function.
pub stream: StreamFn,
}
/// Routes requests through a resolved failover plan.
#[derive(Clone)]
pub struct Router {
catalog: Catalog,
endpoints: Arc<HashMap<String, Endpoint>>,
/// Transient retries per endpoint before failing over.
transient_retries: u32,
/// Injected delay for retry backoff; tests pass a no-op.
sleeper: Arc<dyn Fn(u64) + Send + Sync>,
}
impl Router {
/// Build a router over a catalog and named endpoints.
///
/// # Errors
///
/// Returns a [`SmithError::Config`] fault when an endpoint model is not
/// in the catalog.
pub fn new(
catalog: Catalog,
endpoints: HashMap<String, Endpoint>,
transient_retries: u32,
) -> Result<Self> {
for model in endpoints.keys() {
catalog.resolve_plan(model)?;
}
Ok(Self {
catalog,
endpoints: Arc::new(endpoints),
transient_retries,
sleeper: Arc::new(|_| {}),
})
}
/// Replace the retry delay function; returns the router for chaining.
#[must_use]
pub fn with_sleeper(mut self, sleeper: Arc<dyn Fn(u64) + Send + Sync>) -> Self {
self.sleeper = sleeper;
self
}
/// Send `request` through the failover plan for its target.
///
/// The returned stream carries the first endpoint's normalized events;
/// faults that occur after a stream starts are terminal for the caller,
/// because a partially consumed response cannot be retried honestly.
///
/// # Errors
///
/// Returns the last provider fault when every endpoint failed, or the
/// first non-failover fault (invalid request, authentication).
pub fn route(
&self,
target: &str,
request: &ProviderRequest,
cancel: &CancelHandle,
) -> Result<ProviderStream> {
let plan = self.catalog.resolve_plan(target)?;
let mut last_fault: Option<SmithError> = None;
for entry in plan {
let Some(endpoint) = self.endpoints.get(&entry.id) else {
continue;
};
let mut attempt: u32 = 0;
let mut endpoint_failed = false;
while !endpoint_failed {
if cancel.is_cancelled() {
return Err(SmithError::Cancelled);
}
self.catalog.validate_request(&entry.id, request)?;
let mut routed = request.clone();
routed.model.clone_from(&endpoint.model);
let stream = (endpoint.stream)(routed, cancel.clone());
// Establishment faults surface as the first stream item;
// classify them before the caller consumes anything.
match first_item(stream) {
Ok(events) => return Ok(events),
Err(err) => match &err {
SmithError::Provider {
fault: ProviderFault::Transient { .. },
} if attempt < self.transient_retries => {
attempt += 1;
(self.sleeper)(50 * u64::from(attempt));
}
SmithError::Provider { fault } if fault.may_fail_over() => {
last_fault = Some(err);
endpoint_failed = true;
}
_ => return Err(err),
},
}
}
}
Err(last_fault.unwrap_or_else(|| SmithError::Provider {
fault: ProviderFault::Protocol {
message: format!("no endpoint serves '{target}'"),
},
}))
}
}
/// Split off the first stream item so establishment faults classify before
/// the caller consumes anything.
fn first_item(stream: ProviderStream) -> std::result::Result<ProviderStream, SmithError> {
let mut stream = stream;
let mut buffer: Vec<smith::provider::StreamItem> = Vec::new();
match futures::executor::block_on(stream.next()) {
Some(item) => {
if let Err(err) = &item {
return Err(err.clone());
}
buffer.push(item);
Ok(Box::pin(
futures::stream::iter(buffer).chain(ProviderStreamRest(stream)),
))
}
None => Ok(Box::pin(futures::stream::iter(buffer))),
}
}
/// Adapter that yields the remainder of a partially consumed stream.
struct ProviderStreamRest(ProviderStream);
impl futures::Stream for ProviderStreamRest {
type Item = smith::provider::StreamItem;
fn poll_next(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
std::pin::Pin::new(&mut self.0).poll_next(cx)
}
}
#[cfg(test)]
mod tests {
use super::*;
use futures::StreamExt;
use smith::id::MessageId;
use smith::message::{Message, Role};
use smith::stream::{StopReason, StreamEvent, Usage};
use std::sync::atomic::{AtomicUsize, Ordering};
fn catalog() -> Catalog {
Catalog::from_json(
r#"{
"models": {
"primary": {"id": "primary", "provider": "OpenAiCompatible", "capabilities": {
"context_window": 8192, "max_output_tokens": 1024, "input_images": false,
"tool_use": true, "thinking": false, "strict_tool_input": false, "cache_control": false}},
"backup": {"id": "backup", "provider": "OpenAiCompatible", "capabilities": {
"context_window": 8192, "max_output_tokens": 1024, "input_images": false,
"tool_use": true, "thinking": false, "strict_tool_input": false, "cache_control": false}}
},
"aliases": {},
"groups": { "plan": ["primary", "backup"] }
}"#,
)
.unwrap()
}
fn scripted(script: Vec<Result<StreamEvent>>) -> StreamFn {
let outcome = std::sync::Mutex::new(script);
Arc::new(move |_request, _cancel| {
let item = { outcome.lock().unwrap().remove(0) };
futures::stream::iter(vec![item]).boxed()
})
}
/// The endpoint `name`, serving model `name` from `script` one item per call.
fn endpoint(name: &str, script: Vec<Result<StreamEvent>>) -> (String, Endpoint) {
(
name.to_string(),
Endpoint {
model: name.to_string(),
stream: scripted(script),
},
)
}
fn request() -> ProviderRequest {
ProviderRequest::new("plan", vec![Message::with_text(Role::User, "hi")])
}
fn rate_limit() -> SmithError {
SmithError::Provider {
fault: ProviderFault::RateLimit {
retry_after_ms: None,
},
}
}
fn ok_stream() -> Vec<Result<StreamEvent>> {
vec![Ok(StreamEvent::stop(
StopReason::EndTurn,
Some(Usage::new()),
))]
}
fn drain(mut stream: ProviderStream) -> Result<StreamEvent> {
let _ = MessageId::new();
futures::executor::block_on(stream.next()).unwrap()
}
fn route_err(
router: &Router,
target: &str,
request: &ProviderRequest,
cancel: &CancelHandle,
) -> SmithError {
match router.route(target, request, cancel) {
Err(err) => err,
Ok(_) => unreachable!("routing must fail in this scenario"),
}
}
#[test]
fn rate_limits_fail_over_immediately() {
let primary_calls = Arc::new(AtomicUsize::new(0));
let primary_seen = Arc::clone(&primary_calls);
let primary: StreamFn = Arc::new(move |_request, _cancel| {
primary_seen.fetch_add(1, Ordering::SeqCst);
futures::stream::iter(vec![Err(rate_limit())]).boxed()
});
let endpoints = HashMap::from([
(
"primary".to_string(),
Endpoint {
model: "primary".to_string(),
stream: primary,
},
),
endpoint("backup", vec![text("backup")]),
]);
let router = Router::new(catalog(), endpoints, 2).unwrap();
let stream = router
.route("plan", &request(), &CancelHandle::new())
.unwrap();
assert!(matches!(drain(stream), Ok(StreamEvent::TextDelta { .. })));
// Rate limit fails over without retrying the same endpoint.
assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
}
fn transient() -> SmithError {
SmithError::Provider {
fault: ProviderFault::Transient {
message: "reset".to_string(),
},
}
}
fn text(served_by: &str) -> Result<StreamEvent> {
Ok(StreamEvent::text_delta(MessageId::new(), served_by))
}
#[test]
fn transient_faults_retry_bounded_then_fail_over() {
let authentication = || SmithError::Provider {
fault: ProviderFault::Authentication {
message: "denied".to_string(),
},
};
// (primary faults before it succeeds, outcome, backoff delays in ms)
let cases: [(Vec<SmithError>, Result<&str>, Vec<u64>); 4] = [
(vec![transient(), transient()], Ok("primary"), vec![50, 100]),
(
vec![transient(), transient(), transient()],
Ok("backup"),
vec![50, 100],
),
(vec![rate_limit()], Ok("backup"), vec![]),
(vec![authentication()], Err(authentication()), vec![]),
];
for (faults, outcome, delays) in cases {
let mut primary: Vec<Result<StreamEvent>> = faults.into_iter().map(Err).collect();
primary.push(text("primary"));
let endpoints = HashMap::from([
endpoint("primary", primary),
endpoint("backup", vec![text("backup")]),
]);
let slept = Arc::new(std::sync::Mutex::new(Vec::new()));
let record = Arc::clone(&slept);
let router = Router::new(catalog(), endpoints, 2)
.unwrap()
.with_sleeper(Arc::new(move |ms| record.lock().unwrap().push(ms)));
let served = router
.route("plan", &request(), &CancelHandle::new())
.and_then(drain)
.map(|event| match event {
StreamEvent::TextDelta { delta, .. } => delta,
other => format!("{other:?}"),
});
assert_eq!(served, outcome.map(str::to_string), "delays {delays:?}");
assert_eq!(*slept.lock().unwrap(), delays);
}
}
#[test]
fn invalid_requests_fail_without_failover() {
let endpoints = HashMap::from([
endpoint("primary", ok_stream()),
endpoint("backup", ok_stream()),
]);
let router = Router::new(catalog(), endpoints, 2).unwrap();
let mut invalid = request();
invalid.params.thinking = true;
let err = route_err(&router, "plan", &invalid, &CancelHandle::new());
assert_eq!(err.code(), "PROVIDER_INVALID");
}
#[test]
fn exhausting_the_plan_surfaces_the_last_fault() {
let endpoints = HashMap::from([
endpoint("primary", vec![Err(rate_limit())]),
endpoint("backup", vec![Err(rate_limit())]),
]);
let router = Router::new(catalog(), endpoints, 2).unwrap();
let err = route_err(&router, "plan", &request(), &CancelHandle::new());
assert_eq!(err.code(), "PROVIDER_RATE_LIMIT");
}
}