1use std::{future::Future, panic::AssertUnwindSafe, pin::Pin, sync::Arc, time::Instant};
8
9use async_trait::async_trait;
10use futures::{FutureExt, Stream, StreamExt};
11use parking_lot::Mutex;
12use serde_json::Value;
13use tokio::sync::{mpsc, oneshot};
14use tokio_stream::wrappers::ReceiverStream;
15use tracing::Instrument;
16
17use switchyard_protocol::{ModelId, Request, Response};
25
26use crate::{DriverError, LibsyError, Result, observability};
27
28pub type StepStream = Pin<Box<dyn Stream<Item = Result<Step>> + Send>>;
32
33pub struct CallModel {
43 pub algorithm: String,
46 pub request: Request,
48 pub models: Vec<ModelId>,
50 reply: oneshot::Sender<Result<Response>>,
52}
53
54impl CallModel {
55 pub fn respond(self, result: Result<Response>) -> Result<()> {
59 self.reply
60 .send(result)
61 .map_err(|_| DriverError::ResponseDropped.into())
62 }
63}
64
65pub struct RoutingOutcome {
67 pub selected_model_ids: Vec<ModelId>,
69 pub request: Request,
71 pub response: Option<Response>,
73 pub metadata: Option<crate::OutcomeMetadata>,
78}
79
80impl RoutingOutcome {
81 pub fn selected_model_id(&self) -> Result<&ModelId> {
84 self.selected_model_ids.first().ok_or(LibsyError::NoTargets)
85 }
86
87 pub fn route_to(
91 selected_model_id: ModelId,
92 fallback_models: Vec<ModelId>,
93 mut request: Request,
94 ) -> Self {
95 request.llm_request.model = Some(selected_model_id.to_string());
96 let mut selected_model_ids = Vec::with_capacity(1 + fallback_models.len());
97 selected_model_ids.push(selected_model_id);
98 selected_model_ids.extend(fallback_models);
99 Self {
100 selected_model_ids,
101 request,
102 response: None,
103 metadata: None,
104 }
105 }
106
107 pub fn answered(selected_model_id: ModelId, mut request: Request, response: Response) -> Self {
110 request.llm_request.model = Some(selected_model_id.to_string());
111 Self {
112 selected_model_ids: vec![selected_model_id],
113 request,
114 response: Some(response),
115 metadata: None,
116 }
117 }
118}
119
120#[derive(Clone)]
122pub struct Driver {
123 step_tx: mpsc::Sender<Result<Step>>,
124 algorithm: String,
126 evidence: Arc<Mutex<Option<Value>>>,
128}
129
130impl Driver {
131 pub(crate) fn new(algorithm: &str) -> (Self, mpsc::Receiver<Result<Step>>) {
134 let (step_tx, step_rx) = mpsc::channel(1);
139 (
140 Self {
141 step_tx,
142 algorithm: algorithm.to_string(),
143 evidence: Arc::new(Mutex::new(None)),
144 },
145 step_rx,
146 )
147 }
148
149 pub(crate) fn set_evidence(&self, evidence: Value) {
151 *self.evidence.lock() = Some(evidence);
152 }
153
154 pub(crate) fn set_evidence_if_empty(&self, evidence: Value) {
156 let mut current = self.evidence.lock();
157 if current.is_none() {
158 *current = Some(evidence);
159 }
160 }
161
162 #[tracing::instrument(
171 target = "libsy",
172 name = "libsy.llm_call",
173 skip_all,
174 fields(
175 algorithm = self.algorithm,
176 selected_model = %models.first().map(ModelId::as_str).unwrap_or("NoTargets"),
177 openinference.span.kind = "CHAIN",
178 outcome = tracing::field::Empty,
179 error = tracing::field::Empty,
180 input_tokens = tracing::field::Empty,
181 output_tokens = tracing::field::Empty,
182 total_tokens = tracing::field::Empty,
183 reasoning_tokens = tracing::field::Empty,
184 )
185 )]
186 pub async fn call_model(&self, mut request: Request, models: Vec<ModelId>) -> Result<Response> {
187 let Some(selected_model_id) = models.first().cloned() else {
188 return Err(LibsyError::NoTargets);
189 };
190 request.llm_request.model = Some(selected_model_id.to_string());
191 let started = Instant::now();
192 let (reply, response) = oneshot::channel::<Result<Response>>();
193 let call = CallModel {
194 algorithm: self.algorithm.clone(),
195 request,
196 models,
197 reply,
198 };
199 let result = async {
200 self.step_tx
201 .send(Ok(Step::CallModel(Box::new(call))))
202 .await
203 .map_err(|_| DriverError::StreamClosed)?;
204 response
205 .await
206 .map_err(|_| LibsyError::from(DriverError::ResponseDropped))?
207 }
208 .await;
209 let elapsed = started.elapsed();
210 observability::record_llm_call(
211 &self.algorithm,
212 selected_model_id.as_str(),
213 elapsed,
214 &result,
215 &tracing::Span::current(),
216 );
217 result
218 }
219
220 pub(crate) async fn finish(&self, result: Result<RoutingOutcome>) -> Result<()> {
224 let result = result.map(|mut outcome| {
225 let metadata = outcome.metadata.get_or_insert_with(|| {
226 crate::OutcomeMetadata::new(self.algorithm.clone(), self.evidence.lock().take())
227 });
228 tracing::Span::current().record("outcome_id", metadata.outcome_id());
229 outcome
230 });
231 let selected_model = result
232 .as_ref()
233 .ok()
234 .and_then(|outcome| outcome.selected_model_id().ok().cloned());
235 let step = result.map(|outcome| Step::Done(Box::new(outcome)));
236 self.step_tx
237 .send(step)
238 .await
239 .map_err(|_| DriverError::StreamClosed)?;
240 if let Some(selected_model) = selected_model {
241 observability::record_decision(&self.algorithm, &selected_model);
242 }
243 Ok(())
244 }
245}
246
247pub enum Step {
249 CallModel(Box<CallModel>),
252 Done(Box<RoutingOutcome>),
254}
255
256pub async fn drive<F, Fut>(
269 algorithm: Arc<dyn Algorithm>,
270 request: Request,
271 serve: F,
272) -> Result<RoutingOutcome>
273where
274 F: Fn(CallModel) -> Fut,
275 Fut: Future<Output = Result<()>>,
276{
277 let stream = algorithm.run_stream(request);
278 tokio::pin!(stream);
279
280 let mut in_flight = futures::stream::FuturesUnordered::new();
281 let mut final_outcome: Option<RoutingOutcome> = None;
282
283 loop {
284 tokio::select! {
285 Some(result) = in_flight.next() => match result {
286 Ok(()) => {}, Err(err) => return Err(err), },
289 step = stream.next() => {
290 match step {
291 None => break, Some(item) => match item? {
293 Step::CallModel(call) => in_flight.push(serve(*call)),
294 Step::Done(outcome) => {
295 final_outcome = Some(*outcome);
296 break;
297 }
298 }
299 }
300 },
301 }
302 }
303 final_outcome.ok_or(LibsyError::MissingFinalResponse)
304}
305
306fn panic_message(payload: &(dyn std::any::Any + Send)) -> String {
308 payload
309 .downcast_ref::<&'static str>()
310 .map(|message| (*message).to_string())
311 .or_else(|| payload.downcast_ref::<String>().cloned())
312 .unwrap_or_else(|| "unknown panic payload".to_string())
313}
314
315struct AbortOnDrop(tokio::task::AbortHandle);
317
318impl Drop for AbortOnDrop {
319 fn drop(&mut self) {
320 self.0.abort();
321 }
322}
323
324pub(crate) fn ensure_model_is_target(targets: &[ModelId], name: &ModelId) -> Result<()> {
329 targets
330 .iter()
331 .any(|target| target == name)
332 .then_some(())
333 .ok_or_else(|| LibsyError::TargetNotFound {
334 target: name.clone(),
335 })
336}
337
338#[derive(Clone, Hash, PartialEq, Eq)]
341pub(crate) enum RoutingIdentity {
342 Session(String),
344 Subagent { session: String, agent: String },
346}
347
348impl RoutingIdentity {
349 pub(crate) fn from_request(request: &Request) -> Option<Self> {
354 let metadata = request.metadata.as_ref()?;
355 let session = metadata.session_id.as_deref().filter(|id| !id.is_empty())?;
356 if metadata.is_subagent {
357 let agent = metadata.agent_id.as_deref().filter(|id| !id.is_empty())?;
358 Some(Self::Subagent {
359 session: session.to_string(),
360 agent: agent.to_string(),
361 })
362 } else {
363 Some(Self::Session(session.to_string()))
364 }
365 }
366}
367
368#[async_trait]
390pub trait Algorithm: Send + Sync + 'static {
391 fn name(&self) -> &str;
395
396 async fn route(self: Arc<Self>, driver: Driver, request: Request) -> Result<RoutingOutcome>;
400
401 fn run_stream(self: Arc<Self>, request: Request) -> StepStream {
410 let (driver, step_rx) = Driver::new(self.name());
411 let span = observability::run_span(self.name(), &request);
412 let handle = tokio::spawn(
413 async move {
414 let algorithm = self.name().to_string();
415 let route = AssertUnwindSafe(self.route(driver.clone(), request)).catch_unwind();
417 let result = observability::observe_run(&algorithm, async move {
418 route.await.unwrap_or_else(|payload| {
419 Err(LibsyError::AlgorithmError {
420 message: format!(
421 "algorithm task panicked: {}",
422 panic_message(payload.as_ref())
423 ),
424 })
425 })
426 })
427 .await;
428
429 let _ = driver.finish(result).await;
430 }
431 .instrument(span),
432 );
433 let abort_guard = AbortOnDrop(handle.abort_handle());
435 Box::pin(ReceiverStream::new(step_rx).map(move |step| {
436 let _keep_alive = &abort_guard;
438 step
439 }))
440 }
441}
442
443#[cfg(test)]
444mod tests {
445 use std::collections::HashMap;
446
447 use super::*;
448 use crate::core::testing::{Serve, ServeResult, echo, reply, test_drive};
449 use futures::StreamExt;
450 use switchyard_protocol::{
451 LlmResponse, LlmResponseChunk, completion_text, text_request, text_response,
452 };
453
454 #[derive(Debug, thiserror::Error)]
455 #[error("{0}")]
456 struct TestError(&'static str);
457
458 fn test_error(message: &'static str) -> LibsyError {
459 LibsyError::external("test", TestError(message))
460 }
461
462 struct TestAlgo {
465 target_set: Vec<ModelId>,
466 }
467
468 #[async_trait]
469 impl Algorithm for TestAlgo {
470 fn name(&self) -> &str {
471 "test"
472 }
473
474 async fn route(
475 self: Arc<Self>,
476 driver: Driver,
477 request: Request,
478 ) -> Result<RoutingOutcome> {
479 let target = self
480 .target_set
481 .first()
482 .ok_or(LibsyError::NoTargets)?
483 .clone();
484 let response = driver
485 .call_model(request.clone(), vec![target.clone()])
486 .await?;
487 driver.set_evidence(serde_json::json!({"source": "test"}));
488 driver.set_evidence_if_empty(serde_json::json!({"source": "ignored"}));
489 Ok(RoutingOutcome::answered(target, request, response))
490 }
491 }
492
493 fn orch(target_set: Vec<ModelId>) -> Arc<dyn Algorithm> {
495 Arc::new(TestAlgo { target_set })
496 }
497
498 fn request() -> Request {
499 Request {
500 llm_request: text_request(Some("auto".to_string()), "hi".to_string()),
501 raw_request: None,
502 metadata: None,
503 }
504 }
505
506 #[test]
507 fn routing_outcome_constructors_stamp_selection_and_preserve_payloads() {
508 let outcome = RoutingOutcome::route_to(
509 "selected".into(),
510 target_set(&["fallback-one", "fallback-two"]),
511 request(),
512 );
513
514 assert_eq!(
515 outcome.selected_model_ids,
516 target_set(&["selected", "fallback-one", "fallback-two"])
517 );
518 assert_eq!(outcome.request.model_id().as_deref(), Some("selected"));
519 assert!(outcome.response.is_none());
520 assert!(outcome.metadata.is_none());
521
522 let outcome = RoutingOutcome::route_to("only".into(), Vec::new(), request());
523 assert_eq!(outcome.selected_model_ids, target_set(&["only"]));
524
525 let outcome = RoutingOutcome::answered(
526 "answered".into(),
527 request(),
528 Response {
529 llm_response: LlmResponse::Agg(text_response(None, "existing")),
530 metadata: None,
531 },
532 );
533
534 assert_eq!(outcome.selected_model_ids, target_set(&["answered"]));
535 assert_eq!(outcome.request.model_id().as_deref(), Some("answered"));
536 assert_eq!(
537 outcome
538 .response
539 .as_ref()
540 .and_then(|response| response.llm_response.as_agg())
541 .map(completion_text),
542 Some("existing".to_string())
543 );
544 }
545
546 fn target_set(names: &[&str]) -> Vec<ModelId> {
547 names.iter().map(|name| ModelId::from(*name)).collect()
548 }
549
550 #[tokio::test]
551 async fn typed_driver_preserves_call_and_stream_boundaries() -> Result<()> {
552 tokio::time::timeout(std::time::Duration::from_secs(1), async {
553 let (driver, mut step_rx) = Driver::new("test");
556 let first_driver = driver.clone();
557 let mut first = tokio::spawn(async move {
558 first_driver
559 .call_model(request(), vec![ModelId::from("first")])
560 .await
561 });
562 let second = tokio::spawn(async move {
563 driver
564 .call_model(request(), vec![ModelId::from("second")])
565 .await
566 });
567
568 let mut calls = HashMap::new();
569 for _ in 0..2 {
570 let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
571 let Step::CallModel(call) = step else {
572 return Err(test_error("expected a CallModel step"));
573 };
574 let selected_model = call
575 .models
576 .first()
577 .ok_or_else(|| test_error("model call has no candidates"))?
578 .to_string();
579 calls.insert(selected_model, call);
580 }
581 assert!(
582 tokio::time::timeout(std::time::Duration::from_millis(20), &mut first)
583 .await
584 .is_err(),
585 "call completed before the host responded"
586 );
587 calls
588 .remove("second")
589 .ok_or_else(|| test_error("missing second call"))?
590 .respond(Ok(reply("second response")))?;
591 calls
592 .remove("first")
593 .ok_or_else(|| test_error("missing first call"))?
594 .respond(Ok(reply("first response")))?;
595
596 let first_response = first
597 .await
598 .map_err(|source| LibsyError::external("joining a test task", source))??;
599 let second_response = second
600 .await
601 .map_err(|source| LibsyError::external("joining a test task", source))??;
602 assert_eq!(
603 first_response.llm_response.as_agg().map(completion_text),
604 Some("first response".to_string())
605 );
606 assert_eq!(
607 second_response.llm_response.as_agg().map(completion_text),
608 Some("second response".to_string())
609 );
610
611 let (driver, mut step_rx) = Driver::new("test");
613 let producer = tokio::spawn(async move {
614 driver
615 .call_model(request(), vec![ModelId::from("dropped")])
616 .await
617 });
618 let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
619 let Step::CallModel(call) = step else {
620 return Err(test_error("expected a CallModel step"));
621 };
622 drop(call);
623 let result = producer
624 .await
625 .map_err(|source| LibsyError::external("joining a test task", source))?;
626 assert!(matches!(
627 result,
628 Err(LibsyError::Driver(DriverError::ResponseDropped))
629 ));
630
631 let (driver, step_rx) = Driver::new("test");
633 drop(step_rx);
634 let result = driver
635 .call_model(request(), vec![ModelId::from("closed")])
636 .await;
637 assert!(matches!(
638 result,
639 Err(LibsyError::Driver(DriverError::StreamClosed))
640 ));
641 Ok(())
642 })
643 .await
644 .map_err(|error| LibsyError::external("waiting for typed driver boundaries", error))?
645 }
646
647 #[test]
648 fn target_lookup_returns_the_missing_target() {
649 let error = ensure_model_is_target(&target_set(&[]), &ModelId::from("missing")).err();
650 assert!(matches!(
651 error,
652 Some(LibsyError::TargetNotFound { target }) if target == "missing"
653 ));
654 }
655
656 fn streaming_orch(chunks: Vec<LlmResponseChunk>) -> (Arc<dyn Algorithm>, impl Serve) {
659 let algo = orch(target_set(&["stream/model"]));
660 let serve = move |_target: ModelId, _request: Request| {
661 let chunks = chunks.clone();
662 async move {
663 let stream =
664 futures::stream::iter(chunks.into_iter().map(|chunk| Ok(chunk.into()))).boxed();
665 Ok(Response {
666 llm_response: LlmResponse::Stream(stream),
667 metadata: None,
668 })
669 }
670 };
671 (algo, serve)
672 }
673
674 #[tokio::test]
675 async fn run_returns_a_streamed_response_the_caller_aggregates() -> Result<()> {
676 let (orch, serve) = streaming_orch(vec![
679 LlmResponseChunk::MessageStart {
680 id: Some("m1".to_string()),
681 model: Some("stream/model".to_string()),
682 },
683 LlmResponseChunk::TextDelta {
684 index: 0,
685 text: "hel".to_string(),
686 },
687 LlmResponseChunk::TextDelta {
688 index: 0,
689 text: "lo".to_string(),
690 },
691 LlmResponseChunk::MessageStop {
692 reason: Some("stop".to_string()),
693 },
694 ]);
695 let (selected_model, response) = test_drive(orch, request(), serve).await?;
696 let agg = response
698 .llm_response
699 .into_agg()
700 .await
701 .map_err(|error| LibsyError::external("aggregating response stream", error))?;
702 assert_eq!(completion_text(&agg), "hello");
703 assert_eq!(agg.model.as_deref(), Some("stream/model"));
704 assert_eq!(selected_model, "stream/model");
705 Ok(())
706 }
707
708 #[tokio::test]
709 async fn aggregating_a_streamed_response_propagates_a_mid_stream_error() -> Result<()> {
710 let (orch, serve) = streaming_orch(vec![
713 LlmResponseChunk::TextDelta {
714 index: 0,
715 text: "partial".to_string(),
716 },
717 LlmResponseChunk::StreamError {
718 message: "upstream exploded".to_string(),
719 },
720 ]);
721 let (_, response) = test_drive(orch, request(), serve).await?;
722 match response.llm_response.into_agg().await {
723 Ok(_) => panic!("expected a mid-stream error, got an aggregate"),
724 Err(err) => {
725 assert!(err.to_string().contains("upstream exploded"));
726 Ok(())
727 }
728 }
729 }
730
731 #[tokio::test]
732 async fn run_offloads_via_promise_then_finishes() -> Result<()> {
733 let stream = orch(target_set(&["offload/model"])).run_stream(request());
736 tokio::pin!(stream);
737
738 let mut saw_call = false;
739 let mut final_completion = None;
740 while let Some(step) = stream.next().await {
741 match step? {
742 Step::CallModel(call) => {
743 saw_call = true;
744 assert_eq!(call.models, vec![ModelId::from("offload/model")]);
745 call.respond(Ok(Response {
747 llm_response: LlmResponse::Agg(text_response(
748 None,
749 "fulfilled".to_string(),
750 )),
751 metadata: None,
752 }))?;
753 }
754 Step::Done(outcome) => {
755 let metadata = outcome
756 .metadata
757 .as_ref()
758 .expect("run_stream should attach outcome metadata");
759 assert_eq!(metadata.algorithm, "test");
760 assert_eq!(
761 uuid::Uuid::parse_str(metadata.outcome_id())
762 .expect("outcome id should be a UUID")
763 .get_version_num(),
764 7
765 );
766 assert_eq!(
767 metadata.evidence,
768 Some(serde_json::json!({"source": "test"}))
769 );
770 let response = outcome
771 .response
772 .ok_or_else(|| test_error("expected an answered outcome"))?;
773 final_completion = Some(
774 response
775 .llm_response
776 .as_agg()
777 .map(completion_text)
778 .unwrap_or_default(),
779 );
780 }
781 }
782 }
783
784 assert!(saw_call, "expected a CallModel step before Done");
785 assert_eq!(
786 final_completion.ok_or_else(|| test_error("no Done step"))?,
787 "fulfilled"
788 );
789 Ok(())
790 }
791
792 #[tokio::test(flavor = "multi_thread", worker_threads = 12)]
793 async fn requests_are_processed_in_parallel() -> Result<()> {
794 use std::time::Duration;
795 use tokio::sync::Barrier;
796
797 const N: usize = 12;
798
799 let barrier = Arc::new(Barrier::new(N));
804 let algo = orch(target_set(&["m"]));
806
807 let mut handles = Vec::new();
808 for _ in 0..N {
809 let algo = algo.clone();
810 let barrier = barrier.clone();
811 let serve = move |target: ModelId, _request: Request| {
812 let barrier = barrier.clone();
813 async move {
814 barrier.wait().await;
815 Ok(reply(target))
816 }
817 };
818 handles.push(tokio::spawn(async move {
819 test_drive(algo, request(), serve)
820 .await
821 .map(|(_, response)| {
822 response
823 .llm_response
824 .as_agg()
825 .map(completion_text)
826 .unwrap_or_default()
827 })
828 }));
829 }
830
831 for handle in handles {
832 let completion = tokio::time::timeout(Duration::from_secs(5), handle)
834 .await
835 .map_err(|error| LibsyError::external("waiting for test task", error))?
836 .map_err(|source| LibsyError::external("joining a test task", source))??;
837 assert_eq!(completion, "m");
838 }
839 Ok(())
840 }
841
842 #[tokio::test]
843 async fn offload_error_propagates_back_to_the_algorithm() -> Result<()> {
844 let stream = orch(target_set(&["offload/model"])).run_stream(request());
848 tokio::pin!(stream);
849
850 let mut saw_error = false;
851 while let Some(step) = stream.next().await {
852 match step {
853 Ok(Step::CallModel(call)) => {
854 call.respond(Err(test_error("upstream model call failed")))?;
855 }
856 Ok(Step::Done(..)) => {
857 return Err(test_error(
858 "expected the offload error to propagate, got a response",
859 ));
860 }
861 Err(err) => {
862 assert!(err.to_string().contains("upstream model call failed"));
864 saw_error = true;
865 }
866 }
867 }
868
869 assert!(saw_error, "expected an error step");
870 Ok(())
871 }
872
873 #[tokio::test]
874 async fn dropping_the_stream_cancels_the_algorithm_task() -> Result<()> {
875 use std::sync::atomic::{AtomicBool, Ordering};
876 use std::time::Duration;
877 use tokio::sync::mpsc;
878
879 struct DropGuard(Arc<AtomicBool>);
882 impl Drop for DropGuard {
883 fn drop(&mut self) {
884 self.0.store(true, Ordering::SeqCst);
885 }
886 }
887
888 struct StuckAlgo {
889 started: mpsc::UnboundedSender<()>,
890 dropped: Arc<AtomicBool>,
891 }
892
893 #[async_trait]
894 impl Algorithm for StuckAlgo {
895 fn name(&self) -> &str {
896 "stuck"
897 }
898
899 async fn route(
900 self: Arc<Self>,
901 _driver: Driver,
902 _request: Request,
903 ) -> Result<RoutingOutcome> {
904 let _guard = DropGuard(self.dropped.clone());
905 let _ = self.started.send(());
906 std::future::pending::<()>().await;
908 unreachable!()
909 }
910 }
911
912 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
913 let dropped = Arc::new(AtomicBool::new(false));
914 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
915 started: started_tx,
916 dropped: dropped.clone(),
917 });
918
919 let stream = algo.run_stream(request());
920 started_rx
921 .recv()
922 .await
923 .ok_or_else(|| test_error("task never started"))?;
924 drop(stream);
925 tokio::time::sleep(Duration::from_millis(100)).await;
926
927 assert!(
928 dropped.load(Ordering::SeqCst),
929 "algorithm task was NOT cancelled after dropping the stream"
930 );
931 Ok(())
932 }
933
934 #[tokio::test]
935 async fn route_panic_surfaces_as_a_stream_error() -> Result<()> {
936 struct Panicky;
939
940 #[async_trait]
941 impl Algorithm for Panicky {
942 fn name(&self) -> &str {
943 "panicky"
944 }
945
946 async fn route(
947 self: Arc<Self>,
948 _driver: Driver,
949 _request: Request,
950 ) -> Result<RoutingOutcome> {
951 panic!("boom");
952 }
953 }
954
955 let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
956 let stream = algo.run_stream(request());
957 tokio::pin!(stream);
958
959 let mut saw_error = false;
960 while let Some(step) = stream.next().await {
961 match step {
962 Err(err) => {
963 assert!(err.to_string().contains("algorithm task panicked: boom"));
965 saw_error = true;
966 }
967 Ok(_) => return Err(test_error("expected the panic to surface as an error step")),
968 }
969 }
970
971 assert!(saw_error, "expected an error step from the panicked task");
972 Ok(())
973 }
974
975 #[tokio::test]
979 async fn a_panic_with_a_leaked_driver_clone_still_terminates_the_run() -> Result<()> {
980 struct LeakyPanic;
981
982 #[async_trait]
983 impl Algorithm for LeakyPanic {
984 fn name(&self) -> &str {
985 "leaky_panic"
986 }
987
988 async fn route(
989 self: Arc<Self>,
990 driver: Driver,
991 _request: Request,
992 ) -> Result<RoutingOutcome> {
993 tokio::spawn(async move {
994 let _keep_alive = driver;
996 std::future::pending::<()>().await;
997 });
998 tokio::task::yield_now().await;
999 panic!("boom");
1000 }
1001 }
1002
1003 let algo: Arc<dyn Algorithm> = Arc::new(LeakyPanic);
1004 let result = tokio::time::timeout(
1006 std::time::Duration::from_secs(1),
1007 test_drive(algo, request(), echo()),
1008 )
1009 .await
1010 .map_err(|error| LibsyError::external("waiting for the panicked run to end", error))?;
1011
1012 match result {
1013 Ok(_) => Err(test_error(
1014 "expected the panic to end the run with an error",
1015 )),
1016 Err(err) => {
1017 assert!(err.to_string().contains("algorithm task panicked: boom"));
1018 Ok(())
1019 }
1020 }
1021 }
1022
1023 #[tokio::test]
1024 async fn cancelling_run_cancels_the_algorithm_task() -> Result<()> {
1025 use std::sync::atomic::{AtomicBool, Ordering};
1026 use std::time::Duration;
1027 use tokio::sync::mpsc;
1028
1029 struct DropGuard(Arc<AtomicBool>);
1032 impl Drop for DropGuard {
1033 fn drop(&mut self) {
1034 self.0.store(true, Ordering::SeqCst);
1035 }
1036 }
1037
1038 struct StuckAlgo {
1039 started: mpsc::UnboundedSender<()>,
1040 dropped: Arc<AtomicBool>,
1041 }
1042
1043 #[async_trait]
1044 impl Algorithm for StuckAlgo {
1045 fn name(&self) -> &str {
1046 "stuck"
1047 }
1048
1049 async fn route(
1050 self: Arc<Self>,
1051 _driver: Driver,
1052 _request: Request,
1053 ) -> Result<RoutingOutcome> {
1054 let _guard = DropGuard(self.dropped.clone());
1055 let _ = self.started.send(());
1056 std::future::pending::<()>().await;
1059 unreachable!()
1060 }
1061 }
1062
1063 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1064 let dropped = Arc::new(AtomicBool::new(false));
1065 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1066 started: started_tx,
1067 dropped: dropped.clone(),
1068 });
1069
1070 let run_task = tokio::spawn(async move { test_drive(algo, request(), echo()).await });
1073 started_rx
1074 .recv()
1075 .await
1076 .ok_or_else(|| test_error("task never started"))?;
1077 run_task.abort();
1078 tokio::time::sleep(Duration::from_millis(100)).await;
1079
1080 assert!(
1081 dropped.load(Ordering::SeqCst),
1082 "algorithm task was NOT cancelled after cancelling run"
1083 );
1084 Ok(())
1085 }
1086
1087 struct Hedge {
1092 winner: String,
1093 loser: String,
1094 }
1095
1096 #[async_trait]
1097 impl Algorithm for Hedge {
1098 fn name(&self) -> &str {
1099 "hedge"
1100 }
1101
1102 async fn route(
1103 self: Arc<Self>,
1104 driver: Driver,
1105 request: Request,
1106 ) -> Result<RoutingOutcome> {
1107 let outcome_request = request.clone();
1108 let win = driver.call_model(request.clone(), vec![self.winner.clone().into()]);
1109 let lose = driver.call_model(request, vec![self.loser.clone().into()]);
1110 tokio::select! {
1112 res = win => Ok(RoutingOutcome::answered(
1113 self.winner.clone().into(),
1114 outcome_request,
1115 res?,
1116 )),
1117 res = lose => Ok(RoutingOutcome::answered(
1118 self.loser.clone().into(),
1119 outcome_request,
1120 res?,
1121 )),
1122 }
1123 }
1124 }
1125
1126 fn hedge(loser_delay: Option<std::time::Duration>) -> (Arc<dyn Algorithm>, impl Serve) {
1130 let started = Arc::new(tokio::sync::Notify::new());
1131 let algo = Arc::new(Hedge {
1132 winner: "winner".to_string(),
1133 loser: "loser".to_string(),
1134 });
1135 let serve = move |target: ModelId, _request: Request| {
1136 let started = started.clone();
1137 async move {
1138 if target == "loser" {
1139 started.notify_one();
1140 match loser_delay {
1141 Some(delay) => tokio::time::sleep(delay).await,
1142 None => std::future::pending::<()>().await,
1143 }
1144 } else {
1145 started.notified().await;
1146 }
1147 Ok(reply(target))
1148 }
1149 };
1150 (algo, serve)
1151 }
1152
1153 #[tokio::test]
1154 async fn run_returns_the_winner_without_a_late_loser_overwriting_it() -> Result<()> {
1155 let (algo, serve) = hedge(Some(std::time::Duration::from_millis(50)));
1158 let (_, response) = test_drive(algo, request(), serve).await?;
1159 assert_eq!(
1160 response
1161 .llm_response
1162 .as_agg()
1163 .map(completion_text)
1164 .unwrap_or_default(),
1165 "winner"
1166 );
1167 Ok(())
1168 }
1169
1170 #[tokio::test]
1171 async fn run_returns_the_winner_without_hanging_on_a_pending_loser() -> Result<()> {
1172 let (algo, serve) = hedge(None);
1175 let run = test_drive(algo, request(), serve);
1176 let (_, response) = tokio::time::timeout(std::time::Duration::from_secs(1), run)
1177 .await
1178 .map_err(|error| LibsyError::external("waiting for pending loser", error))??;
1179 assert_eq!(
1180 response
1181 .llm_response
1182 .as_agg()
1183 .map(completion_text)
1184 .unwrap_or_default(),
1185 "winner"
1186 );
1187 Ok(())
1188 }
1189
1190 #[tokio::test]
1191 async fn run_surfaces_a_terminal_error_with_many_calls_in_flight() -> Result<()> {
1192 use std::sync::atomic::{AtomicUsize, Ordering};
1193
1194 const N: usize = 10;
1197
1198 struct FanOutThenError {
1201 all_started: Arc<tokio::sync::Notify>,
1202 n: usize,
1203 }
1204
1205 #[async_trait]
1206 impl Algorithm for FanOutThenError {
1207 fn name(&self) -> &str {
1208 "fan_out_then_error"
1209 }
1210
1211 async fn route(
1212 self: Arc<Self>,
1213 driver: Driver,
1214 request: Request,
1215 ) -> Result<RoutingOutcome> {
1216 let offloads = futures::future::join_all(
1217 (0..self.n)
1218 .map(|i| driver.call_model(request.clone(), vec![format!("m{i}").into()])),
1219 );
1220 tokio::select! {
1221 _ = offloads => Err(test_error("offloads unexpectedly completed")),
1222 _ = self.all_started.notified() => {
1223 Err(test_error("terminal error while calls pending"))
1224 }
1225 }
1226 }
1227 }
1228
1229 let all_started = Arc::new(tokio::sync::Notify::new());
1230 let algo: Arc<dyn Algorithm> = Arc::new(FanOutThenError {
1231 all_started: all_started.clone(),
1232 n: N,
1233 });
1234
1235 let started = Arc::new(AtomicUsize::new(0));
1237 let serve = move |_target: ModelId, _request: Request| {
1238 let started = started.clone();
1239 let all_started = all_started.clone();
1240 async move {
1241 if started.fetch_add(1, Ordering::SeqCst) + 1 == N {
1242 all_started.notify_one();
1243 }
1244 std::future::pending::<ServeResult>().await
1245 }
1246 };
1247
1248 let run = test_drive(algo, request(), serve);
1251 let result = tokio::time::timeout(std::time::Duration::from_millis(500), run)
1252 .await
1253 .map_err(|error| {
1254 LibsyError::external("waiting for terminal error with full call cap", error)
1255 })?;
1256 match result {
1257 Ok(_) => Err(test_error("expected the terminal error, got a response")),
1258 Err(err) => {
1259 assert!(
1260 err.to_string()
1261 .contains("terminal error while calls pending")
1262 );
1263 Ok(())
1264 }
1265 }
1266 }
1267}