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