Skip to main content

switchyard_libsy/core/
algorithm.rs

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