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    /// Replay locally stored messages before routing a Responses continuation.
623    fn needs_history_replay(&self, _request: &Request) -> bool {
624        false
625    }
626
627    /// Select a route, requesting external work through [`Driver`] as needed.
628    /// [`run_stream`](Self::run_stream) runs this method and exposes its work to the host.
629    async fn route(self: Arc<Self>, driver: Driver, request: Request) -> Result<RoutingOutcome>;
630
631    /// Expose work requests and the final outcome so the host can control execution.
632    ///
633    /// Each call waits for its host reply. A run ends with one terminal item — [`Step::Done`] on
634    /// success, an `Err` item on failure, including when the algorithm panics. Dropping
635    /// the stream aborts the spawned algorithm task.
636    ///
637    /// Every invocation owns a separate [`Driver`].
638    fn run_stream(self: Arc<Self>, request: Request, models: Arc<RuntimeModels>) -> StepStream {
639        let (driver, step_rx) = Driver::new(self.name(), models);
640        let span = observability::run_span(self.name(), &request);
641        let handle = tokio::spawn(
642            async move {
643                let algorithm = self.name().to_string();
644                // Catch a panicking algorithm so the run still publishes a terminal step.
645                let route = AssertUnwindSafe(self.route(driver.clone(), request)).catch_unwind();
646                let result = observability::observe_run(&algorithm, async move {
647                    route.await.unwrap_or_else(|payload| {
648                        Err(LibsyError::AlgorithmError {
649                            message: format!(
650                                "algorithm task panicked: {}",
651                                panic_message(payload.as_ref())
652                            ),
653                        })
654                    })
655                })
656                .await;
657
658                let _ = driver.finish(result).await;
659            }
660            .instrument(span),
661        );
662        // Dropping the stream aborts the algorithm task when its consumer goes away.
663        let abort_guard = AbortOnDrop(handle.abort_handle());
664        Box::pin(ReceiverStream::new(step_rx).map(move |step| {
665            // link abort guard to stream
666            let _keep_alive = &abort_guard;
667            step
668        }))
669    }
670}
671
672#[cfg(test)]
673mod tests {
674    use std::collections::HashMap;
675
676    use super::*;
677    use crate::core::testing::{Serve, ServeResult, echo, reply, serve_decision, test_drive};
678    use futures::StreamExt;
679    use switchyard_protocol::{
680        LlmResponse, LlmResponseChunk, completion_text, text_request, text_response,
681    };
682
683    #[derive(Debug, thiserror::Error)]
684    #[error("{0}")]
685    struct TestError(&'static str);
686
687    fn test_error(message: &'static str) -> LibsyError {
688        LibsyError::external("test", TestError(message))
689    }
690
691    /// Trivial algo used only to exercise the orchestrator: calls the first target
692    /// and returns its response as the routing outcome.
693    struct TestAlgo {
694        target_set: Vec<ModelId>,
695    }
696
697    #[async_trait]
698    impl Algorithm for TestAlgo {
699        fn name(&self) -> &str {
700            "test"
701        }
702
703        async fn route(
704            self: Arc<Self>,
705            driver: Driver,
706            request: Request,
707        ) -> Result<RoutingOutcome> {
708            let target = self
709                .target_set
710                .first()
711                .ok_or(LibsyError::NoTargets)?
712                .clone();
713            let response = driver
714                .call_model(request.clone(), vec![target.clone()])
715                .await?;
716            driver.set_evidence(serde_json::json!({"source": "test"}));
717            driver.set_evidence_if_empty(serde_json::json!({"source": "ignored"}));
718            Ok(RoutingOutcome::answered(target, request, response))
719        }
720    }
721
722    /// Build a shared `TestAlgo` over the given target set.
723    fn orch(target_set: Vec<ModelId>) -> Arc<dyn Algorithm> {
724        Arc::new(TestAlgo { target_set })
725    }
726
727    fn request() -> Request {
728        Request {
729            llm_request: text_request(Some("auto".to_string()), "hi".to_string()),
730            raw_request: None,
731            metadata: None,
732        }
733    }
734
735    #[test]
736    fn routing_outcome_constructors_stamp_selection_and_preserve_payloads() {
737        let outcome = RoutingOutcome::route_to(
738            "selected".into(),
739            target_set(&["fallback-one", "fallback-two"]),
740            request(),
741        );
742
743        assert_eq!(
744            outcome.selected_model_ids,
745            target_set(&["selected", "fallback-one", "fallback-two"])
746        );
747        assert_eq!(outcome.request.model_id().as_deref(), Some("selected"));
748        assert!(outcome.response.is_none());
749        assert!(outcome.metadata.is_none());
750
751        let outcome = RoutingOutcome::route_to("only".into(), Vec::new(), request());
752        assert_eq!(outcome.selected_model_ids, target_set(&["only"]));
753
754        let outcome = RoutingOutcome::answered(
755            "answered".into(),
756            request(),
757            Response {
758                llm_response: LlmResponse::Agg(text_response(None, "existing")),
759                metadata: None,
760                upstream_headers: http::HeaderMap::new(),
761            },
762        );
763
764        assert_eq!(outcome.selected_model_ids, target_set(&["answered"]));
765        assert_eq!(outcome.request.model_id().as_deref(), Some("answered"));
766        assert_eq!(
767            outcome
768                .response
769                .as_ref()
770                .and_then(|response| response.llm_response.as_agg())
771                .map(completion_text),
772            Some("existing".to_string())
773        );
774    }
775
776    fn target_set(names: &[&str]) -> Vec<ModelId> {
777        names.iter().map(|name| ModelId::from(*name)).collect()
778    }
779
780    #[tokio::test]
781    async fn typed_driver_preserves_call_and_stream_boundaries() -> Result<()> {
782        tokio::time::timeout(std::time::Duration::from_secs(1), async {
783            // Distinct oneshots keep reverse-order replies paired with their producers, and a
784            // retained call remains pending until the host responds.
785            let (driver, mut step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
786            let first_driver = driver.clone();
787            let mut first = tokio::spawn(async move {
788                first_driver
789                    .call_model(request(), vec![ModelId::from("first")])
790                    .await
791            });
792            let second = tokio::spawn(async move {
793                driver
794                    .call_model(request(), vec![ModelId::from("second")])
795                    .await
796            });
797
798            let mut calls = HashMap::new();
799            for _ in 0..2 {
800                let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
801                let Step::CallModel(call) = step else {
802                    return Err(test_error("expected a CallModel step"));
803                };
804                let selected_model = call
805                    .models
806                    .first()
807                    .ok_or_else(|| test_error("model call has no candidates"))?
808                    .to_string();
809                calls.insert(selected_model, call);
810            }
811            assert!(
812                tokio::time::timeout(std::time::Duration::from_millis(20), &mut first)
813                    .await
814                    .is_err(),
815                "call completed before the host responded"
816            );
817            calls
818                .remove("second")
819                .ok_or_else(|| test_error("missing second call"))?
820                .respond(Ok(reply("second response")))?;
821            calls
822                .remove("first")
823                .ok_or_else(|| test_error("missing first call"))?
824                .respond(Ok(reply("first response")))?;
825
826            let first_response = first
827                .await
828                .map_err(|source| LibsyError::external("joining a test task", source))??;
829            let second_response = second
830                .await
831                .map_err(|source| LibsyError::external("joining a test task", source))??;
832            assert_eq!(
833                first_response.llm_response.as_agg().map(completion_text),
834                Some("first response".to_string())
835            );
836            assert_eq!(
837                second_response.llm_response.as_agg().map(completion_text),
838                Some("second response".to_string())
839            );
840
841            // Dropping the host-facing promise closes only that call's reply channel.
842            let (driver, mut step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
843            let producer = tokio::spawn(async move {
844                driver
845                    .call_model(request(), vec![ModelId::from("dropped")])
846                    .await
847            });
848            let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
849            let Step::CallModel(call) = step else {
850                return Err(test_error("expected a CallModel step"));
851            };
852            drop(call);
853            let result = producer
854                .await
855                .map_err(|source| LibsyError::external("joining a test task", source))?;
856            assert!(matches!(
857                result,
858                Err(LibsyError::Driver(DriverError::ResponseDropped))
859            ));
860
861            // A standalone driver reports the typed step receiver disappearing at its next call.
862            let (driver, step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
863            drop(step_rx);
864            let result = driver
865                .call_model(request(), vec![ModelId::from("closed")])
866                .await;
867            assert!(matches!(
868                result,
869                Err(LibsyError::Driver(DriverError::StreamClosed))
870            ));
871
872            fn decision_request() -> DecisionRequest {
873                DecisionRequest {
874                    model: Some("overwritten".into()),
875                    context: serde_json::json!({"task": "choose a route"}),
876                    questions: Default::default(),
877                }
878            }
879
880            fn decision_response() -> DecisionResponse {
881                DecisionResponse {
882                    id: Some("decision-1".to_string()),
883                    model: Some("provider-model".into()),
884                    answers: [(
885                        "p_solve".to_string(),
886                        switchyard_protocol::DecisionAnswer {
887                            value: switchyard_protocol::DecisionValue::Boolean(
888                                switchyard_protocol::BooleanEstimate::ProbabilityTrue(
889                                    switchyard_protocol::Probability(0.8),
890                                ),
891                            ),
892                            provider_confidence: None,
893                        },
894                    )]
895                    .into(),
896                    usage: Default::default(),
897                }
898            }
899
900            struct MixedCalls(&'static str);
901
902            #[async_trait]
903            impl Algorithm for MixedCalls {
904                fn name(&self) -> &str {
905                    "mixed"
906                }
907
908                async fn route(
909                    self: Arc<Self>,
910                    driver: Driver,
911                    request: Request,
912                ) -> Result<RoutingOutcome> {
913                    let (llm, decision) = tokio::join!(
914                        driver.call_model(request.clone(), vec!["llm".into()]),
915                        driver.call_decision(decision_request(), "decision".into()),
916                    );
917                    assert_eq!(
918                        llm?.llm_response.as_agg().map(completion_text),
919                        Some("llm reply".into())
920                    );
921                    match self.0 {
922                        "mock" => assert_eq!(decision?, DecisionResponse {
923                            id: None,
924                            model: Some("decision".into()),
925                            answers: Default::default(),
926                            usage: Default::default(),
927                        }),
928                        "reply" => assert_eq!(decision?, decision_response()),
929                        "error" => assert!(matches!(
930                            decision,
931                            Err(LibsyError::AlgorithmError { message }) if message == "provider failed"
932                        )),
933                        "drop" => assert!(matches!(
934                            decision,
935                            Err(LibsyError::Driver(DriverError::ResponseDropped))
936                        )),
937                        "abort" => return std::future::pending().await,
938                        _ => unreachable!(),
939                    }
940                    Ok(RoutingOutcome::route_to("answer".into(), vec![], request))
941                }
942            }
943
944            for mode in ["mock", "reply", "error", "drop", "abort"] {
945                // Serial dispatch would deadlock here and fail the enclosing timeout.
946                let barrier = Arc::new(tokio::sync::Barrier::new(2));
947                let outcome = drive(
948                    Arc::new(MixedCalls(mode)),
949                    request(),
950                    Arc::new(RuntimeModels::default()),
951                    move |call| {
952                        let barrier = barrier.clone();
953                        async move {
954                            barrier.wait().await;
955                            let call = match call {
956                                Call::Model(call) => return call.respond(Ok(reply("llm reply"))),
957                                Call::Decision(call) => *call,
958                            };
959                            assert_eq!(call.algorithm, "mixed");
960                            assert_eq!(call.model, "decision");
961                            assert_eq!(call.request.model, Some("decision".into()));
962                            assert_eq!(call.request.context, decision_request().context);
963                            match mode {
964                                "mock" => serve_decision(call).await,
965                                "reply" => call.respond(Ok(decision_response())),
966                                "error" => call.respond(Err(LibsyError::AlgorithmError {
967                                    message: "provider failed".into(),
968                                })),
969                                "drop" => {
970                                    drop(call);
971                                    Ok(())
972                                }
973                                "abort" => call.fail(test_error("host aborted")),
974                                _ => unreachable!(),
975                            }
976                        }
977                    },
978                )
979                .await;
980                if mode == "abort" {
981                    assert!(matches!(outcome, Err(LibsyError::External { .. })));
982                } else {
983                    assert_eq!(outcome?.selected_model_id()?, "answer");
984                }
985            }
986
987            let (driver, mut steps) = Driver::new("test", Arc::new(RuntimeModels::default()));
988            let mut pending = Box::pin(driver.call_decision(decision_request(), "decision".into()));
989            assert!(futures::poll!(&mut pending).is_pending());
990            let Some(Ok(Step::CallDecision(call))) = steps.recv().await else {
991                return Err(test_error("expected a decision call"));
992            };
993            drop(pending);
994            assert!(matches!(
995                call.respond(Ok(decision_response())),
996                Err(LibsyError::Driver(DriverError::ResponseDropped))
997            ));
998            drop(steps);
999            assert!(matches!(
1000                driver.call_decision(decision_request(), "decision".into()).await,
1001                Err(LibsyError::Driver(DriverError::StreamClosed))
1002            ));
1003            Ok(())
1004        })
1005        .await
1006        .map_err(|error| LibsyError::external("waiting for typed driver boundaries", error))?
1007    }
1008
1009    #[test]
1010    fn target_lookup_returns_the_missing_target() {
1011        let error = ensure_model_is_target(&target_set(&[]), &ModelId::from("missing")).err();
1012        assert!(matches!(
1013            error,
1014            Some(LibsyError::TargetNotFound { target }) if target == "missing"
1015        ));
1016    }
1017
1018    /// Build a single-target algo, plus a `serve` that answers it as a token stream
1019    /// replaying `chunks` in order (as `Ok` items).
1020    fn streaming_orch(chunks: Vec<LlmResponseChunk>) -> (Arc<dyn Algorithm>, impl Serve) {
1021        let algo = orch(target_set(&["stream/model"]));
1022        let serve = move |_target: ModelId, _request: Request| {
1023            let chunks = chunks.clone();
1024            async move {
1025                let stream =
1026                    futures::stream::iter(chunks.into_iter().map(|chunk| Ok(chunk.into()))).boxed();
1027                Ok(Response {
1028                    llm_response: LlmResponse::Stream(stream),
1029                    metadata: None,
1030                    upstream_headers: http::HeaderMap::new(),
1031                })
1032            }
1033        };
1034        (algo, serve)
1035    }
1036
1037    #[tokio::test]
1038    async fn run_returns_a_streamed_response_the_caller_aggregates() -> Result<()> {
1039        // A streaming client -> its chunks flow through the promise and `Done`,
1040        // and `run` returns the live stream untouched for the caller to fold.
1041        let (orch, serve) = streaming_orch(vec![
1042            LlmResponseChunk::MessageStart {
1043                id: Some("m1".to_string()),
1044                model: Some("stream/model".to_string()),
1045            },
1046            LlmResponseChunk::TextDelta {
1047                index: 0,
1048                text: "hel".to_string(),
1049            },
1050            LlmResponseChunk::TextDelta {
1051                index: 0,
1052                text: "lo".to_string(),
1053            },
1054            LlmResponseChunk::MessageStop {
1055                reason: Some("stop".to_string()),
1056            },
1057        ]);
1058        let (selected_model, response) = test_drive(orch, request(), serve).await?;
1059        // The run handed back the live stream; the caller folds it to a buffered aggregate.
1060        let agg = response
1061            .llm_response
1062            .into_agg()
1063            .await
1064            .map_err(|error| LibsyError::external("aggregating response stream", error))?;
1065        assert_eq!(completion_text(&agg), "hello");
1066        assert_eq!(agg.model.as_deref(), Some("stream/model"));
1067        assert_eq!(selected_model, "stream/model");
1068        Ok(())
1069    }
1070
1071    #[tokio::test]
1072    async fn aggregating_a_streamed_response_propagates_a_mid_stream_error() -> Result<()> {
1073        // The run succeeds and returns the stream; the in-band `Error` chunk surfaces only
1074        // when the caller aggregates it.
1075        let (orch, serve) = streaming_orch(vec![
1076            LlmResponseChunk::TextDelta {
1077                index: 0,
1078                text: "partial".to_string(),
1079            },
1080            LlmResponseChunk::StreamError {
1081                message: "upstream exploded".to_string(),
1082            },
1083        ]);
1084        let (_, response) = test_drive(orch, request(), serve).await?;
1085        match response.llm_response.into_agg().await {
1086            Ok(_) => panic!("expected a mid-stream error, got an aggregate"),
1087            Err(err) => {
1088                assert!(err.to_string().contains("upstream exploded"));
1089                Ok(())
1090            }
1091        }
1092    }
1093
1094    #[tokio::test]
1095    async fn run_offloads_via_promise_then_finishes() -> Result<()> {
1096        // Every call is offloaded via a promise the orchestrator surfaces as a
1097        // `CallModel` step for us to fulfill.
1098        let stream = orch(target_set(&["offload/model"]))
1099            .run_stream(request(), Arc::new(RuntimeModels::default()));
1100        tokio::pin!(stream);
1101
1102        let mut saw_call = false;
1103        let mut final_completion = None;
1104        while let Some(step) = stream.next().await {
1105            match step? {
1106                Step::CallDecision(_) => return Err(test_error("unexpected decision call")),
1107                Step::CallModel(call) => {
1108                    saw_call = true;
1109                    assert_eq!(call.models, vec![ModelId::from("offload/model")]);
1110                    // Fulfilling the promise is the "real" model call the caller makes.
1111                    call.respond(Ok(Response {
1112                        llm_response: LlmResponse::Agg(text_response(
1113                            None,
1114                            "fulfilled".to_string(),
1115                        )),
1116                        metadata: None,
1117                        upstream_headers: http::HeaderMap::new(),
1118                    }))?;
1119                }
1120                Step::Done(outcome) => {
1121                    let metadata = outcome
1122                        .metadata
1123                        .as_ref()
1124                        .expect("run_stream should attach outcome metadata");
1125                    assert_eq!(metadata.algorithm, "test");
1126                    assert_eq!(
1127                        uuid::Uuid::parse_str(metadata.outcome_id())
1128                            .expect("outcome id should be a UUID")
1129                            .get_version_num(),
1130                        7
1131                    );
1132                    assert_eq!(
1133                        metadata.evidence,
1134                        Some(serde_json::json!({"source": "test"}))
1135                    );
1136                    let response = outcome
1137                        .response
1138                        .ok_or_else(|| test_error("expected an answered outcome"))?;
1139                    final_completion = Some(
1140                        response
1141                            .llm_response
1142                            .as_agg()
1143                            .map(completion_text)
1144                            .unwrap_or_default(),
1145                    );
1146                }
1147            }
1148        }
1149
1150        assert!(saw_call, "expected a CallModel step before Done");
1151        assert_eq!(
1152            final_completion.ok_or_else(|| test_error("no Done step"))?,
1153            "fulfilled"
1154        );
1155        Ok(())
1156    }
1157
1158    #[tokio::test(flavor = "multi_thread", worker_threads = 12)]
1159    async fn requests_are_processed_in_parallel() -> Result<()> {
1160        use std::time::Duration;
1161        use tokio::sync::Barrier;
1162
1163        const N: usize = 12;
1164
1165        // Serving blocks until all N concurrent calls have arrived. If requests were
1166        // serialized (one algorithm behind a `Mutex`), only one call could be in flight,
1167        // the barrier would never reach N, and the test would time out. It passes only
1168        // because the shared algorithm is driven concurrently across requests.
1169        let barrier = Arc::new(Barrier::new(N));
1170        // One shared algorithm driven by many concurrent requests.
1171        let algo = orch(target_set(&["m"]));
1172
1173        let mut handles = Vec::new();
1174        for _ in 0..N {
1175            let algo = algo.clone();
1176            let barrier = barrier.clone();
1177            let serve = move |target: ModelId, _request: Request| {
1178                let barrier = barrier.clone();
1179                async move {
1180                    barrier.wait().await;
1181                    Ok(reply(target))
1182                }
1183            };
1184            handles.push(tokio::spawn(async move {
1185                test_drive(algo, request(), serve)
1186                    .await
1187                    .map(|(_, response)| {
1188                        response
1189                            .llm_response
1190                            .as_agg()
1191                            .map(completion_text)
1192                            .unwrap_or_default()
1193                    })
1194            }));
1195        }
1196
1197        for handle in handles {
1198            // The timeout turns a serialization deadlock into a failure, not a hang.
1199            let completion = tokio::time::timeout(Duration::from_secs(5), handle)
1200                .await
1201                .map_err(|error| LibsyError::external("waiting for test task", error))?
1202                .map_err(|source| LibsyError::external("joining a test task", source))??;
1203            assert_eq!(completion, "m");
1204        }
1205        Ok(())
1206    }
1207
1208    #[tokio::test]
1209    async fn offload_error_propagates_back_to_the_algorithm() -> Result<()> {
1210        // A client-less target offloads its call; we fulfill the promise with an
1211        // Err, which must flow back through `call_model_target` into the algorithm and
1212        // out as an error step — not a response.
1213        let stream = orch(target_set(&["offload/model"]))
1214            .run_stream(request(), Arc::new(RuntimeModels::default()));
1215        tokio::pin!(stream);
1216
1217        let mut saw_error = false;
1218        while let Some(step) = stream.next().await {
1219            match step {
1220                Ok(Step::CallDecision(_)) => return Err(test_error("unexpected decision call")),
1221                Ok(Step::CallModel(call)) => {
1222                    call.respond(Err(test_error("upstream model call failed")))?;
1223                }
1224                Ok(Step::Done(..)) => {
1225                    return Err(test_error(
1226                        "expected the offload error to propagate, got a response",
1227                    ));
1228                }
1229                Err(err) => {
1230                    // The algorithm's `call_model_target` saw the error via the promise.
1231                    assert!(err.to_string().contains("upstream model call failed"));
1232                    saw_error = true;
1233                }
1234            }
1235        }
1236
1237        assert!(saw_error, "expected an error step");
1238        Ok(())
1239    }
1240
1241    #[tokio::test]
1242    async fn dropping_the_stream_cancels_the_algorithm_task() -> Result<()> {
1243        use std::sync::atomic::{AtomicBool, Ordering};
1244        use std::time::Duration;
1245        use tokio::sync::mpsc;
1246
1247        // Sets a flag when dropped, so we can observe whether the algorithm task was
1248        // cancelled/dropped.
1249        struct DropGuard(Arc<AtomicBool>);
1250        impl Drop for DropGuard {
1251            fn drop(&mut self) {
1252                self.0.store(true, Ordering::SeqCst);
1253            }
1254        }
1255
1256        struct StuckAlgo {
1257            started: mpsc::UnboundedSender<()>,
1258            dropped: Arc<AtomicBool>,
1259        }
1260
1261        #[async_trait]
1262        impl Algorithm for StuckAlgo {
1263            fn name(&self) -> &str {
1264                "stuck"
1265            }
1266
1267            async fn route(
1268                self: Arc<Self>,
1269                _driver: Driver,
1270                _request: Request,
1271            ) -> Result<RoutingOutcome> {
1272                let _guard = DropGuard(self.dropped.clone());
1273                let _ = self.started.send(());
1274                // Await forever without ever touching the driver.
1275                std::future::pending::<()>().await;
1276                unreachable!()
1277            }
1278        }
1279
1280        let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1281        let dropped = Arc::new(AtomicBool::new(false));
1282        let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1283            started: started_tx,
1284            dropped: dropped.clone(),
1285        });
1286
1287        let stream = algo.run_stream(request(), Arc::new(RuntimeModels::default()));
1288        started_rx
1289            .recv()
1290            .await
1291            .ok_or_else(|| test_error("task never started"))?;
1292        drop(stream);
1293        tokio::time::sleep(Duration::from_millis(100)).await;
1294
1295        assert!(
1296            dropped.load(Ordering::SeqCst),
1297            "algorithm task was NOT cancelled after dropping the stream"
1298        );
1299        Ok(())
1300    }
1301
1302    #[tokio::test]
1303    async fn route_panic_surfaces_as_a_stream_error() -> Result<()> {
1304        // An algorithm whose task panics must surface an `Err` step carrying the panic
1305        // message, not abort the process from an unobserved detached task.
1306        struct Panicky;
1307
1308        #[async_trait]
1309        impl Algorithm for Panicky {
1310            fn name(&self) -> &str {
1311                "panicky"
1312            }
1313
1314            async fn route(
1315                self: Arc<Self>,
1316                _driver: Driver,
1317                _request: Request,
1318            ) -> Result<RoutingOutcome> {
1319                panic!("boom");
1320            }
1321        }
1322
1323        let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
1324        let stream = algo.run_stream(request(), Arc::new(RuntimeModels::default()));
1325        tokio::pin!(stream);
1326
1327        let mut saw_error = false;
1328        while let Some(step) = stream.next().await {
1329            match step {
1330                Err(err) => {
1331                    // The panic message is preserved, not flattened into an opaque failure.
1332                    assert!(err.to_string().contains("algorithm task panicked: boom"));
1333                    saw_error = true;
1334                }
1335                Ok(_) => return Err(test_error("expected the panic to surface as an error step")),
1336            }
1337        }
1338
1339        assert!(saw_error, "expected an error step from the panicked task");
1340        Ok(())
1341    }
1342
1343    /// A panicking algorithm must publish its terminal step even when it left a `Driver`
1344    /// clone alive in another task. That clone holds the step channel open, so a run that
1345    /// merely unwound would never terminate and the consumer would wait forever.
1346    #[tokio::test]
1347    async fn a_panic_with_a_leaked_driver_clone_still_terminates_the_run() -> Result<()> {
1348        struct LeakyPanic;
1349
1350        #[async_trait]
1351        impl Algorithm for LeakyPanic {
1352            fn name(&self) -> &str {
1353                "leaky_panic"
1354            }
1355
1356            async fn route(
1357                self: Arc<Self>,
1358                driver: Driver,
1359                _request: Request,
1360            ) -> Result<RoutingOutcome> {
1361                tokio::spawn(async move {
1362                    // Outlives the panic below, keeping a sender clone alive.
1363                    let _keep_alive = driver;
1364                    std::future::pending::<()>().await;
1365                });
1366                tokio::task::yield_now().await;
1367                panic!("boom");
1368            }
1369        }
1370
1371        let algo: Arc<dyn Algorithm> = Arc::new(LeakyPanic);
1372        // The timeout turns the hang this guards against into a failure rather than a hang.
1373        let result = tokio::time::timeout(
1374            std::time::Duration::from_secs(1),
1375            test_drive(algo, request(), echo()),
1376        )
1377        .await
1378        .map_err(|error| LibsyError::external("waiting for the panicked run to end", error))?;
1379
1380        match result {
1381            Ok(_) => Err(test_error(
1382                "expected the panic to end the run with an error",
1383            )),
1384            Err(err) => {
1385                assert!(err.to_string().contains("algorithm task panicked: boom"));
1386                Ok(())
1387            }
1388        }
1389    }
1390
1391    #[tokio::test]
1392    async fn cancelling_run_cancels_the_algorithm_task() -> Result<()> {
1393        use std::sync::atomic::{AtomicBool, Ordering};
1394        use std::time::Duration;
1395        use tokio::sync::mpsc;
1396
1397        // Sets a flag when dropped, so we can observe whether the algorithm task was
1398        // cancelled once the `run` future driving it is dropped.
1399        struct DropGuard(Arc<AtomicBool>);
1400        impl Drop for DropGuard {
1401            fn drop(&mut self) {
1402                self.0.store(true, Ordering::SeqCst);
1403            }
1404        }
1405
1406        struct StuckAlgo {
1407            started: mpsc::UnboundedSender<()>,
1408            dropped: Arc<AtomicBool>,
1409        }
1410
1411        #[async_trait]
1412        impl Algorithm for StuckAlgo {
1413            fn name(&self) -> &str {
1414                "stuck"
1415            }
1416
1417            async fn route(
1418                self: Arc<Self>,
1419                _driver: Driver,
1420                _request: Request,
1421            ) -> Result<RoutingOutcome> {
1422                let _guard = DropGuard(self.dropped.clone());
1423                let _ = self.started.send(());
1424                // Hang forever without ever touching the driver, so only cancellation
1425                // (not a dropped step channel) can stop this task.
1426                std::future::pending::<()>().await;
1427                unreachable!()
1428            }
1429        }
1430
1431        let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1432        let dropped = Arc::new(AtomicBool::new(false));
1433        let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1434            started: started_tx,
1435            dropped: dropped.clone(),
1436        });
1437
1438        // Drive the run on its own task, wait until the algorithm task is up, then cancel
1439        // it — dropping its future (and the `run_stream` stream it holds).
1440        let run_task = tokio::spawn(async move { test_drive(algo, request(), echo()).await });
1441        started_rx
1442            .recv()
1443            .await
1444            .ok_or_else(|| test_error("task never started"))?;
1445        run_task.abort();
1446        tokio::time::sleep(Duration::from_millis(100)).await;
1447
1448        assert!(
1449            dropped.load(Ordering::SeqCst),
1450            "algorithm task was NOT cancelled after cancelling run"
1451        );
1452        Ok(())
1453    }
1454
1455    // --- first-wins hedging: `run` must not wait on losing speculative calls -------------
1456
1457    /// Offloads two targets concurrently and returns the first to resolve, dropping the
1458    /// loser's call (first-wins hedging).
1459    struct Hedge {
1460        winner: String,
1461        loser: String,
1462    }
1463
1464    #[async_trait]
1465    impl Algorithm for Hedge {
1466        fn name(&self) -> &str {
1467            "hedge"
1468        }
1469
1470        async fn route(
1471            self: Arc<Self>,
1472            driver: Driver,
1473            request: Request,
1474        ) -> Result<RoutingOutcome> {
1475            let outcome_request = request.clone();
1476            let win = driver.call_model(request.clone(), vec![self.winner.clone().into()]);
1477            let lose = driver.call_model(request, vec![self.loser.clone().into()]);
1478            // First to resolve wins; `select!` drops the losing future (and its promise).
1479            tokio::select! {
1480                res = win => Ok(RoutingOutcome::answered(
1481                    self.winner.clone().into(),
1482                    outcome_request,
1483                    res?,
1484                )),
1485                res = lose => Ok(RoutingOutcome::answered(
1486                    self.loser.clone().into(),
1487                    outcome_request,
1488                    res?,
1489                )),
1490            }
1491        }
1492    }
1493
1494    /// Builds a hedging algo and the `serve` that drives it: the winner is gated behind the
1495    /// loser starting (so the loser's serve is guaranteed in flight when the winner wins),
1496    /// and the loser finishes after `loser_delay` — or never, when `None`.
1497    fn hedge(loser_delay: Option<std::time::Duration>) -> (Arc<dyn Algorithm>, impl Serve) {
1498        let started = Arc::new(tokio::sync::Notify::new());
1499        let algo = Arc::new(Hedge {
1500            winner: "winner".to_string(),
1501            loser: "loser".to_string(),
1502        });
1503        let serve = move |target: ModelId, _request: Request| {
1504            let started = started.clone();
1505            async move {
1506                if target == "loser" {
1507                    started.notify_one();
1508                    match loser_delay {
1509                        Some(delay) => tokio::time::sleep(delay).await,
1510                        None => std::future::pending::<()>().await,
1511                    }
1512                } else {
1513                    started.notified().await;
1514                }
1515                Ok(reply(target))
1516            }
1517        };
1518        (algo, serve)
1519    }
1520
1521    #[tokio::test]
1522    async fn run_returns_the_winner_without_a_late_loser_overwriting_it() -> Result<()> {
1523        // The loser responds 50ms after the winner has already won. `run` must return the
1524        // winner, not the loser's `respond`-to-a-dropped-receiver error.
1525        let (algo, serve) = hedge(Some(std::time::Duration::from_millis(50)));
1526        let (_, response) = test_drive(algo, request(), serve).await?;
1527        assert_eq!(
1528            response
1529                .llm_response
1530                .as_agg()
1531                .map(completion_text)
1532                .unwrap_or_default(),
1533            "winner"
1534        );
1535        Ok(())
1536    }
1537
1538    #[tokio::test]
1539    async fn run_returns_the_winner_without_hanging_on_a_pending_loser() -> Result<()> {
1540        // The loser never resolves. `run` must return the winner promptly, not hang
1541        // waiting for the in-flight loser.
1542        let (algo, serve) = hedge(None);
1543        let run = test_drive(algo, request(), serve);
1544        let (_, response) = tokio::time::timeout(std::time::Duration::from_secs(1), run)
1545            .await
1546            .map_err(|error| LibsyError::external("waiting for pending loser", error))??;
1547        assert_eq!(
1548            response
1549                .llm_response
1550                .as_agg()
1551                .map(completion_text)
1552                .unwrap_or_default(),
1553            "winner"
1554        );
1555        Ok(())
1556    }
1557
1558    #[tokio::test]
1559    async fn run_surfaces_a_terminal_error_with_many_calls_in_flight() -> Result<()> {
1560        use std::sync::atomic::{AtomicUsize, Ordering};
1561
1562        // A large fan-out (10 matched the old, now-removed concurrency cap). The terminal
1563        // error must still reach the caller with all of these calls pending.
1564        const N: usize = 10;
1565
1566        // Fans out N calls, then errors as soon as all N are in flight — exercising a
1567        // terminal failure emitted while the offloaded calls are still pending.
1568        struct FanOutThenError {
1569            all_started: Arc<tokio::sync::Notify>,
1570            n: usize,
1571        }
1572
1573        #[async_trait]
1574        impl Algorithm for FanOutThenError {
1575            fn name(&self) -> &str {
1576                "fan_out_then_error"
1577            }
1578
1579            async fn route(
1580                self: Arc<Self>,
1581                driver: Driver,
1582                request: Request,
1583            ) -> Result<RoutingOutcome> {
1584                let offloads = futures::future::join_all(
1585                    (0..self.n)
1586                        .map(|i| driver.call_model(request.clone(), vec![format!("m{i}").into()])),
1587                );
1588                tokio::select! {
1589                    _ = offloads => Err(test_error("offloads unexpectedly completed")),
1590                    _ = self.all_started.notified() => {
1591                        Err(test_error("terminal error while calls pending"))
1592                    }
1593                }
1594            }
1595        }
1596
1597        let all_started = Arc::new(tokio::sync::Notify::new());
1598        let algo: Arc<dyn Algorithm> = Arc::new(FanOutThenError {
1599            all_started: all_started.clone(),
1600            n: N,
1601        });
1602
1603        // Serving enters each call; once all N are in flight it signals, then pends forever.
1604        let started = Arc::new(AtomicUsize::new(0));
1605        let serve = move |_target: ModelId, _request: Request| {
1606            let started = started.clone();
1607            let all_started = all_started.clone();
1608            async move {
1609                if started.fetch_add(1, Ordering::SeqCst) + 1 == N {
1610                    all_started.notify_one();
1611                }
1612                std::future::pending::<ServeResult>().await
1613            }
1614        };
1615
1616        // With the cap gone, the driver keeps polling the stream even with N calls in
1617        // flight, so the terminal error surfaces promptly instead of hanging.
1618        let run = test_drive(algo, request(), serve);
1619        let result = tokio::time::timeout(std::time::Duration::from_millis(500), run)
1620            .await
1621            .map_err(|error| {
1622                LibsyError::external("waiting for terminal error with full call cap", error)
1623            })?;
1624        match result {
1625            Ok(_) => Err(test_error("expected the terminal error, got a response")),
1626            Err(err) => {
1627                assert!(
1628                    err.to_string()
1629                        .contains("terminal error while calls pending")
1630                );
1631                Ok(())
1632            }
1633        }
1634    }
1635}