1use 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
20use switchyard_protocol::{Category, ModelId, Request, Response};
28
29use crate::{DriverError, LibsyError, Result, observability};
30
31pub type StepStream = Pin<Box<dyn Stream<Item = Result<Step>> + Send>>;
35
36#[derive(Clone, Debug, Default)]
47pub struct RuntimeModels {
48 by_category: HashMap<Category, Vec<ModelId>>,
50 subagent: Option<HashMap<Category, Vec<ModelId>>>,
51}
52
53impl RuntimeModels {
54 pub fn new(by_category: HashMap<Category, Vec<ModelId>>) -> Self {
56 Self {
57 by_category,
58 subagent: None,
59 }
60 }
61
62 pub fn with_subagent(mut self, models: HashMap<Category, Vec<ModelId>>) -> Self {
64 self.subagent = Some(models);
65 self
66 }
67
68 pub fn models_for(&self, category: &Category) -> &[ModelId] {
70 self.by_category.get(category).map_or(&[], Vec::as_slice)
71 }
72
73 pub fn subagent_models_for(&self, category: &Category) -> &[ModelId] {
75 self.subagent
76 .as_ref()
77 .and_then(|models| models.get(category))
78 .map_or(&[], Vec::as_slice)
79 }
80}
81
82impl From<HashMap<Category, Vec<ModelId>>> for RuntimeModels {
83 fn from(by_category: HashMap<Category, Vec<ModelId>>) -> Self {
84 Self::new(by_category)
85 }
86}
87
88#[derive(Clone, Copy)]
90enum Scope {
91 Parent,
92 Subagent,
93}
94
95pub struct CallModel {
105 pub algorithm: String,
108 pub request: Request,
110 pub models: Vec<ModelId>,
112 pub recover_errors: bool,
114 reply: Option<oneshot::Sender<Result<Response>>>,
116 started: Instant,
117}
118
119impl CallModel {
120 pub fn respond(mut self, result: Result<Response>) -> Result<()> {
124 self.record(result.is_ok());
125 self.reply
126 .take()
127 .ok_or(DriverError::ResponseDropped)?
128 .send(result)
129 .map_err(|_| DriverError::ResponseDropped.into())
130 }
131
132 pub fn fail(mut self, error: LibsyError) -> Result<()> {
135 self.reply = None;
136 self.record(false);
137 Err(error)
138 }
139
140 fn record(&self, is_ok: bool) {
141 observability::record_llm_call(
142 &self.algorithm,
143 self.models
144 .first()
145 .map(ModelId::as_str)
146 .unwrap_or("NoTargets"),
147 self.started.elapsed(),
148 is_ok,
149 );
150 }
151}
152
153impl Drop for CallModel {
154 fn drop(&mut self) {
155 if self.reply.is_some() {
156 self.record(false);
157 }
158 }
159}
160
161pub struct RoutingOutcome {
163 pub selected_model_ids: Vec<ModelId>,
165 pub request: Request,
167 pub response: Option<Response>,
169 pub metadata: Option<crate::OutcomeMetadata>,
174}
175
176impl RoutingOutcome {
177 pub fn selected_model_id(&self) -> Result<&ModelId> {
180 self.selected_model_ids.first().ok_or(LibsyError::NoTargets)
181 }
182
183 pub fn route_to(
187 selected_model_id: ModelId,
188 fallback_models: Vec<ModelId>,
189 mut request: Request,
190 ) -> Self {
191 request.llm_request.model = Some(selected_model_id.to_string());
192 let mut selected_model_ids = Vec::with_capacity(1 + fallback_models.len());
193 selected_model_ids.push(selected_model_id);
194 selected_model_ids.extend(fallback_models);
195 Self {
196 selected_model_ids,
197 request,
198 response: None,
199 metadata: None,
200 }
201 }
202
203 pub fn answered(selected_model_id: ModelId, mut request: Request, response: Response) -> Self {
206 request.llm_request.model = Some(selected_model_id.to_string());
207 Self {
208 selected_model_ids: vec![selected_model_id],
209 request,
210 response: Some(response),
211 metadata: None,
212 }
213 }
214}
215
216#[derive(Clone)]
218pub struct Driver {
219 step_tx: mpsc::Sender<Result<Step>>,
220
221 algorithm: String,
223
224 evidence: Arc<Mutex<Option<Value>>>,
226
227 models: Arc<RuntimeModels>,
229
230 scope: Scope,
232}
233
234impl Driver {
235 pub(crate) fn new(
238 algorithm: &str,
239 models: Arc<RuntimeModels>,
240 ) -> (Self, mpsc::Receiver<Result<Step>>) {
241 let (step_tx, step_rx) = mpsc::channel(1);
246 (
247 Self {
248 step_tx,
249 algorithm: algorithm.to_string(),
250 evidence: Arc::new(Mutex::new(None)),
251 models,
252 scope: Scope::Parent,
253 },
254 step_rx,
255 )
256 }
257
258 pub(crate) fn set_evidence(&self, evidence: Value) {
260 *self.evidence.lock() = Some(evidence);
261 }
262
263 pub(crate) fn set_evidence_if_empty(&self, evidence: Value) {
265 let mut current = self.evidence.lock();
266 if current.is_none() {
267 *current = Some(evidence);
268 }
269 }
270
271 pub async fn call_model(&self, request: Request, models: Vec<ModelId>) -> Result<Response> {
281 self.call_model_with_error_recovery(request, models, false)
282 .await
283 }
284
285 #[tracing::instrument(
287 target = "libsy",
288 name = "libsy.llm_call",
289 skip_all,
290 fields(
291 algorithm = self.algorithm,
292 selected_model = %models.first().map(ModelId::as_str).unwrap_or("NoTargets"),
293 openinference.span.kind = "CHAIN",
294 outcome = tracing::field::Empty,
295 input_tokens = tracing::field::Empty,
296 output_tokens = tracing::field::Empty,
297 total_tokens = tracing::field::Empty,
298 reasoning_tokens = tracing::field::Empty,
299 )
300 )]
301 pub(crate) async fn call_model_with_error_recovery(
302 &self,
303 mut request: Request,
304 models: Vec<ModelId>,
305 recover_errors: bool,
306 ) -> Result<Response> {
307 let Some(selected_model_id) = models.first() else {
308 return Err(LibsyError::NoTargets);
309 };
310 request.llm_request.model = Some(selected_model_id.to_string());
311 let started = Instant::now();
312 let (reply, response) = oneshot::channel::<Result<Response>>();
313 let call = CallModel {
314 algorithm: self.algorithm.clone(),
315 request,
316 models,
317 recover_errors,
318 reply: Some(reply),
319 started,
320 };
321 let result = async {
322 self.step_tx
323 .send(Ok(Step::CallModel(Box::new(call))))
324 .await
325 .map_err(|_| DriverError::StreamClosed)?;
326 response
327 .await
328 .map_err(|_| LibsyError::from(DriverError::ResponseDropped))?
329 }
330 .await;
331 observability::record_llm_call_span(&result, &tracing::Span::current());
332 result
333 }
334
335 pub fn models_for(&self, category: &Category) -> &[ModelId] {
337 match self.scope {
338 Scope::Parent => self.models.models_for(category),
339 Scope::Subagent => self.models.subagent_models_for(category),
340 }
341 }
342
343 pub fn first_model_for(&self, category: &Category) -> Result<&ModelId> {
345 self.models_for(category)
346 .first()
347 .ok_or_else(|| LibsyError::AlgorithmError {
348 message: format!("no models available for category {}", category.as_str()),
349 })
350 }
351
352 pub fn for_subagent(&self) -> Result<Self> {
355 if self.models.subagent.is_none() {
356 return Err(LibsyError::AlgorithmError {
357 message: "delegated work has no sub-agent models".to_string(),
358 });
359 }
360 Ok(Self {
361 scope: Scope::Subagent,
362 ..self.clone()
363 })
364 }
365
366 pub(crate) async fn finish(&self, result: Result<RoutingOutcome>) -> Result<()> {
370 let result = result.map(|mut outcome| {
371 let metadata = outcome.metadata.get_or_insert_with(|| {
372 crate::OutcomeMetadata::new(self.algorithm.clone(), self.evidence.lock().take())
373 });
374 observability::record_outcome(metadata, &outcome.selected_model_ids);
375 outcome
376 });
377 let selected_model = result
378 .as_ref()
379 .ok()
380 .and_then(|outcome| outcome.selected_model_id().ok().cloned());
381 let step = result.map(|outcome| Step::Done(Box::new(outcome)));
382 self.step_tx
383 .send(step)
384 .await
385 .map_err(|_| DriverError::StreamClosed)?;
386 if let Some(selected_model) = selected_model {
387 observability::record_decision(&self.algorithm, &selected_model);
388 }
389 Ok(())
390 }
391}
392
393pub enum Step {
395 CallModel(Box<CallModel>),
398 Done(Box<RoutingOutcome>),
400}
401
402pub async fn drive<F, Fut>(
416 algorithm: Arc<dyn Algorithm>,
417 request: Request,
418 models: Arc<RuntimeModels>,
419 serve: F,
420) -> Result<RoutingOutcome>
421where
422 F: Fn(CallModel) -> Fut,
423 Fut: Future<Output = Result<()>>,
424{
425 let stream = algorithm.run_stream(request, models);
426 tokio::pin!(stream);
427
428 let mut in_flight = futures::stream::FuturesUnordered::new();
429 let mut final_outcome: Option<RoutingOutcome> = None;
430
431 loop {
432 tokio::select! {
433 Some(result) = in_flight.next() => match result {
434 Ok(()) => {}, Err(err) => return Err(err), },
437 step = stream.next() => {
438 match step {
439 None => break, Some(item) => match item? {
441 Step::CallModel(call) => in_flight.push(serve(*call)),
442 Step::Done(outcome) => {
443 final_outcome = Some(*outcome);
444 break;
445 }
446 }
447 }
448 },
449 }
450 }
451 final_outcome.ok_or(LibsyError::MissingFinalResponse)
452}
453
454fn panic_message(payload: &(dyn std::any::Any + Send)) -> String {
456 payload
457 .downcast_ref::<&'static str>()
458 .map(|message| (*message).to_string())
459 .or_else(|| payload.downcast_ref::<String>().cloned())
460 .unwrap_or_else(|| "unknown panic payload".to_string())
461}
462
463struct AbortOnDrop(tokio::task::AbortHandle);
465
466impl Drop for AbortOnDrop {
467 fn drop(&mut self) {
468 self.0.abort();
469 }
470}
471
472pub(crate) fn ensure_model_is_target(targets: &[ModelId], name: &ModelId) -> Result<()> {
477 targets
478 .iter()
479 .any(|target| target == name)
480 .then_some(())
481 .ok_or_else(|| LibsyError::TargetNotFound {
482 target: name.clone(),
483 })
484}
485
486#[derive(Clone, Hash, PartialEq, Eq)]
489pub(crate) enum RoutingIdentity {
490 Session(String),
492 Subagent { session: String, agent: String },
494}
495
496impl RoutingIdentity {
497 pub(crate) fn from_request(request: &Request) -> Option<Self> {
502 let metadata = request.metadata.as_ref()?;
503 let session = metadata.session_id.as_deref().filter(|id| !id.is_empty())?;
504 if metadata.is_subagent {
505 let agent = metadata.agent_id.as_deref().filter(|id| !id.is_empty())?;
506 Some(Self::Subagent {
507 session: session.to_string(),
508 agent: agent.to_string(),
509 })
510 } else {
511 Some(Self::Session(session.to_string()))
512 }
513 }
514}
515
516#[async_trait]
549pub trait Algorithm: Send + Sync + 'static {
550 fn name(&self) -> &str;
554
555 async fn route(self: Arc<Self>, driver: Driver, request: Request) -> Result<RoutingOutcome>;
559
560 fn run_stream(self: Arc<Self>, request: Request, models: Arc<RuntimeModels>) -> StepStream {
569 let (driver, step_rx) = Driver::new(self.name(), models);
570 let span = observability::run_span(self.name(), &request);
571 let handle = tokio::spawn(
572 async move {
573 let algorithm = self.name().to_string();
574 let route = AssertUnwindSafe(self.route(driver.clone(), request)).catch_unwind();
576 let result = observability::observe_run(&algorithm, async move {
577 route.await.unwrap_or_else(|payload| {
578 Err(LibsyError::AlgorithmError {
579 message: format!(
580 "algorithm task panicked: {}",
581 panic_message(payload.as_ref())
582 ),
583 })
584 })
585 })
586 .await;
587
588 let _ = driver.finish(result).await;
589 }
590 .instrument(span),
591 );
592 let abort_guard = AbortOnDrop(handle.abort_handle());
594 Box::pin(ReceiverStream::new(step_rx).map(move |step| {
595 let _keep_alive = &abort_guard;
597 step
598 }))
599 }
600}
601
602#[cfg(test)]
603mod tests {
604 use std::collections::HashMap;
605
606 use super::*;
607 use crate::core::testing::{Serve, ServeResult, echo, reply, test_drive};
608 use futures::StreamExt;
609 use switchyard_protocol::{
610 LlmResponse, LlmResponseChunk, completion_text, text_request, text_response,
611 };
612
613 #[derive(Debug, thiserror::Error)]
614 #[error("{0}")]
615 struct TestError(&'static str);
616
617 fn test_error(message: &'static str) -> LibsyError {
618 LibsyError::external("test", TestError(message))
619 }
620
621 struct TestAlgo {
624 target_set: Vec<ModelId>,
625 }
626
627 #[async_trait]
628 impl Algorithm for TestAlgo {
629 fn name(&self) -> &str {
630 "test"
631 }
632
633 async fn route(
634 self: Arc<Self>,
635 driver: Driver,
636 request: Request,
637 ) -> Result<RoutingOutcome> {
638 let target = self
639 .target_set
640 .first()
641 .ok_or(LibsyError::NoTargets)?
642 .clone();
643 let response = driver
644 .call_model(request.clone(), vec![target.clone()])
645 .await?;
646 driver.set_evidence(serde_json::json!({"source": "test"}));
647 driver.set_evidence_if_empty(serde_json::json!({"source": "ignored"}));
648 Ok(RoutingOutcome::answered(target, request, response))
649 }
650 }
651
652 fn orch(target_set: Vec<ModelId>) -> Arc<dyn Algorithm> {
654 Arc::new(TestAlgo { target_set })
655 }
656
657 fn request() -> Request {
658 Request {
659 llm_request: text_request(Some("auto".to_string()), "hi".to_string()),
660 raw_request: None,
661 metadata: None,
662 }
663 }
664
665 #[test]
666 fn routing_outcome_constructors_stamp_selection_and_preserve_payloads() {
667 let outcome = RoutingOutcome::route_to(
668 "selected".into(),
669 target_set(&["fallback-one", "fallback-two"]),
670 request(),
671 );
672
673 assert_eq!(
674 outcome.selected_model_ids,
675 target_set(&["selected", "fallback-one", "fallback-two"])
676 );
677 assert_eq!(outcome.request.model_id().as_deref(), Some("selected"));
678 assert!(outcome.response.is_none());
679 assert!(outcome.metadata.is_none());
680
681 let outcome = RoutingOutcome::route_to("only".into(), Vec::new(), request());
682 assert_eq!(outcome.selected_model_ids, target_set(&["only"]));
683
684 let outcome = RoutingOutcome::answered(
685 "answered".into(),
686 request(),
687 Response {
688 llm_response: LlmResponse::Agg(text_response(None, "existing")),
689 metadata: None,
690 upstream_headers: http::HeaderMap::new(),
691 },
692 );
693
694 assert_eq!(outcome.selected_model_ids, target_set(&["answered"]));
695 assert_eq!(outcome.request.model_id().as_deref(), Some("answered"));
696 assert_eq!(
697 outcome
698 .response
699 .as_ref()
700 .and_then(|response| response.llm_response.as_agg())
701 .map(completion_text),
702 Some("existing".to_string())
703 );
704 }
705
706 fn target_set(names: &[&str]) -> Vec<ModelId> {
707 names.iter().map(|name| ModelId::from(*name)).collect()
708 }
709
710 #[tokio::test]
711 async fn typed_driver_preserves_call_and_stream_boundaries() -> Result<()> {
712 tokio::time::timeout(std::time::Duration::from_secs(1), async {
713 let (driver, mut step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
716 let first_driver = driver.clone();
717 let mut first = tokio::spawn(async move {
718 first_driver
719 .call_model(request(), vec![ModelId::from("first")])
720 .await
721 });
722 let second = tokio::spawn(async move {
723 driver
724 .call_model(request(), vec![ModelId::from("second")])
725 .await
726 });
727
728 let mut calls = HashMap::new();
729 for _ in 0..2 {
730 let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
731 let Step::CallModel(call) = step else {
732 return Err(test_error("expected a CallModel step"));
733 };
734 let selected_model = call
735 .models
736 .first()
737 .ok_or_else(|| test_error("model call has no candidates"))?
738 .to_string();
739 calls.insert(selected_model, call);
740 }
741 assert!(
742 tokio::time::timeout(std::time::Duration::from_millis(20), &mut first)
743 .await
744 .is_err(),
745 "call completed before the host responded"
746 );
747 calls
748 .remove("second")
749 .ok_or_else(|| test_error("missing second call"))?
750 .respond(Ok(reply("second response")))?;
751 calls
752 .remove("first")
753 .ok_or_else(|| test_error("missing first call"))?
754 .respond(Ok(reply("first response")))?;
755
756 let first_response = first
757 .await
758 .map_err(|source| LibsyError::external("joining a test task", source))??;
759 let second_response = second
760 .await
761 .map_err(|source| LibsyError::external("joining a test task", source))??;
762 assert_eq!(
763 first_response.llm_response.as_agg().map(completion_text),
764 Some("first response".to_string())
765 );
766 assert_eq!(
767 second_response.llm_response.as_agg().map(completion_text),
768 Some("second response".to_string())
769 );
770
771 let (driver, mut step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
773 let producer = tokio::spawn(async move {
774 driver
775 .call_model(request(), vec![ModelId::from("dropped")])
776 .await
777 });
778 let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
779 let Step::CallModel(call) = step else {
780 return Err(test_error("expected a CallModel step"));
781 };
782 drop(call);
783 let result = producer
784 .await
785 .map_err(|source| LibsyError::external("joining a test task", source))?;
786 assert!(matches!(
787 result,
788 Err(LibsyError::Driver(DriverError::ResponseDropped))
789 ));
790
791 let (driver, step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
793 drop(step_rx);
794 let result = driver
795 .call_model(request(), vec![ModelId::from("closed")])
796 .await;
797 assert!(matches!(
798 result,
799 Err(LibsyError::Driver(DriverError::StreamClosed))
800 ));
801 Ok(())
802 })
803 .await
804 .map_err(|error| LibsyError::external("waiting for typed driver boundaries", error))?
805 }
806
807 #[test]
808 fn target_lookup_returns_the_missing_target() {
809 let error = ensure_model_is_target(&target_set(&[]), &ModelId::from("missing")).err();
810 assert!(matches!(
811 error,
812 Some(LibsyError::TargetNotFound { target }) if target == "missing"
813 ));
814 }
815
816 fn streaming_orch(chunks: Vec<LlmResponseChunk>) -> (Arc<dyn Algorithm>, impl Serve) {
819 let algo = orch(target_set(&["stream/model"]));
820 let serve = move |_target: ModelId, _request: Request| {
821 let chunks = chunks.clone();
822 async move {
823 let stream =
824 futures::stream::iter(chunks.into_iter().map(|chunk| Ok(chunk.into()))).boxed();
825 Ok(Response {
826 llm_response: LlmResponse::Stream(stream),
827 metadata: None,
828 upstream_headers: http::HeaderMap::new(),
829 })
830 }
831 };
832 (algo, serve)
833 }
834
835 #[tokio::test]
836 async fn run_returns_a_streamed_response_the_caller_aggregates() -> Result<()> {
837 let (orch, serve) = streaming_orch(vec![
840 LlmResponseChunk::MessageStart {
841 id: Some("m1".to_string()),
842 model: Some("stream/model".to_string()),
843 },
844 LlmResponseChunk::TextDelta {
845 index: 0,
846 text: "hel".to_string(),
847 },
848 LlmResponseChunk::TextDelta {
849 index: 0,
850 text: "lo".to_string(),
851 },
852 LlmResponseChunk::MessageStop {
853 reason: Some("stop".to_string()),
854 },
855 ]);
856 let (selected_model, response) = test_drive(orch, request(), serve).await?;
857 let agg = response
859 .llm_response
860 .into_agg()
861 .await
862 .map_err(|error| LibsyError::external("aggregating response stream", error))?;
863 assert_eq!(completion_text(&agg), "hello");
864 assert_eq!(agg.model.as_deref(), Some("stream/model"));
865 assert_eq!(selected_model, "stream/model");
866 Ok(())
867 }
868
869 #[tokio::test]
870 async fn aggregating_a_streamed_response_propagates_a_mid_stream_error() -> Result<()> {
871 let (orch, serve) = streaming_orch(vec![
874 LlmResponseChunk::TextDelta {
875 index: 0,
876 text: "partial".to_string(),
877 },
878 LlmResponseChunk::StreamError {
879 message: "upstream exploded".to_string(),
880 },
881 ]);
882 let (_, response) = test_drive(orch, request(), serve).await?;
883 match response.llm_response.into_agg().await {
884 Ok(_) => panic!("expected a mid-stream error, got an aggregate"),
885 Err(err) => {
886 assert!(err.to_string().contains("upstream exploded"));
887 Ok(())
888 }
889 }
890 }
891
892 #[tokio::test]
893 async fn run_offloads_via_promise_then_finishes() -> Result<()> {
894 let stream = orch(target_set(&["offload/model"]))
897 .run_stream(request(), Arc::new(RuntimeModels::default()));
898 tokio::pin!(stream);
899
900 let mut saw_call = false;
901 let mut final_completion = None;
902 while let Some(step) = stream.next().await {
903 match step? {
904 Step::CallModel(call) => {
905 saw_call = true;
906 assert_eq!(call.models, vec![ModelId::from("offload/model")]);
907 call.respond(Ok(Response {
909 llm_response: LlmResponse::Agg(text_response(
910 None,
911 "fulfilled".to_string(),
912 )),
913 metadata: None,
914 upstream_headers: http::HeaderMap::new(),
915 }))?;
916 }
917 Step::Done(outcome) => {
918 let metadata = outcome
919 .metadata
920 .as_ref()
921 .expect("run_stream should attach outcome metadata");
922 assert_eq!(metadata.algorithm, "test");
923 assert_eq!(
924 uuid::Uuid::parse_str(metadata.outcome_id())
925 .expect("outcome id should be a UUID")
926 .get_version_num(),
927 7
928 );
929 assert_eq!(
930 metadata.evidence,
931 Some(serde_json::json!({"source": "test"}))
932 );
933 let response = outcome
934 .response
935 .ok_or_else(|| test_error("expected an answered outcome"))?;
936 final_completion = Some(
937 response
938 .llm_response
939 .as_agg()
940 .map(completion_text)
941 .unwrap_or_default(),
942 );
943 }
944 }
945 }
946
947 assert!(saw_call, "expected a CallModel step before Done");
948 assert_eq!(
949 final_completion.ok_or_else(|| test_error("no Done step"))?,
950 "fulfilled"
951 );
952 Ok(())
953 }
954
955 #[tokio::test(flavor = "multi_thread", worker_threads = 12)]
956 async fn requests_are_processed_in_parallel() -> Result<()> {
957 use std::time::Duration;
958 use tokio::sync::Barrier;
959
960 const N: usize = 12;
961
962 let barrier = Arc::new(Barrier::new(N));
967 let algo = orch(target_set(&["m"]));
969
970 let mut handles = Vec::new();
971 for _ in 0..N {
972 let algo = algo.clone();
973 let barrier = barrier.clone();
974 let serve = move |target: ModelId, _request: Request| {
975 let barrier = barrier.clone();
976 async move {
977 barrier.wait().await;
978 Ok(reply(target))
979 }
980 };
981 handles.push(tokio::spawn(async move {
982 test_drive(algo, request(), serve)
983 .await
984 .map(|(_, response)| {
985 response
986 .llm_response
987 .as_agg()
988 .map(completion_text)
989 .unwrap_or_default()
990 })
991 }));
992 }
993
994 for handle in handles {
995 let completion = tokio::time::timeout(Duration::from_secs(5), handle)
997 .await
998 .map_err(|error| LibsyError::external("waiting for test task", error))?
999 .map_err(|source| LibsyError::external("joining a test task", source))??;
1000 assert_eq!(completion, "m");
1001 }
1002 Ok(())
1003 }
1004
1005 #[tokio::test]
1006 async fn offload_error_propagates_back_to_the_algorithm() -> Result<()> {
1007 let stream = orch(target_set(&["offload/model"]))
1011 .run_stream(request(), Arc::new(RuntimeModels::default()));
1012 tokio::pin!(stream);
1013
1014 let mut saw_error = false;
1015 while let Some(step) = stream.next().await {
1016 match step {
1017 Ok(Step::CallModel(call)) => {
1018 call.respond(Err(test_error("upstream model call failed")))?;
1019 }
1020 Ok(Step::Done(..)) => {
1021 return Err(test_error(
1022 "expected the offload error to propagate, got a response",
1023 ));
1024 }
1025 Err(err) => {
1026 assert!(err.to_string().contains("upstream model call failed"));
1028 saw_error = true;
1029 }
1030 }
1031 }
1032
1033 assert!(saw_error, "expected an error step");
1034 Ok(())
1035 }
1036
1037 #[tokio::test]
1038 async fn dropping_the_stream_cancels_the_algorithm_task() -> Result<()> {
1039 use std::sync::atomic::{AtomicBool, Ordering};
1040 use std::time::Duration;
1041 use tokio::sync::mpsc;
1042
1043 struct DropGuard(Arc<AtomicBool>);
1046 impl Drop for DropGuard {
1047 fn drop(&mut self) {
1048 self.0.store(true, Ordering::SeqCst);
1049 }
1050 }
1051
1052 struct StuckAlgo {
1053 started: mpsc::UnboundedSender<()>,
1054 dropped: Arc<AtomicBool>,
1055 }
1056
1057 #[async_trait]
1058 impl Algorithm for StuckAlgo {
1059 fn name(&self) -> &str {
1060 "stuck"
1061 }
1062
1063 async fn route(
1064 self: Arc<Self>,
1065 _driver: Driver,
1066 _request: Request,
1067 ) -> Result<RoutingOutcome> {
1068 let _guard = DropGuard(self.dropped.clone());
1069 let _ = self.started.send(());
1070 std::future::pending::<()>().await;
1072 unreachable!()
1073 }
1074 }
1075
1076 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1077 let dropped = Arc::new(AtomicBool::new(false));
1078 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1079 started: started_tx,
1080 dropped: dropped.clone(),
1081 });
1082
1083 let stream = algo.run_stream(request(), Arc::new(RuntimeModels::default()));
1084 started_rx
1085 .recv()
1086 .await
1087 .ok_or_else(|| test_error("task never started"))?;
1088 drop(stream);
1089 tokio::time::sleep(Duration::from_millis(100)).await;
1090
1091 assert!(
1092 dropped.load(Ordering::SeqCst),
1093 "algorithm task was NOT cancelled after dropping the stream"
1094 );
1095 Ok(())
1096 }
1097
1098 #[tokio::test]
1099 async fn route_panic_surfaces_as_a_stream_error() -> Result<()> {
1100 struct Panicky;
1103
1104 #[async_trait]
1105 impl Algorithm for Panicky {
1106 fn name(&self) -> &str {
1107 "panicky"
1108 }
1109
1110 async fn route(
1111 self: Arc<Self>,
1112 _driver: Driver,
1113 _request: Request,
1114 ) -> Result<RoutingOutcome> {
1115 panic!("boom");
1116 }
1117 }
1118
1119 let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
1120 let stream = algo.run_stream(request(), Arc::new(RuntimeModels::default()));
1121 tokio::pin!(stream);
1122
1123 let mut saw_error = false;
1124 while let Some(step) = stream.next().await {
1125 match step {
1126 Err(err) => {
1127 assert!(err.to_string().contains("algorithm task panicked: boom"));
1129 saw_error = true;
1130 }
1131 Ok(_) => return Err(test_error("expected the panic to surface as an error step")),
1132 }
1133 }
1134
1135 assert!(saw_error, "expected an error step from the panicked task");
1136 Ok(())
1137 }
1138
1139 #[tokio::test]
1143 async fn a_panic_with_a_leaked_driver_clone_still_terminates_the_run() -> Result<()> {
1144 struct LeakyPanic;
1145
1146 #[async_trait]
1147 impl Algorithm for LeakyPanic {
1148 fn name(&self) -> &str {
1149 "leaky_panic"
1150 }
1151
1152 async fn route(
1153 self: Arc<Self>,
1154 driver: Driver,
1155 _request: Request,
1156 ) -> Result<RoutingOutcome> {
1157 tokio::spawn(async move {
1158 let _keep_alive = driver;
1160 std::future::pending::<()>().await;
1161 });
1162 tokio::task::yield_now().await;
1163 panic!("boom");
1164 }
1165 }
1166
1167 let algo: Arc<dyn Algorithm> = Arc::new(LeakyPanic);
1168 let result = tokio::time::timeout(
1170 std::time::Duration::from_secs(1),
1171 test_drive(algo, request(), echo()),
1172 )
1173 .await
1174 .map_err(|error| LibsyError::external("waiting for the panicked run to end", error))?;
1175
1176 match result {
1177 Ok(_) => Err(test_error(
1178 "expected the panic to end the run with an error",
1179 )),
1180 Err(err) => {
1181 assert!(err.to_string().contains("algorithm task panicked: boom"));
1182 Ok(())
1183 }
1184 }
1185 }
1186
1187 #[tokio::test]
1188 async fn cancelling_run_cancels_the_algorithm_task() -> Result<()> {
1189 use std::sync::atomic::{AtomicBool, Ordering};
1190 use std::time::Duration;
1191 use tokio::sync::mpsc;
1192
1193 struct DropGuard(Arc<AtomicBool>);
1196 impl Drop for DropGuard {
1197 fn drop(&mut self) {
1198 self.0.store(true, Ordering::SeqCst);
1199 }
1200 }
1201
1202 struct StuckAlgo {
1203 started: mpsc::UnboundedSender<()>,
1204 dropped: Arc<AtomicBool>,
1205 }
1206
1207 #[async_trait]
1208 impl Algorithm for StuckAlgo {
1209 fn name(&self) -> &str {
1210 "stuck"
1211 }
1212
1213 async fn route(
1214 self: Arc<Self>,
1215 _driver: Driver,
1216 _request: Request,
1217 ) -> Result<RoutingOutcome> {
1218 let _guard = DropGuard(self.dropped.clone());
1219 let _ = self.started.send(());
1220 std::future::pending::<()>().await;
1223 unreachable!()
1224 }
1225 }
1226
1227 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1228 let dropped = Arc::new(AtomicBool::new(false));
1229 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1230 started: started_tx,
1231 dropped: dropped.clone(),
1232 });
1233
1234 let run_task = tokio::spawn(async move { test_drive(algo, request(), echo()).await });
1237 started_rx
1238 .recv()
1239 .await
1240 .ok_or_else(|| test_error("task never started"))?;
1241 run_task.abort();
1242 tokio::time::sleep(Duration::from_millis(100)).await;
1243
1244 assert!(
1245 dropped.load(Ordering::SeqCst),
1246 "algorithm task was NOT cancelled after cancelling run"
1247 );
1248 Ok(())
1249 }
1250
1251 struct Hedge {
1256 winner: String,
1257 loser: String,
1258 }
1259
1260 #[async_trait]
1261 impl Algorithm for Hedge {
1262 fn name(&self) -> &str {
1263 "hedge"
1264 }
1265
1266 async fn route(
1267 self: Arc<Self>,
1268 driver: Driver,
1269 request: Request,
1270 ) -> Result<RoutingOutcome> {
1271 let outcome_request = request.clone();
1272 let win = driver.call_model(request.clone(), vec![self.winner.clone().into()]);
1273 let lose = driver.call_model(request, vec![self.loser.clone().into()]);
1274 tokio::select! {
1276 res = win => Ok(RoutingOutcome::answered(
1277 self.winner.clone().into(),
1278 outcome_request,
1279 res?,
1280 )),
1281 res = lose => Ok(RoutingOutcome::answered(
1282 self.loser.clone().into(),
1283 outcome_request,
1284 res?,
1285 )),
1286 }
1287 }
1288 }
1289
1290 fn hedge(loser_delay: Option<std::time::Duration>) -> (Arc<dyn Algorithm>, impl Serve) {
1294 let started = Arc::new(tokio::sync::Notify::new());
1295 let algo = Arc::new(Hedge {
1296 winner: "winner".to_string(),
1297 loser: "loser".to_string(),
1298 });
1299 let serve = move |target: ModelId, _request: Request| {
1300 let started = started.clone();
1301 async move {
1302 if target == "loser" {
1303 started.notify_one();
1304 match loser_delay {
1305 Some(delay) => tokio::time::sleep(delay).await,
1306 None => std::future::pending::<()>().await,
1307 }
1308 } else {
1309 started.notified().await;
1310 }
1311 Ok(reply(target))
1312 }
1313 };
1314 (algo, serve)
1315 }
1316
1317 #[tokio::test]
1318 async fn run_returns_the_winner_without_a_late_loser_overwriting_it() -> Result<()> {
1319 let (algo, serve) = hedge(Some(std::time::Duration::from_millis(50)));
1322 let (_, response) = test_drive(algo, request(), serve).await?;
1323 assert_eq!(
1324 response
1325 .llm_response
1326 .as_agg()
1327 .map(completion_text)
1328 .unwrap_or_default(),
1329 "winner"
1330 );
1331 Ok(())
1332 }
1333
1334 #[tokio::test]
1335 async fn run_returns_the_winner_without_hanging_on_a_pending_loser() -> Result<()> {
1336 let (algo, serve) = hedge(None);
1339 let run = test_drive(algo, request(), serve);
1340 let (_, response) = tokio::time::timeout(std::time::Duration::from_secs(1), run)
1341 .await
1342 .map_err(|error| LibsyError::external("waiting for pending loser", error))??;
1343 assert_eq!(
1344 response
1345 .llm_response
1346 .as_agg()
1347 .map(completion_text)
1348 .unwrap_or_default(),
1349 "winner"
1350 );
1351 Ok(())
1352 }
1353
1354 #[tokio::test]
1355 async fn run_surfaces_a_terminal_error_with_many_calls_in_flight() -> Result<()> {
1356 use std::sync::atomic::{AtomicUsize, Ordering};
1357
1358 const N: usize = 10;
1361
1362 struct FanOutThenError {
1365 all_started: Arc<tokio::sync::Notify>,
1366 n: usize,
1367 }
1368
1369 #[async_trait]
1370 impl Algorithm for FanOutThenError {
1371 fn name(&self) -> &str {
1372 "fan_out_then_error"
1373 }
1374
1375 async fn route(
1376 self: Arc<Self>,
1377 driver: Driver,
1378 request: Request,
1379 ) -> Result<RoutingOutcome> {
1380 let offloads = futures::future::join_all(
1381 (0..self.n)
1382 .map(|i| driver.call_model(request.clone(), vec![format!("m{i}").into()])),
1383 );
1384 tokio::select! {
1385 _ = offloads => Err(test_error("offloads unexpectedly completed")),
1386 _ = self.all_started.notified() => {
1387 Err(test_error("terminal error while calls pending"))
1388 }
1389 }
1390 }
1391 }
1392
1393 let all_started = Arc::new(tokio::sync::Notify::new());
1394 let algo: Arc<dyn Algorithm> = Arc::new(FanOutThenError {
1395 all_started: all_started.clone(),
1396 n: N,
1397 });
1398
1399 let started = Arc::new(AtomicUsize::new(0));
1401 let serve = move |_target: ModelId, _request: Request| {
1402 let started = started.clone();
1403 let all_started = all_started.clone();
1404 async move {
1405 if started.fetch_add(1, Ordering::SeqCst) + 1 == N {
1406 all_started.notify_one();
1407 }
1408 std::future::pending::<ServeResult>().await
1409 }
1410 };
1411
1412 let run = test_drive(algo, request(), serve);
1415 let result = tokio::time::timeout(std::time::Duration::from_millis(500), run)
1416 .await
1417 .map_err(|error| {
1418 LibsyError::external("waiting for terminal error with full call cap", error)
1419 })?;
1420 match result {
1421 Ok(_) => Err(test_error("expected the terminal error, got a response")),
1422 Err(err) => {
1423 assert!(
1424 err.to_string()
1425 .contains("terminal error while calls pending")
1426 );
1427 Ok(())
1428 }
1429 }
1430 }
1431}