Luigit
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");
    }
}