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