Skip to main content

switchyard_libsy/core/
algorithm.rs

1// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! Routing algorithms request external work through [`Driver`] and receive typed replies.
5//! Hosts execute that work, keeping routing policy independent of transport.
6
7use std::{
8    collections::HashMap, future::Future, panic::AssertUnwindSafe, pin::Pin, sync::Arc,
9    time::Instant,
10};
11
12use async_trait::async_trait;
13use futures::{FutureExt, Stream, StreamExt};
14use parking_lot::Mutex;
15use serde_json::Value;
16use tokio::sync::{mpsc, oneshot};
17use tokio_stream::wrappers::ReceiverStream;
18use tracing::Instrument;
19
20/// The request/response protocol types come from [`switchyard_protocol`].
21/// [`switchyard_protocol::LlmRequest`] is the normalized request;
22/// [`switchyard_protocol::AggLlmResponse`] is the buffered response;
23/// [`switchyard_protocol::LlmResponseChunk`] is normalized streaming content;
24/// [`switchyard_protocol::LlmResponseStreamEvent`] is its host/algorithm envelope; and
25/// [`switchyard_protocol::LlmResponse`] carries either a live
26/// [`switchyard_protocol::LlmResponseStream`] or the terminal aggregate.
27use switchyard_protocol::{
28    Category, DecisionRequest, DecisionResponse, ModelId, Request, Response,
29};
30
31use crate::{DriverError, LibsyError, Result, observability};
32
33/// A boxed, `Send` stream of [`Step`]s — the output of
34/// [`Algorithm::run_stream`]. Boxed so the trait method that produces it keeps
35/// `Arc<dyn Algorithm>` object-safe.
36pub type StepStream = Pin<Box<dyn Stream<Item = Result<Step>> + Send>>;
37
38/// The models one algorithm run may use, grouped by [`Category`]. Within a
39/// category they are ordered best-first.
40///
41/// Delegated sub-agent work gets its own groups, reachable only through
42/// [`Driver::for_subagent`]. Keeping them separate is what stops a sub-agent's
43/// `capable` from resolving to the parent's, and stops the parent falling back
44/// onto a model only its sub-agents were given.
45///
46/// One run's driver clones all read the same value, so it is passed as
47/// `Arc<RuntimeModels>` rather than cloned per driver.
48#[derive(Clone, Debug, Default)]
49pub struct RuntimeModels {
50    /// When using subagents this is the parent agent category.
51    by_category: HashMap<Category, Vec<ModelId>>,
52    subagent: Option<HashMap<Category, Vec<ModelId>>>,
53}
54
55impl RuntimeModels {
56    /// The models available to the algorithm itself.
57    pub fn new(by_category: HashMap<Category, Vec<ModelId>>) -> Self {
58        Self {
59            by_category,
60            subagent: None,
61        }
62    }
63
64    /// Adds the groups used for delegated sub-agent work.
65    pub fn with_subagent(mut self, models: HashMap<Category, Vec<ModelId>>) -> Self {
66        self.subagent = Some(models);
67        self
68    }
69
70    /// The models in `category`, ordered best-first.
71    pub fn models_for(&self, category: &Category) -> &[ModelId] {
72        self.by_category.get(category).map_or(&[], Vec::as_slice)
73    }
74
75    /// The models delegated sub-agent work uses for `category`, ordered best-first.
76    pub fn subagent_models_for(&self, category: &Category) -> &[ModelId] {
77        self.subagent
78            .as_ref()
79            .and_then(|models| models.get(category))
80            .map_or(&[], Vec::as_slice)
81    }
82}
83
84impl From<HashMap<Category, Vec<ModelId>>> for RuntimeModels {
85    fn from(by_category: HashMap<Category, Vec<ModelId>>) -> Self {
86        Self::new(by_category)
87    }
88}
89
90/// Which of a [`RuntimeModels`]' groups a driver reads.
91#[derive(Clone, Copy)]
92enum Scope {
93    Parent,
94    Subagent,
95}
96
97/// An offloaded model call, surfaced inside [`Step::CallModel`].
98///
99/// The host reads the public fields, performs (or delegates) the model call, and fulfills it
100/// with [`respond`](Self::respond) — unblocking the algorithm's [`Driver::call_model`] on the
101/// other side. `switchyard-llm-client`'s `run` is the ready-made consumer that does this for
102/// you.
103///
104/// [`Driver::call_model`] stamps the first candidate model onto the request before publishing
105/// the call. A consumer that falls through to a later candidate must re-stamp it.
106pub struct CallModel {
107    /// The name of the algorithm that produced this call, so a host instrumenting the
108    /// calls it serves can attribute its own spans to the algorithm behind them.
109    pub algorithm: String,
110    /// The request to serve; its `model` is stamped with the first candidate.
111    pub request: Request,
112    /// Candidate models, tried in order until one answers. Never empty.
113    pub models: Vec<ModelId>,
114    /// Return client errors to the algorithm so it can apply its fallback policy.
115    pub recover_errors: bool,
116    /// How to send the response back to the algorithm. `None` once the call is recorded.
117    reply: Option<oneshot::Sender<Result<Response>>>,
118    started: Instant,
119}
120
121impl CallModel {
122    /// Fulfill the promise with the caller's model-call result. Pass `Err(..)` to
123    /// propagate a failed model call back to the algorithm. Consumes the promise: it
124    /// can only be fulfilled once.
125    pub fn respond(mut self, result: Result<Response>) -> Result<()> {
126        self.record(result.is_ok());
127        self.reply
128            .take()
129            .ok_or(DriverError::ResponseDropped)?
130            .send(result)
131            .map_err(|_| DriverError::ResponseDropped.into())
132    }
133
134    /// Record a failed call and return its error to stop [`drive`].
135    /// Leaves the promise unfulfilled so the driver can cancel the algorithm.
136    pub fn fail(mut self, error: LibsyError) -> Result<()> {
137        self.reply = None;
138        self.record(false);
139        Err(error)
140    }
141
142    fn record(&self, is_ok: bool) {
143        observability::record_llm_call(
144            &self.algorithm,
145            self.models
146                .first()
147                .map(ModelId::as_str)
148                .unwrap_or("NoTargets"),
149            self.started.elapsed(),
150            is_ok,
151        );
152    }
153}
154
155impl Drop for CallModel {
156    fn drop(&mut self) {
157        if self.reply.is_some() {
158            self.record(false);
159        }
160    }
161}
162
163/// A host-owned decision call. Dropping it without replying yields [`DriverError::ResponseDropped`].
164pub struct CallDecision {
165    /// Algorithm name for attributing host telemetry.
166    pub algorithm: String,
167    /// Its model is set to [`Self::model`] before dispatch.
168    pub request: DecisionRequest,
169    /// Target ID for client lookup.
170    pub model: ModelId,
171    reply: oneshot::Sender<Result<DecisionResponse>>,
172}
173
174impl CallDecision {
175    /// Return a response or provider error to the algorithm so it can continue or fall back.
176    pub fn respond(self, result: Result<DecisionResponse>) -> Result<()> {
177        self.reply
178            .send(result)
179            .map_err(|_| DriverError::ResponseDropped.into())
180    }
181
182    /// Returning this error from the host handler aborts [`drive`].
183    pub fn fail(self, error: LibsyError) -> Result<()> {
184        Err(error)
185    }
186}
187
188/// The terminal result of routing.
189pub struct RoutingOutcome {
190    /// Models selected by the algorithm, ordered best model first.
191    pub selected_model_ids: Vec<ModelId>,
192    /// The request after all routing-time rewrites, stamped with the selected model.
193    pub request: Request,
194    /// A response produced while routing, or `None` when the client must make the answer call.
195    pub response: Option<Response>,
196    /// Outcome identity and optional algorithm evidence.
197    ///
198    /// Constructors leave this empty; [`Algorithm::run_stream`] fills it before publishing a
199    /// successful outcome.
200    pub metadata: Option<crate::OutcomeMetadata>,
201}
202
203impl RoutingOutcome {
204    /// The model the algorithm recommends, the best model for this request.
205    /// `LibsyError::NoTargets` if the algorithm selected no models, which should be impossible.
206    pub fn selected_model_id(&self) -> Result<&ModelId> {
207        self.selected_model_ids.first().ok_or(LibsyError::NoTargets)
208    }
209
210    /// The decision is that client should send this `request`. The `selected_model_id`
211    /// will be written into it by this function.
212    /// If that fails client should try the `fallback_models` in order.
213    pub fn route_to(
214        selected_model_id: ModelId,
215        fallback_models: Vec<ModelId>,
216        mut request: Request,
217    ) -> Self {
218        request.llm_request.model = Some(selected_model_id.to_string());
219        let mut selected_model_ids = Vec::with_capacity(1 + fallback_models.len());
220        selected_model_ids.push(selected_model_id);
221        selected_model_ids.extend(fallback_models);
222        Self {
223            selected_model_ids,
224            request,
225            response: None,
226            metadata: None,
227        }
228    }
229
230    /// Algorithm generated the response as part of the routing decision. Here it is.
231    /// The `request` will have the `selected_model_id` written into it by this function.
232    pub fn answered(selected_model_id: ModelId, mut request: Request, response: Response) -> Self {
233        request.llm_request.model = Some(selected_model_id.to_string());
234        Self {
235            selected_model_ids: vec![selected_model_id],
236            request,
237            response: Some(response),
238            metadata: None,
239        }
240    }
241}
242
243/// An algorithm's handle for requesting external work and awaiting the host's typed replies.
244#[derive(Clone)]
245pub struct Driver {
246    step_tx: mpsc::Sender<Result<Step>>,
247
248    /// The owning algorithm's telemetry label, stamped onto every call this driver publishes.
249    algorithm: String,
250
251    /// Run-scoped evidence shared by driver clones and attached only to a successful outcome.
252    evidence: Arc<Mutex<Option<Value>>>,
253
254    /// Every group this run may route over, shared by all driver clones.
255    models: Arc<RuntimeModels>,
256
257    /// Which of those groups this driver reads.
258    scope: Scope,
259}
260
261impl Driver {
262    /// Build an empty driver with its step channel ready. Created per call by
263    /// [`run_stream`](Algorithm::run_stream). Also returns the Step receiver.
264    pub(crate) fn new(
265        algorithm: &str,
266        models: Arc<RuntimeModels>,
267    ) -> (Self, mpsc::Receiver<Result<Step>>) {
268        // Capacity one keeps the algorithm paced by the stream consumer. It limits queued steps,
269        // not model calls already pulled from the stream, which can still run at the same time.
270        // A larger buffer would use more memory and let the algorithm run farther ahead with
271        // little benefit because reading a step is cheap compared with serving a model call.
272        let (step_tx, step_rx) = mpsc::channel(1);
273        (
274            Self {
275                step_tx,
276                algorithm: algorithm.to_string(),
277                evidence: Arc::new(Mutex::new(None)),
278                models,
279                scope: Scope::Parent,
280            },
281            step_rx,
282        )
283    }
284
285    /// Replace the current run's evidence when a component makes the final decision.
286    pub(crate) fn set_evidence(&self, evidence: Value) {
287        *self.evidence.lock() = Some(evidence);
288    }
289
290    /// Supply fallback evidence without replacing a decision made earlier in the cascade.
291    pub(crate) fn set_evidence_if_empty(&self, evidence: Value) {
292        let mut current = self.evidence.lock();
293        if current.is_none() {
294            *current = Some(evidence);
295        }
296    }
297
298    /// Publish a model call and await the consumer's response.
299    ///
300    /// Errors if the stream is closed or the call failed.
301    /// The await is wrapped in a `libsy.llm_call` span measuring *fulfillment* as
302    /// the algorithm observes it (host queueing/serving included; a streamed
303    /// response resolves when its stream handle arrives). The host records call metrics
304    /// through [`CallModel::respond`] or [`CallModel::fail`]; outcome and token usage
305    /// are recorded on the span when the promise resolves. The provider call itself is the
306    /// host's, and is instrumented by whoever makes it.
307    pub async fn call_model(&self, request: Request, models: Vec<ModelId>) -> Result<Response> {
308        self.call_model_with_error_recovery(request, models, false)
309            .await
310    }
311
312    /// Allows a routing policy to handle a failed call when recovery is enabled.
313    #[tracing::instrument(
314        target = "libsy",
315        name = "libsy.llm_call",
316        skip_all,
317        fields(
318            algorithm = self.algorithm,
319            selected_model = %models.first().map(ModelId::as_str).unwrap_or("NoTargets"),
320            openinference.span.kind = "CHAIN",
321            outcome = tracing::field::Empty,
322            input_tokens = tracing::field::Empty,
323            output_tokens = tracing::field::Empty,
324            total_tokens = tracing::field::Empty,
325            reasoning_tokens = tracing::field::Empty,
326        )
327    )]
328    pub(crate) async fn call_model_with_error_recovery(
329        &self,
330        mut request: Request,
331        models: Vec<ModelId>,
332        recover_errors: bool,
333    ) -> Result<Response> {
334        let Some(selected_model_id) = models.first() else {
335            return Err(LibsyError::NoTargets);
336        };
337        request.llm_request.model = Some(selected_model_id.to_string());
338        let started = Instant::now();
339        let (reply, response) = oneshot::channel::<Result<Response>>();
340        let call = CallModel {
341            algorithm: self.algorithm.clone(),
342            request,
343            models,
344            recover_errors,
345            reply: Some(reply),
346            started,
347        };
348        let result = async {
349            self.step_tx
350                .send(Ok(Step::CallModel(Box::new(call))))
351                .await
352                .map_err(|_| DriverError::StreamClosed)?;
353            response
354                .await
355                .map_err(|_| LibsyError::from(DriverError::ResponseDropped))?
356        }
357        .await;
358        observability::record_llm_call_span(&result, &tracing::Span::current());
359        result
360    }
361
362    /// Override the request's model with `model` and wait for the host to reply.
363    #[tracing::instrument(
364        target = "libsy",
365        name = "libsy.decision_call",
366        skip_all,
367        fields(algorithm = self.algorithm, selected_model = %model),
368    )]
369    pub async fn call_decision(
370        &self,
371        mut request: DecisionRequest,
372        model: ModelId,
373    ) -> Result<DecisionResponse> {
374        request.model = Some(model.clone());
375        let (reply, response) = oneshot::channel();
376        let call = CallDecision {
377            algorithm: self.algorithm.clone(),
378            request,
379            model,
380            reply,
381        };
382        self.step_tx
383            .send(Ok(Step::CallDecision(Box::new(call))))
384            .await
385            .map_err(|_| DriverError::StreamClosed)?;
386        response
387            .await
388            .map_err(|_| LibsyError::from(DriverError::ResponseDropped))?
389    }
390
391    /// The available models for this category, typically ordered best-first.
392    pub fn models_for(&self, category: &Category) -> &[ModelId] {
393        match self.scope {
394            Scope::Parent => self.models.models_for(category),
395            Scope::Subagent => self.models.subagent_models_for(category),
396        }
397    }
398
399    /// The first available model for `category`.
400    pub fn first_model_for(&self, category: &Category) -> Result<&ModelId> {
401        self.models_for(category)
402            .first()
403            .ok_or_else(|| LibsyError::AlgorithmError {
404                message: format!("no models available for category {}", category.as_str()),
405            })
406    }
407
408    /// A driver scoped to delegated sub-agent work: its categories are the
409    /// sub-agent's own, and the parent's are no longer reachable through it.
410    pub fn for_subagent(&self) -> Result<Self> {
411        if self.models.subagent.is_none() {
412            return Err(LibsyError::AlgorithmError {
413                message: "delegated work has no sub-agent models".to_string(),
414            });
415        }
416        Ok(Self {
417            scope: Scope::Subagent,
418            ..self.clone()
419        })
420    }
421
422    /// Emit the terminal step: [`Step::Done`] on `Ok`, or an `Err` stream
423    /// item on failure. Internal: called once by [`run_stream`](Algorithm::run_stream)
424    /// when the algorithm finishes.
425    pub(crate) async fn finish(&self, result: Result<RoutingOutcome>) -> Result<()> {
426        let result = result.map(|mut outcome| {
427            let metadata = outcome.metadata.get_or_insert_with(|| {
428                crate::OutcomeMetadata::new(self.algorithm.clone(), self.evidence.lock().take())
429            });
430            observability::record_outcome(metadata, &outcome.selected_model_ids);
431            outcome
432        });
433        let selected_model = result
434            .as_ref()
435            .ok()
436            .and_then(|outcome| outcome.selected_model_id().ok().cloned());
437        let step = result.map(|outcome| Step::Done(Box::new(outcome)));
438        self.step_tx
439            .send(step)
440            .await
441            .map_err(|_| DriverError::StreamClosed)?;
442        if let Some(selected_model) = selected_model {
443            observability::record_decision(&self.algorithm, &selected_model);
444        }
445        Ok(())
446    }
447}
448
449/// Work requests and the terminal outcome of [`Algorithm::run_stream`].
450/// Each call carries its own reply channel, so the host can execute calls concurrently.
451pub enum Step {
452    /// The algorithm needs this model call performed. The host serves it and fulfills
453    /// it with [`CallModel::respond`]. Boxed: it is by far the largest variant.
454    CallModel(Box<CallModel>),
455    /// The host fulfills this decision call with [`CallDecision::respond`].
456    CallDecision(Box<CallDecision>),
457    /// The algorithm finished with its routing outcome — the last step of a run.
458    Done(Box<RoutingOutcome>),
459}
460
461/// Work dispatched through one host handler, with a typed reply channel per variant.
462/// [`drive`] handles terminal outcomes separately; the handler only executes calls.
463pub enum Call {
464    /// Fulfilled with an LLM [`Response`].
465    Model(Box<CallModel>),
466    /// Fulfilled with a [`DecisionResponse`].
467    Decision(Box<CallDecision>),
468}
469
470/// Drive [`Algorithm::run_stream`] to its routing outcome, serving calls concurrently.
471///
472/// The shared execution loop lets hosts supply clients or test doubles through `serve`
473/// without reimplementing stream handling and cancellation.
474///
475/// `serve` owns each [`Call`] and performs its I/O. Passing an error to `respond`
476/// lets the algorithm fall back; returning an error from `serve` aborts the run.
477///
478/// Dropping this future cancels the algorithm and pending handler futures.
479pub async fn drive<F, Fut>(
480    algorithm: Arc<dyn Algorithm>,
481    request: Request,
482    models: Arc<RuntimeModels>,
483    serve: F,
484) -> Result<RoutingOutcome>
485where
486    F: Fn(Call) -> Fut,
487    Fut: Future<Output = Result<()>>,
488{
489    let stream = algorithm.run_stream(request, models);
490    tokio::pin!(stream);
491
492    let mut in_flight = futures::stream::FuturesUnordered::new();
493    let mut final_outcome: Option<RoutingOutcome> = None;
494
495    loop {
496        tokio::select! {
497            Some(result) = in_flight.next() => match result {
498                Ok(()) => {}, // Call completed successfully
499                Err(err) => return Err(err), // Call failed, propagate the error
500            },
501            step = stream.next() => {
502                match step {
503                    None => break, // stream has ended, no more steps
504                    Some(item) => match item? {
505                        Step::CallModel(call) => in_flight.push(serve(Call::Model(call))),
506                        Step::CallDecision(call) => {
507                            in_flight.push(serve(Call::Decision(call)));
508                        }
509                        Step::Done(outcome) => {
510                            final_outcome = Some(*outcome);
511                            break;
512                        }
513                    }
514                }
515            },
516        }
517    }
518    final_outcome.ok_or(LibsyError::MissingFinalResponse)
519}
520
521/// Recover the message from an algorithm's panic.
522fn panic_message(payload: &(dyn std::any::Any + Send)) -> String {
523    payload
524        .downcast_ref::<&'static str>()
525        .map(|message| (*message).to_string())
526        .or_else(|| payload.downcast_ref::<String>().cloned())
527        .unwrap_or_else(|| "unknown panic payload".to_string())
528}
529
530/// Abort guard
531struct AbortOnDrop(tokio::task::AbortHandle);
532
533impl Drop for AbortOnDrop {
534    fn drop(&mut self) {
535        self.0.abort();
536    }
537}
538
539/// Errors unless `targets` contains `name`.
540///
541/// Config target names must be resolved before an algorithm is built. This list contains
542/// model IDs, not target names.
543pub(crate) fn ensure_model_is_target(targets: &[ModelId], name: &ModelId) -> Result<()> {
544    targets
545        .iter()
546        .any(|target| target == name)
547        .then_some(())
548        .ok_or_else(|| LibsyError::TargetNotFound {
549            target: name.clone(),
550        })
551}
552
553/// Key for routing affinity: a root request by its session, a child request by its session
554/// and agent.
555#[derive(Clone, Hash, PartialEq, Eq)]
556pub(crate) enum RoutingIdentity {
557    /// Root request, keyed by session ID.
558    Session(String),
559    /// Child request, keyed by session and agent IDs.
560    Subagent { session: String, agent: String },
561}
562
563impl RoutingIdentity {
564    /// Builds a root or child identity from non-empty request metadata.
565    ///
566    /// A child request missing either ID returns `None`, so it keeps no routing history
567    /// rather than sharing the parent's.
568    pub(crate) fn from_request(request: &Request) -> Option<Self> {
569        let metadata = request.metadata.as_ref()?;
570        let session = metadata.session_id.as_deref().filter(|id| !id.is_empty())?;
571        if metadata.is_subagent {
572            let agent = metadata.agent_id.as_deref().filter(|id| !id.is_empty())?;
573            Some(Self::Subagent {
574                session: session.to_string(),
575                agent: agent.to_string(),
576            })
577        } else {
578            Some(Self::Session(session.to_string()))
579        }
580    }
581}
582
583/// Routing policy independent of how external calls are executed.
584/// Implement [`route`](Self::route), using [`Driver`] to request work from the host.
585/// Hosts can use [`drive`] for the shared execution loop or consume
586/// [`run_stream`](Self::run_stream) for custom execution.
587///
588/// Methods take `self: Arc<Self>`: one algorithm (`Arc<dyn Algorithm>`) is shared across
589/// requests and run concurrently, so it owns its thread-safety and any shared state.
590///
591/// # Concurrency
592///
593/// A host may run the same algorithm concurrently for many requests. Implementations
594/// must synchronize their own mutable shared state. Each call to [`run_stream`](Self::run_stream)
595/// creates an independent [`Driver`], so model-call promises and emitted [`Step`]s cannot
596/// cross between runs.
597///
598/// # Observability
599///
600/// [`run_stream`](Self::run_stream) creates a `libsy.run` span, and each offloaded model
601/// call creates a nested `libsy.llm_call` span. Successful outcomes record their
602/// [`OutcomeMetadata::outcome_id`](crate::OutcomeMetadata::outcome_id) on `libsy.run`,
603/// alongside `selected_model_ids` (an ordered OpenTelemetry string array).
604/// `algorithm` and `switchyard.algorithm` retain the run's [`Algorithm::name`].
605/// Optional `evidence.source`, `evidence.verdict`, `evidence.trigger`, and
606/// `evidence.reason_code` are strings; `evidence.score`, `evidence.confidence`, and
607/// `evidence.threshold` are numbers. Unknown evidence fields are not exported.
608/// These fields are span attributes, never metric labels.
609///
610/// The run/call observability helpers retain `outcome` status and operational metrics,
611/// but omit error details and arbitrary request extra metadata. Algorithms and hosts
612/// may emit their own logs. Errors still reach the caller unchanged.
613/// The host controls the tracing subscriber and global OpenTelemetry
614/// meter provider; libsy installs no exporter and performs no telemetry network I/O.
615#[async_trait]
616pub trait Algorithm: Send + Sync + 'static {
617    /// Stable, low-cardinality name identifying this algorithm — the
618    /// `algorithm` attribute on every span, metric, and log line the crate
619    /// emits for its runs.
620    fn name(&self) -> &str;
621
622    /// Select a route, requesting external work through [`Driver`] as needed.
623    /// [`run_stream`](Self::run_stream) runs this method and exposes its work to the host.
624    async fn route(self: Arc<Self>, driver: Driver, request: Request) -> Result<RoutingOutcome>;
625
626    /// Expose work requests and the final outcome so the host can control execution.
627    ///
628    /// Each call waits for its host reply. A run ends with one terminal item — [`Step::Done`] on
629    /// success, an `Err` item on failure, including when the algorithm panics. Dropping
630    /// the stream aborts the spawned algorithm task.
631    ///
632    /// Every invocation owns a separate [`Driver`].
633    fn run_stream(self: Arc<Self>, request: Request, models: Arc<RuntimeModels>) -> StepStream {
634        let (driver, step_rx) = Driver::new(self.name(), models);
635        let span = observability::run_span(self.name(), &request);
636        let handle = tokio::spawn(
637            async move {
638                let algorithm = self.name().to_string();
639                // Catch a panicking algorithm so the run still publishes a terminal step.
640                let route = AssertUnwindSafe(self.route(driver.clone(), request)).catch_unwind();
641                let result = observability::observe_run(&algorithm, async move {
642                    route.await.unwrap_or_else(|payload| {
643                        Err(LibsyError::AlgorithmError {
644                            message: format!(
645                                "algorithm task panicked: {}",
646                                panic_message(payload.as_ref())
647                            ),
648                        })
649                    })
650                })
651                .await;
652
653                let _ = driver.finish(result).await;
654            }
655            .instrument(span),
656        );
657        // Dropping the stream aborts the algorithm task when its consumer goes away.
658        let abort_guard = AbortOnDrop(handle.abort_handle());
659        Box::pin(ReceiverStream::new(step_rx).map(move |step| {
660            // link abort guard to stream
661            let _keep_alive = &abort_guard;
662            step
663        }))
664    }
665}
666
667#[cfg(test)]
668mod tests {
669    use std::collections::HashMap;
670
671    use super::*;
672    use crate::core::testing::{Serve, ServeResult, echo, reply, serve_decision, test_drive};
673    use futures::StreamExt;
674    use switchyard_protocol::{
675        LlmResponse, LlmResponseChunk, completion_text, text_request, text_response,
676    };
677
678    #[derive(Debug, thiserror::Error)]
679    #[error("{0}")]
680    struct TestError(&'static str);
681
682    fn test_error(message: &'static str) -> LibsyError {
683        LibsyError::external("test", TestError(message))
684    }
685
686    /// Trivial algo used only to exercise the orchestrator: calls the first target
687    /// and returns its response as the routing outcome.
688    struct TestAlgo {
689        target_set: Vec<ModelId>,
690    }
691
692    #[async_trait]
693    impl Algorithm for TestAlgo {
694        fn name(&self) -> &str {
695            "test"
696        }
697
698        async fn route(
699            self: Arc<Self>,
700            driver: Driver,
701            request: Request,
702        ) -> Result<RoutingOutcome> {
703            let target = self
704                .target_set
705                .first()
706                .ok_or(LibsyError::NoTargets)?
707                .clone();
708            let response = driver
709                .call_model(request.clone(), vec![target.clone()])
710                .await?;
711            driver.set_evidence(serde_json::json!({"source": "test"}));
712            driver.set_evidence_if_empty(serde_json::json!({"source": "ignored"}));
713            Ok(RoutingOutcome::answered(target, request, response))
714        }
715    }
716
717    /// Build a shared `TestAlgo` over the given target set.
718    fn orch(target_set: Vec<ModelId>) -> Arc<dyn Algorithm> {
719        Arc::new(TestAlgo { target_set })
720    }
721
722    fn request() -> Request {
723        Request {
724            llm_request: text_request(Some("auto".to_string()), "hi".to_string()),
725            raw_request: None,
726            metadata: None,
727        }
728    }
729
730    #[test]
731    fn routing_outcome_constructors_stamp_selection_and_preserve_payloads() {
732        let outcome = RoutingOutcome::route_to(
733            "selected".into(),
734            target_set(&["fallback-one", "fallback-two"]),
735            request(),
736        );
737
738        assert_eq!(
739            outcome.selected_model_ids,
740            target_set(&["selected", "fallback-one", "fallback-two"])
741        );
742        assert_eq!(outcome.request.model_id().as_deref(), Some("selected"));
743        assert!(outcome.response.is_none());
744        assert!(outcome.metadata.is_none());
745
746        let outcome = RoutingOutcome::route_to("only".into(), Vec::new(), request());
747        assert_eq!(outcome.selected_model_ids, target_set(&["only"]));
748
749        let outcome = RoutingOutcome::answered(
750            "answered".into(),
751            request(),
752            Response {
753                llm_response: LlmResponse::Agg(text_response(None, "existing")),
754                metadata: None,
755                upstream_headers: http::HeaderMap::new(),
756            },
757        );
758
759        assert_eq!(outcome.selected_model_ids, target_set(&["answered"]));
760        assert_eq!(outcome.request.model_id().as_deref(), Some("answered"));
761        assert_eq!(
762            outcome
763                .response
764                .as_ref()
765                .and_then(|response| response.llm_response.as_agg())
766                .map(completion_text),
767            Some("existing".to_string())
768        );
769    }
770
771    fn target_set(names: &[&str]) -> Vec<ModelId> {
772        names.iter().map(|name| ModelId::from(*name)).collect()
773    }
774
775    #[tokio::test]
776    async fn typed_driver_preserves_call_and_stream_boundaries() -> Result<()> {
777        tokio::time::timeout(std::time::Duration::from_secs(1), async {
778            // Distinct oneshots keep reverse-order replies paired with their producers, and a
779            // retained call remains pending until the host responds.
780            let (driver, mut step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
781            let first_driver = driver.clone();
782            let mut first = tokio::spawn(async move {
783                first_driver
784                    .call_model(request(), vec![ModelId::from("first")])
785                    .await
786            });
787            let second = tokio::spawn(async move {
788                driver
789                    .call_model(request(), vec![ModelId::from("second")])
790                    .await
791            });
792
793            let mut calls = HashMap::new();
794            for _ in 0..2 {
795                let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
796                let Step::CallModel(call) = step else {
797                    return Err(test_error("expected a CallModel step"));
798                };
799                let selected_model = call
800                    .models
801                    .first()
802                    .ok_or_else(|| test_error("model call has no candidates"))?
803                    .to_string();
804                calls.insert(selected_model, call);
805            }
806            assert!(
807                tokio::time::timeout(std::time::Duration::from_millis(20), &mut first)
808                    .await
809                    .is_err(),
810                "call completed before the host responded"
811            );
812            calls
813                .remove("second")
814                .ok_or_else(|| test_error("missing second call"))?
815                .respond(Ok(reply("second response")))?;
816            calls
817                .remove("first")
818                .ok_or_else(|| test_error("missing first call"))?
819                .respond(Ok(reply("first response")))?;
820
821            let first_response = first
822                .await
823                .map_err(|source| LibsyError::external("joining a test task", source))??;
824            let second_response = second
825                .await
826                .map_err(|source| LibsyError::external("joining a test task", source))??;
827            assert_eq!(
828                first_response.llm_response.as_agg().map(completion_text),
829                Some("first response".to_string())
830            );
831            assert_eq!(
832                second_response.llm_response.as_agg().map(completion_text),
833                Some("second response".to_string())
834            );
835
836            // Dropping the host-facing promise closes only that call's reply channel.
837            let (driver, mut step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
838            let producer = tokio::spawn(async move {
839                driver
840                    .call_model(request(), vec![ModelId::from("dropped")])
841                    .await
842            });
843            let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
844            let Step::CallModel(call) = step else {
845                return Err(test_error("expected a CallModel step"));
846            };
847            drop(call);
848            let result = producer
849                .await
850                .map_err(|source| LibsyError::external("joining a test task", source))?;
851            assert!(matches!(
852                result,
853                Err(LibsyError::Driver(DriverError::ResponseDropped))
854            ));
855
856            // A standalone driver reports the typed step receiver disappearing at its next call.
857            let (driver, step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
858            drop(step_rx);
859            let result = driver
860                .call_model(request(), vec![ModelId::from("closed")])
861                .await;
862            assert!(matches!(
863                result,
864                Err(LibsyError::Driver(DriverError::StreamClosed))
865            ));
866
867            fn decision_request() -> DecisionRequest {
868                DecisionRequest {
869                    model: Some("overwritten".into()),
870                    context: serde_json::json!({"task": "choose a route"}),
871                    questions: Default::default(),
872                }
873            }
874
875            fn decision_response() -> DecisionResponse {
876                DecisionResponse {
877                    id: Some("decision-1".to_string()),
878                    model: Some("provider-model".into()),
879                    answers: [(
880                        "p_solve".to_string(),
881                        switchyard_protocol::DecisionAnswer {
882                            value: switchyard_protocol::DecisionValue::Boolean(
883                                switchyard_protocol::BooleanEstimate::ProbabilityTrue(
884                                    switchyard_protocol::Probability(0.8),
885                                ),
886                            ),
887                            provider_confidence: None,
888                        },
889                    )]
890                    .into(),
891                    usage: Default::default(),
892                }
893            }
894
895            struct MixedCalls(&'static str);
896
897            #[async_trait]
898            impl Algorithm for MixedCalls {
899                fn name(&self) -> &str {
900                    "mixed"
901                }
902
903                async fn route(
904                    self: Arc<Self>,
905                    driver: Driver,
906                    request: Request,
907                ) -> Result<RoutingOutcome> {
908                    let (llm, decision) = tokio::join!(
909                        driver.call_model(request.clone(), vec!["llm".into()]),
910                        driver.call_decision(decision_request(), "decision".into()),
911                    );
912                    assert_eq!(
913                        llm?.llm_response.as_agg().map(completion_text),
914                        Some("llm reply".into())
915                    );
916                    match self.0 {
917                        "mock" => assert_eq!(decision?, DecisionResponse {
918                            id: None,
919                            model: Some("decision".into()),
920                            answers: Default::default(),
921                            usage: Default::default(),
922                        }),
923                        "reply" => assert_eq!(decision?, decision_response()),
924                        "error" => assert!(matches!(
925                            decision,
926                            Err(LibsyError::AlgorithmError { message }) if message == "provider failed"
927                        )),
928                        "drop" => assert!(matches!(
929                            decision,
930                            Err(LibsyError::Driver(DriverError::ResponseDropped))
931                        )),
932                        "abort" => return std::future::pending().await,
933                        _ => unreachable!(),
934                    }
935                    Ok(RoutingOutcome::route_to("answer".into(), vec![], request))
936                }
937            }
938
939            for mode in ["mock", "reply", "error", "drop", "abort"] {
940                // Serial dispatch would deadlock here and fail the enclosing timeout.
941                let barrier = Arc::new(tokio::sync::Barrier::new(2));
942                let outcome = drive(
943                    Arc::new(MixedCalls(mode)),
944                    request(),
945                    Arc::new(RuntimeModels::default()),
946                    move |call| {
947                        let barrier = barrier.clone();
948                        async move {
949                            barrier.wait().await;
950                            let call = match call {
951                                Call::Model(call) => return call.respond(Ok(reply("llm reply"))),
952                                Call::Decision(call) => *call,
953                            };
954                            assert_eq!(call.algorithm, "mixed");
955                            assert_eq!(call.model, "decision");
956                            assert_eq!(call.request.model, Some("decision".into()));
957                            assert_eq!(call.request.context, decision_request().context);
958                            match mode {
959                                "mock" => serve_decision(call).await,
960                                "reply" => call.respond(Ok(decision_response())),
961                                "error" => call.respond(Err(LibsyError::AlgorithmError {
962                                    message: "provider failed".into(),
963                                })),
964                                "drop" => {
965                                    drop(call);
966                                    Ok(())
967                                }
968                                "abort" => call.fail(test_error("host aborted")),
969                                _ => unreachable!(),
970                            }
971                        }
972                    },
973                )
974                .await;
975                if mode == "abort" {
976                    assert!(matches!(outcome, Err(LibsyError::External { .. })));
977                } else {
978                    assert_eq!(outcome?.selected_model_id()?, "answer");
979                }
980            }
981
982            let (driver, mut steps) = Driver::new("test", Arc::new(RuntimeModels::default()));
983            let mut pending = Box::pin(driver.call_decision(decision_request(), "decision".into()));
984            assert!(futures::poll!(&mut pending).is_pending());
985            let Some(Ok(Step::CallDecision(call))) = steps.recv().await else {
986                return Err(test_error("expected a decision call"));
987            };
988            drop(pending);
989            assert!(matches!(
990                call.respond(Ok(decision_response())),
991                Err(LibsyError::Driver(DriverError::ResponseDropped))
992            ));
993            drop(steps);
994            assert!(matches!(
995                driver.call_decision(decision_request(), "decision".into()).await,
996                Err(LibsyError::Driver(DriverError::StreamClosed))
997            ));
998            Ok(())
999        })
1000        .await
1001        .map_err(|error| LibsyError::external("waiting for typed driver boundaries", error))?
1002    }
1003
1004    #[test]
1005    fn target_lookup_returns_the_missing_target() {
1006        let error = ensure_model_is_target(&target_set(&[]), &ModelId::from("missing")).err();
1007        assert!(matches!(
1008            error,
1009            Some(LibsyError::TargetNotFound { target }) if target == "missing"
1010        ));
1011    }
1012
1013    /// Build a single-target algo, plus a `serve` that answers it as a token stream
1014    /// replaying `chunks` in order (as `Ok` items).
1015    fn streaming_orch(chunks: Vec<LlmResponseChunk>) -> (Arc<dyn Algorithm>, impl Serve) {
1016        let algo = orch(target_set(&["stream/model"]));
1017        let serve = move |_target: ModelId, _request: Request| {
1018            let chunks = chunks.clone();
1019            async move {
1020                let stream =
1021                    futures::stream::iter(chunks.into_iter().map(|chunk| Ok(chunk.into()))).boxed();
1022                Ok(Response {
1023                    llm_response: LlmResponse::Stream(stream),
1024                    metadata: None,
1025                    upstream_headers: http::HeaderMap::new(),
1026                })
1027            }
1028        };
1029        (algo, serve)
1030    }
1031
1032    #[tokio::test]
1033    async fn run_returns_a_streamed_response_the_caller_aggregates() -> Result<()> {
1034        // A streaming client -> its chunks flow through the promise and `Done`,
1035        // and `run` returns the live stream untouched for the caller to fold.
1036        let (orch, serve) = streaming_orch(vec![
1037            LlmResponseChunk::MessageStart {
1038                id: Some("m1".to_string()),
1039                model: Some("stream/model".to_string()),
1040            },
1041            LlmResponseChunk::TextDelta {
1042                index: 0,
1043                text: "hel".to_string(),
1044            },
1045            LlmResponseChunk::TextDelta {
1046                index: 0,
1047                text: "lo".to_string(),
1048            },
1049            LlmResponseChunk::MessageStop {
1050                reason: Some("stop".to_string()),
1051            },
1052        ]);
1053        let (selected_model, response) = test_drive(orch, request(), serve).await?;
1054        // The run handed back the live stream; the caller folds it to a buffered aggregate.
1055        let agg = response
1056            .llm_response
1057            .into_agg()
1058            .await
1059            .map_err(|error| LibsyError::external("aggregating response stream", error))?;
1060        assert_eq!(completion_text(&agg), "hello");
1061        assert_eq!(agg.model.as_deref(), Some("stream/model"));
1062        assert_eq!(selected_model, "stream/model");
1063        Ok(())
1064    }
1065
1066    #[tokio::test]
1067    async fn aggregating_a_streamed_response_propagates_a_mid_stream_error() -> Result<()> {
1068        // The run succeeds and returns the stream; the in-band `Error` chunk surfaces only
1069        // when the caller aggregates it.
1070        let (orch, serve) = streaming_orch(vec![
1071            LlmResponseChunk::TextDelta {
1072                index: 0,
1073                text: "partial".to_string(),
1074            },
1075            LlmResponseChunk::StreamError {
1076                message: "upstream exploded".to_string(),
1077            },
1078        ]);
1079        let (_, response) = test_drive(orch, request(), serve).await?;
1080        match response.llm_response.into_agg().await {
1081            Ok(_) => panic!("expected a mid-stream error, got an aggregate"),
1082            Err(err) => {
1083                assert!(err.to_string().contains("upstream exploded"));
1084                Ok(())
1085            }
1086        }
1087    }
1088
1089    #[tokio::test]
1090    async fn run_offloads_via_promise_then_finishes() -> Result<()> {
1091        // Every call is offloaded via a promise the orchestrator surfaces as a
1092        // `CallModel` step for us to fulfill.
1093        let stream = orch(target_set(&["offload/model"]))
1094            .run_stream(request(), Arc::new(RuntimeModels::default()));
1095        tokio::pin!(stream);
1096
1097        let mut saw_call = false;
1098        let mut final_completion = None;
1099        while let Some(step) = stream.next().await {
1100            match step? {
1101                Step::CallDecision(_) => return Err(test_error("unexpected decision call")),
1102                Step::CallModel(call) => {
1103                    saw_call = true;
1104                    assert_eq!(call.models, vec![ModelId::from("offload/model")]);
1105                    // Fulfilling the promise is the "real" model call the caller makes.
1106                    call.respond(Ok(Response {
1107                        llm_response: LlmResponse::Agg(text_response(
1108                            None,
1109                            "fulfilled".to_string(),
1110                        )),
1111                        metadata: None,
1112                        upstream_headers: http::HeaderMap::new(),
1113                    }))?;
1114                }
1115                Step::Done(outcome) => {
1116                    let metadata = outcome
1117                        .metadata
1118                        .as_ref()
1119                        .expect("run_stream should attach outcome metadata");
1120                    assert_eq!(metadata.algorithm, "test");
1121                    assert_eq!(
1122                        uuid::Uuid::parse_str(metadata.outcome_id())
1123                            .expect("outcome id should be a UUID")
1124                            .get_version_num(),
1125                        7
1126                    );
1127                    assert_eq!(
1128                        metadata.evidence,
1129                        Some(serde_json::json!({"source": "test"}))
1130                    );
1131                    let response = outcome
1132                        .response
1133                        .ok_or_else(|| test_error("expected an answered outcome"))?;
1134                    final_completion = Some(
1135                        response
1136                            .llm_response
1137                            .as_agg()
1138                            .map(completion_text)
1139                            .unwrap_or_default(),
1140                    );
1141                }
1142            }
1143        }
1144
1145        assert!(saw_call, "expected a CallModel step before Done");
1146        assert_eq!(
1147            final_completion.ok_or_else(|| test_error("no Done step"))?,
1148            "fulfilled"
1149        );
1150        Ok(())
1151    }
1152
1153    #[tokio::test(flavor = "multi_thread", worker_threads = 12)]
1154    async fn requests_are_processed_in_parallel() -> Result<()> {
1155        use std::time::Duration;
1156        use tokio::sync::Barrier;
1157
1158        const N: usize = 12;
1159
1160        // Serving blocks until all N concurrent calls have arrived. If requests were
1161        // serialized (one algorithm behind a `Mutex`), only one call could be in flight,
1162        // the barrier would never reach N, and the test would time out. It passes only
1163        // because the shared algorithm is driven concurrently across requests.
1164        let barrier = Arc::new(Barrier::new(N));
1165        // One shared algorithm driven by many concurrent requests.
1166        let algo = orch(target_set(&["m"]));
1167
1168        let mut handles = Vec::new();
1169        for _ in 0..N {
1170            let algo = algo.clone();
1171            let barrier = barrier.clone();
1172            let serve = move |target: ModelId, _request: Request| {
1173                let barrier = barrier.clone();
1174                async move {
1175                    barrier.wait().await;
1176                    Ok(reply(target))
1177                }
1178            };
1179            handles.push(tokio::spawn(async move {
1180                test_drive(algo, request(), serve)
1181                    .await
1182                    .map(|(_, response)| {
1183                        response
1184                            .llm_response
1185                            .as_agg()
1186                            .map(completion_text)
1187                            .unwrap_or_default()
1188                    })
1189            }));
1190        }
1191
1192        for handle in handles {
1193            // The timeout turns a serialization deadlock into a failure, not a hang.
1194            let completion = tokio::time::timeout(Duration::from_secs(5), handle)
1195                .await
1196                .map_err(|error| LibsyError::external("waiting for test task", error))?
1197                .map_err(|source| LibsyError::external("joining a test task", source))??;
1198            assert_eq!(completion, "m");
1199        }
1200        Ok(())
1201    }
1202
1203    #[tokio::test]
1204    async fn offload_error_propagates_back_to_the_algorithm() -> Result<()> {
1205        // A client-less target offloads its call; we fulfill the promise with an
1206        // Err, which must flow back through `call_model_target` into the algorithm and
1207        // out as an error step — not a response.
1208        let stream = orch(target_set(&["offload/model"]))
1209            .run_stream(request(), Arc::new(RuntimeModels::default()));
1210        tokio::pin!(stream);
1211
1212        let mut saw_error = false;
1213        while let Some(step) = stream.next().await {
1214            match step {
1215                Ok(Step::CallDecision(_)) => return Err(test_error("unexpected decision call")),
1216                Ok(Step::CallModel(call)) => {
1217                    call.respond(Err(test_error("upstream model call failed")))?;
1218                }
1219                Ok(Step::Done(..)) => {
1220                    return Err(test_error(
1221                        "expected the offload error to propagate, got a response",
1222                    ));
1223                }
1224                Err(err) => {
1225                    // The algorithm's `call_model_target` saw the error via the promise.
1226                    assert!(err.to_string().contains("upstream model call failed"));
1227                    saw_error = true;
1228                }
1229            }
1230        }
1231
1232        assert!(saw_error, "expected an error step");
1233        Ok(())
1234    }
1235
1236    #[tokio::test]
1237    async fn dropping_the_stream_cancels_the_algorithm_task() -> Result<()> {
1238        use std::sync::atomic::{AtomicBool, Ordering};
1239        use std::time::Duration;
1240        use tokio::sync::mpsc;
1241
1242        // Sets a flag when dropped, so we can observe whether the algorithm task was
1243        // cancelled/dropped.
1244        struct DropGuard(Arc<AtomicBool>);
1245        impl Drop for DropGuard {
1246            fn drop(&mut self) {
1247                self.0.store(true, Ordering::SeqCst);
1248            }
1249        }
1250
1251        struct StuckAlgo {
1252            started: mpsc::UnboundedSender<()>,
1253            dropped: Arc<AtomicBool>,
1254        }
1255
1256        #[async_trait]
1257        impl Algorithm for StuckAlgo {
1258            fn name(&self) -> &str {
1259                "stuck"
1260            }
1261
1262            async fn route(
1263                self: Arc<Self>,
1264                _driver: Driver,
1265                _request: Request,
1266            ) -> Result<RoutingOutcome> {
1267                let _guard = DropGuard(self.dropped.clone());
1268                let _ = self.started.send(());
1269                // Await forever without ever touching the driver.
1270                std::future::pending::<()>().await;
1271                unreachable!()
1272            }
1273        }
1274
1275        let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1276        let dropped = Arc::new(AtomicBool::new(false));
1277        let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1278            started: started_tx,
1279            dropped: dropped.clone(),
1280        });
1281
1282        let stream = algo.run_stream(request(), Arc::new(RuntimeModels::default()));
1283        started_rx
1284            .recv()
1285            .await
1286            .ok_or_else(|| test_error("task never started"))?;
1287        drop(stream);
1288        tokio::time::sleep(Duration::from_millis(100)).await;
1289
1290        assert!(
1291            dropped.load(Ordering::SeqCst),
1292            "algorithm task was NOT cancelled after dropping the stream"
1293        );
1294        Ok(())
1295    }
1296
1297    #[tokio::test]
1298    async fn route_panic_surfaces_as_a_stream_error() -> Result<()> {
1299        // An algorithm whose task panics must surface an `Err` step carrying the panic
1300        // message, not abort the process from an unobserved detached task.
1301        struct Panicky;
1302
1303        #[async_trait]
1304        impl Algorithm for Panicky {
1305            fn name(&self) -> &str {
1306                "panicky"
1307            }
1308
1309            async fn route(
1310                self: Arc<Self>,
1311                _driver: Driver,
1312                _request: Request,
1313            ) -> Result<RoutingOutcome> {
1314                panic!("boom");
1315            }
1316        }
1317
1318        let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
1319        let stream = algo.run_stream(request(), Arc::new(RuntimeModels::default()));
1320        tokio::pin!(stream);
1321
1322        let mut saw_error = false;
1323        while let Some(step) = stream.next().await {
1324            match step {
1325                Err(err) => {
1326                    // The panic message is preserved, not flattened into an opaque failure.
1327                    assert!(err.to_string().contains("algorithm task panicked: boom"));
1328                    saw_error = true;
1329                }
1330                Ok(_) => return Err(test_error("expected the panic to surface as an error step")),
1331            }
1332        }
1333
1334        assert!(saw_error, "expected an error step from the panicked task");
1335        Ok(())
1336    }
1337
1338    /// A panicking algorithm must publish its terminal step even when it left a `Driver`
1339    /// clone alive in another task. That clone holds the step channel open, so a run that
1340    /// merely unwound would never terminate and the consumer would wait forever.
1341    #[tokio::test]
1342    async fn a_panic_with_a_leaked_driver_clone_still_terminates_the_run() -> Result<()> {
1343        struct LeakyPanic;
1344
1345        #[async_trait]
1346        impl Algorithm for LeakyPanic {
1347            fn name(&self) -> &str {
1348                "leaky_panic"
1349            }
1350
1351            async fn route(
1352                self: Arc<Self>,
1353                driver: Driver,
1354                _request: Request,
1355            ) -> Result<RoutingOutcome> {
1356                tokio::spawn(async move {
1357                    // Outlives the panic below, keeping a sender clone alive.
1358                    let _keep_alive = driver;
1359                    std::future::pending::<()>().await;
1360                });
1361                tokio::task::yield_now().await;
1362                panic!("boom");
1363            }
1364        }
1365
1366        let algo: Arc<dyn Algorithm> = Arc::new(LeakyPanic);
1367        // The timeout turns the hang this guards against into a failure rather than a hang.
1368        let result = tokio::time::timeout(
1369            std::time::Duration::from_secs(1),
1370            test_drive(algo, request(), echo()),
1371        )
1372        .await
1373        .map_err(|error| LibsyError::external("waiting for the panicked run to end", error))?;
1374
1375        match result {
1376            Ok(_) => Err(test_error(
1377                "expected the panic to end the run with an error",
1378            )),
1379            Err(err) => {
1380                assert!(err.to_string().contains("algorithm task panicked: boom"));
1381                Ok(())
1382            }
1383        }
1384    }
1385
1386    #[tokio::test]
1387    async fn cancelling_run_cancels_the_algorithm_task() -> Result<()> {
1388        use std::sync::atomic::{AtomicBool, Ordering};
1389        use std::time::Duration;
1390        use tokio::sync::mpsc;
1391
1392        // Sets a flag when dropped, so we can observe whether the algorithm task was
1393        // cancelled once the `run` future driving it is dropped.
1394        struct DropGuard(Arc<AtomicBool>);
1395        impl Drop for DropGuard {
1396            fn drop(&mut self) {
1397                self.0.store(true, Ordering::SeqCst);
1398            }
1399        }
1400
1401        struct StuckAlgo {
1402            started: mpsc::UnboundedSender<()>,
1403            dropped: Arc<AtomicBool>,
1404        }
1405
1406        #[async_trait]
1407        impl Algorithm for StuckAlgo {
1408            fn name(&self) -> &str {
1409                "stuck"
1410            }
1411
1412            async fn route(
1413                self: Arc<Self>,
1414                _driver: Driver,
1415                _request: Request,
1416            ) -> Result<RoutingOutcome> {
1417                let _guard = DropGuard(self.dropped.clone());
1418                let _ = self.started.send(());
1419                // Hang forever without ever touching the driver, so only cancellation
1420                // (not a dropped step channel) can stop this task.
1421                std::future::pending::<()>().await;
1422                unreachable!()
1423            }
1424        }
1425
1426        let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1427        let dropped = Arc::new(AtomicBool::new(false));
1428        let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1429            started: started_tx,
1430            dropped: dropped.clone(),
1431        });
1432
1433        // Drive the run on its own task, wait until the algorithm task is up, then cancel
1434        // it — dropping its future (and the `run_stream` stream it holds).
1435        let run_task = tokio::spawn(async move { test_drive(algo, request(), echo()).await });
1436        started_rx
1437            .recv()
1438            .await
1439            .ok_or_else(|| test_error("task never started"))?;
1440        run_task.abort();
1441        tokio::time::sleep(Duration::from_millis(100)).await;
1442
1443        assert!(
1444            dropped.load(Ordering::SeqCst),
1445            "algorithm task was NOT cancelled after cancelling run"
1446        );
1447        Ok(())
1448    }
1449
1450    // --- first-wins hedging: `run` must not wait on losing speculative calls -------------
1451
1452    /// Offloads two targets concurrently and returns the first to resolve, dropping the
1453    /// loser's call (first-wins hedging).
1454    struct Hedge {
1455        winner: String,
1456        loser: String,
1457    }
1458
1459    #[async_trait]
1460    impl Algorithm for Hedge {
1461        fn name(&self) -> &str {
1462            "hedge"
1463        }
1464
1465        async fn route(
1466            self: Arc<Self>,
1467            driver: Driver,
1468            request: Request,
1469        ) -> Result<RoutingOutcome> {
1470            let outcome_request = request.clone();
1471            let win = driver.call_model(request.clone(), vec![self.winner.clone().into()]);
1472            let lose = driver.call_model(request, vec![self.loser.clone().into()]);
1473            // First to resolve wins; `select!` drops the losing future (and its promise).
1474            tokio::select! {
1475                res = win => Ok(RoutingOutcome::answered(
1476                    self.winner.clone().into(),
1477                    outcome_request,
1478                    res?,
1479                )),
1480                res = lose => Ok(RoutingOutcome::answered(
1481                    self.loser.clone().into(),
1482                    outcome_request,
1483                    res?,
1484                )),
1485            }
1486        }
1487    }
1488
1489    /// Builds a hedging algo and the `serve` that drives it: the winner is gated behind the
1490    /// loser starting (so the loser's serve is guaranteed in flight when the winner wins),
1491    /// and the loser finishes after `loser_delay` — or never, when `None`.
1492    fn hedge(loser_delay: Option<std::time::Duration>) -> (Arc<dyn Algorithm>, impl Serve) {
1493        let started = Arc::new(tokio::sync::Notify::new());
1494        let algo = Arc::new(Hedge {
1495            winner: "winner".to_string(),
1496            loser: "loser".to_string(),
1497        });
1498        let serve = move |target: ModelId, _request: Request| {
1499            let started = started.clone();
1500            async move {
1501                if target == "loser" {
1502                    started.notify_one();
1503                    match loser_delay {
1504                        Some(delay) => tokio::time::sleep(delay).await,
1505                        None => std::future::pending::<()>().await,
1506                    }
1507                } else {
1508                    started.notified().await;
1509                }
1510                Ok(reply(target))
1511            }
1512        };
1513        (algo, serve)
1514    }
1515
1516    #[tokio::test]
1517    async fn run_returns_the_winner_without_a_late_loser_overwriting_it() -> Result<()> {
1518        // The loser responds 50ms after the winner has already won. `run` must return the
1519        // winner, not the loser's `respond`-to-a-dropped-receiver error.
1520        let (algo, serve) = hedge(Some(std::time::Duration::from_millis(50)));
1521        let (_, response) = test_drive(algo, request(), serve).await?;
1522        assert_eq!(
1523            response
1524                .llm_response
1525                .as_agg()
1526                .map(completion_text)
1527                .unwrap_or_default(),
1528            "winner"
1529        );
1530        Ok(())
1531    }
1532
1533    #[tokio::test]
1534    async fn run_returns_the_winner_without_hanging_on_a_pending_loser() -> Result<()> {
1535        // The loser never resolves. `run` must return the winner promptly, not hang
1536        // waiting for the in-flight loser.
1537        let (algo, serve) = hedge(None);
1538        let run = test_drive(algo, request(), serve);
1539        let (_, response) = tokio::time::timeout(std::time::Duration::from_secs(1), run)
1540            .await
1541            .map_err(|error| LibsyError::external("waiting for pending loser", error))??;
1542        assert_eq!(
1543            response
1544                .llm_response
1545                .as_agg()
1546                .map(completion_text)
1547                .unwrap_or_default(),
1548            "winner"
1549        );
1550        Ok(())
1551    }
1552
1553    #[tokio::test]
1554    async fn run_surfaces_a_terminal_error_with_many_calls_in_flight() -> Result<()> {
1555        use std::sync::atomic::{AtomicUsize, Ordering};
1556
1557        // A large fan-out (10 matched the old, now-removed concurrency cap). The terminal
1558        // error must still reach the caller with all of these calls pending.
1559        const N: usize = 10;
1560
1561        // Fans out N calls, then errors as soon as all N are in flight — exercising a
1562        // terminal failure emitted while the offloaded calls are still pending.
1563        struct FanOutThenError {
1564            all_started: Arc<tokio::sync::Notify>,
1565            n: usize,
1566        }
1567
1568        #[async_trait]
1569        impl Algorithm for FanOutThenError {
1570            fn name(&self) -> &str {
1571                "fan_out_then_error"
1572            }
1573
1574            async fn route(
1575                self: Arc<Self>,
1576                driver: Driver,
1577                request: Request,
1578            ) -> Result<RoutingOutcome> {
1579                let offloads = futures::future::join_all(
1580                    (0..self.n)
1581                        .map(|i| driver.call_model(request.clone(), vec![format!("m{i}").into()])),
1582                );
1583                tokio::select! {
1584                    _ = offloads => Err(test_error("offloads unexpectedly completed")),
1585                    _ = self.all_started.notified() => {
1586                        Err(test_error("terminal error while calls pending"))
1587                    }
1588                }
1589            }
1590        }
1591
1592        let all_started = Arc::new(tokio::sync::Notify::new());
1593        let algo: Arc<dyn Algorithm> = Arc::new(FanOutThenError {
1594            all_started: all_started.clone(),
1595            n: N,
1596        });
1597
1598        // Serving enters each call; once all N are in flight it signals, then pends forever.
1599        let started = Arc::new(AtomicUsize::new(0));
1600        let serve = move |_target: ModelId, _request: Request| {
1601            let started = started.clone();
1602            let all_started = all_started.clone();
1603            async move {
1604                if started.fetch_add(1, Ordering::SeqCst) + 1 == N {
1605                    all_started.notify_one();
1606                }
1607                std::future::pending::<ServeResult>().await
1608            }
1609        };
1610
1611        // With the cap gone, the driver keeps polling the stream even with N calls in
1612        // flight, so the terminal error surfaces promptly instead of hanging.
1613        let run = test_drive(algo, request(), serve);
1614        let result = tokio::time::timeout(std::time::Duration::from_millis(500), run)
1615            .await
1616            .map_err(|error| {
1617                LibsyError::external("waiting for terminal error with full call cap", error)
1618            })?;
1619        match result {
1620            Ok(_) => Err(test_error("expected the terminal error, got a response")),
1621            Err(err) => {
1622                assert!(
1623                    err.to_string()
1624                        .contains("terminal error while calls pending")
1625                );
1626                Ok(())
1627            }
1628        }
1629    }
1630}