//! 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>, /// Transient retries per endpoint before failing over. transient_retries: u32, /// Injected delay for retry backoff; tests pass a no-op. sleeper: Arc, } 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, transient_retries: u32, ) -> Result { 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) -> 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 { let plan = self.catalog.resolve_plan(target)?; let mut last_fault: Option = 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 { let mut stream = stream; let mut buffer: Vec = 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> { 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>) -> 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>) -> (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> { vec![Ok(StreamEvent::stop( StopReason::EndTurn, Some(Usage::new()), ))] } fn drain(mut stream: ProviderStream) -> Result { 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 { 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, Result<&str>, Vec); 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> = 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"); } }