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::{
28 Category, DecisionRequest, DecisionResponse, ModelId, Request, Response,
29};
30
31use crate::{DriverError, LibsyError, Result, observability};
32
33pub type StepStream = Pin<Box<dyn Stream<Item = Result<Step>> + Send>>;
37
38#[derive(Clone, Debug, Default)]
49pub struct RuntimeModels {
50 by_category: HashMap<Category, Vec<ModelId>>,
52 subagent: Option<HashMap<Category, Vec<ModelId>>>,
53}
54
55impl RuntimeModels {
56 pub fn new(by_category: HashMap<Category, Vec<ModelId>>) -> Self {
58 Self {
59 by_category,
60 subagent: None,
61 }
62 }
63
64 pub fn with_subagent(mut self, models: HashMap<Category, Vec<ModelId>>) -> Self {
66 self.subagent = Some(models);
67 self
68 }
69
70 pub fn models_for(&self, category: &Category) -> &[ModelId] {
72 self.by_category.get(category).map_or(&[], Vec::as_slice)
73 }
74
75 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#[derive(Clone, Copy)]
92enum Scope {
93 Parent,
94 Subagent,
95}
96
97pub struct CallModel {
107 pub algorithm: String,
110 pub request: Request,
112 pub models: Vec<ModelId>,
114 pub recover_errors: bool,
116 reply: Option<oneshot::Sender<Result<Response>>>,
118 started: Instant,
119 span: tracing::Span,
121}
122
123impl CallModel {
124 pub fn respond(mut self, result: Result<Response>) -> Result<()> {
128 self.record(result.is_ok());
129 if let Ok(response) = &result {
130 observability::record_llm_response(response, &self.span);
131 }
132 self.reply
133 .take()
134 .ok_or(DriverError::ResponseDropped)?
135 .send(result)
136 .map_err(|_| DriverError::ResponseDropped.into())
137 }
138
139 pub fn fail(mut self, error: LibsyError) -> Result<()> {
142 self.reply = None;
143 self.record(false);
144 Err(error)
145 }
146
147 fn record(&self, is_ok: bool) {
148 observability::record_llm_call(
149 &self.algorithm,
150 self.models
151 .first()
152 .map(ModelId::as_str)
153 .unwrap_or("NoTargets"),
154 self.started.elapsed(),
155 is_ok,
156 );
157 self.span
158 .record("outcome", if is_ok { "ok" } else { "error" });
159 }
160}
161
162impl Drop for CallModel {
163 fn drop(&mut self) {
164 if self.reply.is_some() {
165 self.record(false);
166 }
167 }
168}
169
170pub struct CallDecision {
173 pub algorithm: String,
175 pub request: DecisionRequest,
177 pub model: ModelId,
179 reply: Option<oneshot::Sender<Result<DecisionResponse>>>,
180 started: Instant,
181 span: tracing::Span,
183}
184
185impl CallDecision {
186 pub fn respond(mut self, result: Result<DecisionResponse>) -> Result<()> {
188 self.record(result.is_ok());
189 if let Ok(response) = &result {
190 observability::record_decision_response(response, &self.span);
191 }
192 self.reply
193 .take()
194 .ok_or(DriverError::ResponseDropped)?
195 .send(result)
196 .map_err(|_| DriverError::ResponseDropped.into())
197 }
198
199 pub fn fail(mut self, error: LibsyError) -> Result<()> {
201 self.reply = None;
202 self.record(false);
203 Err(error)
204 }
205
206 fn record(&self, is_ok: bool) {
207 observability::record_decision_call(
208 &self.algorithm,
209 &self.model,
210 self.started.elapsed(),
211 is_ok,
212 );
213 self.span
214 .record("outcome", if is_ok { "ok" } else { "error" });
215 }
216}
217
218impl Drop for CallDecision {
219 fn drop(&mut self) {
220 if self.reply.is_some() {
221 self.record(false);
222 }
223 }
224}
225
226pub struct RoutingOutcome {
228 pub selected_model_ids: Vec<ModelId>,
230 pub request: Request,
232 pub response: Option<Response>,
234 pub metadata: Option<crate::OutcomeMetadata>,
239}
240
241impl RoutingOutcome {
242 pub fn selected_model_id(&self) -> Result<&ModelId> {
245 self.selected_model_ids.first().ok_or(LibsyError::NoTargets)
246 }
247
248 pub fn route_to(
252 selected_model_id: ModelId,
253 fallback_models: Vec<ModelId>,
254 mut request: Request,
255 ) -> Self {
256 request.llm_request.model = Some(selected_model_id.to_string());
257 let mut selected_model_ids = Vec::with_capacity(1 + fallback_models.len());
258 selected_model_ids.push(selected_model_id);
259 selected_model_ids.extend(fallback_models);
260 Self {
261 selected_model_ids,
262 request,
263 response: None,
264 metadata: None,
265 }
266 }
267
268 pub fn answered(selected_model_id: ModelId, mut request: Request, response: Response) -> Self {
271 request.llm_request.model = Some(selected_model_id.to_string());
272 Self {
273 selected_model_ids: vec![selected_model_id],
274 request,
275 response: Some(response),
276 metadata: None,
277 }
278 }
279}
280
281#[derive(Clone)]
283pub struct Driver {
284 step_tx: mpsc::Sender<Result<Step>>,
285
286 algorithm: String,
288
289 evidence: Arc<Mutex<Option<Value>>>,
291
292 models: Arc<RuntimeModels>,
294
295 scope: Scope,
297}
298
299impl Driver {
300 pub(crate) fn new(
303 algorithm: &str,
304 models: Arc<RuntimeModels>,
305 ) -> (Self, mpsc::Receiver<Result<Step>>) {
306 let (step_tx, step_rx) = mpsc::channel(1);
311 (
312 Self {
313 step_tx,
314 algorithm: algorithm.to_string(),
315 evidence: Arc::new(Mutex::new(None)),
316 models,
317 scope: Scope::Parent,
318 },
319 step_rx,
320 )
321 }
322
323 pub(crate) fn set_evidence(&self, evidence: Value) {
325 *self.evidence.lock() = Some(evidence);
326 }
327
328 pub(crate) fn set_evidence_if_empty(&self, evidence: Value) {
330 let mut current = self.evidence.lock();
331 if current.is_none() {
332 *current = Some(evidence);
333 }
334 }
335
336 pub async fn call_model(&self, request: Request, models: Vec<ModelId>) -> Result<Response> {
346 self.call_model_with_error_recovery(request, models, false)
347 .await
348 }
349
350 #[tracing::instrument(
352 target = "libsy",
353 name = "libsy.llm_call",
354 skip_all,
355 fields(
356 algorithm = self.algorithm,
357 selected_model = %models.first().map(ModelId::as_str).unwrap_or("NoTargets"),
358 openinference.span.kind = "CHAIN",
359 outcome = tracing::field::Empty,
360 input_tokens = tracing::field::Empty,
361 output_tokens = tracing::field::Empty,
362 total_tokens = tracing::field::Empty,
363 reasoning_tokens = tracing::field::Empty,
364 )
365 )]
366 pub(crate) async fn call_model_with_error_recovery(
367 &self,
368 mut request: Request,
369 models: Vec<ModelId>,
370 recover_errors: bool,
371 ) -> Result<Response> {
372 let Some(selected_model_id) = models.first() else {
373 return Err(LibsyError::NoTargets);
374 };
375 request.llm_request.model = Some(selected_model_id.to_string());
376 let started = Instant::now();
377 let (reply, response) = oneshot::channel::<Result<Response>>();
378 let call = CallModel {
379 algorithm: self.algorithm.clone(),
380 request,
381 models,
382 recover_errors,
383 reply: Some(reply),
384 started,
385 span: tracing::Span::current(),
386 };
387 self.step_tx
388 .send(Ok(Step::CallModel(Box::new(call))))
389 .await
390 .map_err(|_| DriverError::StreamClosed)?;
391 response
392 .await
393 .map_err(|_| LibsyError::from(DriverError::ResponseDropped))?
394 }
395
396 #[tracing::instrument(
398 target = "libsy",
399 name = "libsy.decision_call",
400 skip_all,
401 fields(
402 algorithm = self.algorithm,
403 selected_model = %model,
404 openinference.span.kind = "CHAIN",
405 outcome = tracing::field::Empty,
406 input_tokens = tracing::field::Empty,
407 output_tokens = tracing::field::Empty,
408 total_tokens = tracing::field::Empty,
409 reasoning_tokens = tracing::field::Empty,
410 gen_ai.response.id = tracing::field::Empty,
411 gen_ai.response.model = tracing::field::Empty,
412 ),
413 )]
414 pub async fn call_decision(
415 &self,
416 mut request: DecisionRequest,
417 model: ModelId,
418 ) -> Result<DecisionResponse> {
419 request.model = Some(model.clone());
420 let started = Instant::now();
421 let (reply, response) = oneshot::channel();
422 let call = CallDecision {
423 algorithm: self.algorithm.clone(),
424 request,
425 model,
426 reply: Some(reply),
427 started,
428 span: tracing::Span::current(),
429 };
430 self.step_tx
431 .send(Ok(Step::CallDecision(Box::new(call))))
432 .await
433 .map_err(|_| DriverError::StreamClosed)?;
434 response
435 .await
436 .map_err(|_| LibsyError::from(DriverError::ResponseDropped))?
437 }
438
439 pub fn models_for(&self, category: &Category) -> &[ModelId] {
441 match self.scope {
442 Scope::Parent => self.models.models_for(category),
443 Scope::Subagent => self.models.subagent_models_for(category),
444 }
445 }
446
447 pub fn first_model_for(&self, category: &Category) -> Result<&ModelId> {
449 self.models_for(category)
450 .first()
451 .ok_or_else(|| LibsyError::AlgorithmError {
452 message: format!("no models available for category {}", category.as_str()),
453 })
454 }
455
456 pub fn for_subagent(&self) -> Result<Self> {
459 if self.models.subagent.is_none() {
460 return Err(LibsyError::AlgorithmError {
461 message: "delegated work has no sub-agent models".to_string(),
462 });
463 }
464 Ok(Self {
465 scope: Scope::Subagent,
466 ..self.clone()
467 })
468 }
469
470 pub(crate) async fn finish(&self, result: Result<RoutingOutcome>) -> Result<()> {
474 let result = result.map(|mut outcome| {
475 let metadata = outcome.metadata.get_or_insert_with(|| {
476 crate::OutcomeMetadata::new(self.algorithm.clone(), self.evidence.lock().take())
477 });
478 observability::record_outcome(metadata, &outcome.selected_model_ids);
479 outcome
480 });
481 let selected_model = result
482 .as_ref()
483 .ok()
484 .and_then(|outcome| outcome.selected_model_id().ok().cloned());
485 let step = result.map(|outcome| Step::Done(Box::new(outcome)));
486 self.step_tx
487 .send(step)
488 .await
489 .map_err(|_| DriverError::StreamClosed)?;
490 if let Some(selected_model) = selected_model {
491 observability::record_decision(&self.algorithm, &selected_model);
492 }
493 Ok(())
494 }
495}
496
497pub enum Step {
500 CallModel(Box<CallModel>),
503 CallDecision(Box<CallDecision>),
505 Done(Box<RoutingOutcome>),
507}
508
509pub enum Call {
512 Model(Box<CallModel>),
514 Decision(Box<CallDecision>),
516}
517
518pub async fn drive<F, Fut>(
528 algorithm: Arc<dyn Algorithm>,
529 request: Request,
530 models: Arc<RuntimeModels>,
531 serve: F,
532) -> Result<RoutingOutcome>
533where
534 F: Fn(Call) -> Fut,
535 Fut: Future<Output = Result<()>>,
536{
537 let stream = algorithm.run_stream(request, models);
538 tokio::pin!(stream);
539
540 let mut in_flight = futures::stream::FuturesUnordered::new();
541 let mut final_outcome: Option<RoutingOutcome> = None;
542
543 loop {
544 tokio::select! {
545 Some(result) = in_flight.next() => match result {
546 Ok(()) => {}, Err(err) => return Err(err), },
549 step = stream.next() => {
550 match step {
551 None => break, Some(item) => match item? {
553 Step::CallModel(call) => in_flight.push(serve(Call::Model(call))),
554 Step::CallDecision(call) => {
555 in_flight.push(serve(Call::Decision(call)));
556 }
557 Step::Done(outcome) => {
558 final_outcome = Some(*outcome);
559 break;
560 }
561 }
562 }
563 },
564 }
565 }
566 final_outcome.ok_or(LibsyError::MissingFinalResponse)
567}
568
569fn panic_message(payload: &(dyn std::any::Any + Send)) -> String {
571 payload
572 .downcast_ref::<&'static str>()
573 .map(|message| (*message).to_string())
574 .or_else(|| payload.downcast_ref::<String>().cloned())
575 .unwrap_or_else(|| "unknown panic payload".to_string())
576}
577
578struct AbortOnDrop(tokio::task::AbortHandle);
580
581impl Drop for AbortOnDrop {
582 fn drop(&mut self) {
583 self.0.abort();
584 }
585}
586
587pub(crate) fn ensure_model_is_target(targets: &[ModelId], name: &ModelId) -> Result<()> {
592 targets
593 .iter()
594 .any(|target| target == name)
595 .then_some(())
596 .ok_or_else(|| LibsyError::TargetNotFound {
597 target: name.clone(),
598 })
599}
600
601#[derive(Clone, Hash, PartialEq, Eq)]
604pub(crate) enum RoutingIdentity {
605 Session(String),
607 Subagent { session: String, agent: String },
609}
610
611impl RoutingIdentity {
612 pub(crate) fn from_request(request: &Request) -> Option<Self> {
617 let metadata = request.metadata.as_ref()?;
618 let session = metadata.session_id.as_deref().filter(|id| !id.is_empty())?;
619 if metadata.is_subagent {
620 let agent = metadata.agent_id.as_deref().filter(|id| !id.is_empty())?;
621 Some(Self::Subagent {
622 session: session.to_string(),
623 agent: agent.to_string(),
624 })
625 } else {
626 Some(Self::Session(session.to_string()))
627 }
628 }
629}
630
631#[async_trait]
666pub trait Algorithm: Send + Sync + 'static {
667 fn name(&self) -> &str;
671
672 async fn route(self: Arc<Self>, driver: Driver, request: Request) -> Result<RoutingOutcome>;
675
676 fn run_stream(self: Arc<Self>, request: Request, models: Arc<RuntimeModels>) -> StepStream {
684 let (driver, step_rx) = Driver::new(self.name(), models);
685 let span = observability::run_span(self.name(), &request);
686 let handle = tokio::spawn(
687 async move {
688 let algorithm = self.name().to_string();
689 let route = AssertUnwindSafe(self.route(driver.clone(), request)).catch_unwind();
691 let result = observability::observe_run(&algorithm, async move {
692 route.await.unwrap_or_else(|payload| {
693 Err(LibsyError::AlgorithmError {
694 message: format!(
695 "algorithm task panicked: {}",
696 panic_message(payload.as_ref())
697 ),
698 })
699 })
700 })
701 .await;
702
703 let _ = driver.finish(result).await;
704 }
705 .instrument(span),
706 );
707 let abort_guard = AbortOnDrop(handle.abort_handle());
709 Box::pin(ReceiverStream::new(step_rx).map(move |step| {
710 let _keep_alive = &abort_guard;
712 step
713 }))
714 }
715}
716
717#[cfg(test)]
718mod tests {
719 use std::collections::HashMap;
720
721 use super::*;
722 use crate::core::testing::{Serve, ServeResult, echo, reply, serve_decision, test_drive};
723 use futures::StreamExt;
724 use switchyard_protocol::{
725 LlmResponse, LlmResponseChunk, completion_text, text_request, text_response,
726 };
727
728 #[derive(Debug, thiserror::Error)]
729 #[error("{0}")]
730 struct TestError(&'static str);
731
732 fn test_error(message: &'static str) -> LibsyError {
733 LibsyError::external("test", TestError(message))
734 }
735
736 struct TestAlgo {
739 target_set: Vec<ModelId>,
740 }
741
742 #[async_trait]
743 impl Algorithm for TestAlgo {
744 fn name(&self) -> &str {
745 "test"
746 }
747
748 async fn route(
749 self: Arc<Self>,
750 driver: Driver,
751 request: Request,
752 ) -> Result<RoutingOutcome> {
753 let target = self
754 .target_set
755 .first()
756 .ok_or(LibsyError::NoTargets)?
757 .clone();
758 let response = driver
759 .call_model(request.clone(), vec![target.clone()])
760 .await?;
761 driver.set_evidence(serde_json::json!({"source": "test"}));
762 driver.set_evidence_if_empty(serde_json::json!({"source": "ignored"}));
763 Ok(RoutingOutcome::answered(target, request, response))
764 }
765 }
766
767 fn orch(target_set: Vec<ModelId>) -> Arc<dyn Algorithm> {
769 Arc::new(TestAlgo { target_set })
770 }
771
772 fn request() -> Request {
773 Request {
774 llm_request: text_request(Some("auto".to_string()), "hi".to_string()),
775 raw_request: None,
776 metadata: None,
777 }
778 }
779
780 #[test]
781 fn routing_outcome_constructors_stamp_selection_and_preserve_payloads() {
782 let outcome = RoutingOutcome::route_to(
783 "selected".into(),
784 target_set(&["fallback-one", "fallback-two"]),
785 request(),
786 );
787
788 assert_eq!(
789 outcome.selected_model_ids,
790 target_set(&["selected", "fallback-one", "fallback-two"])
791 );
792 assert_eq!(outcome.request.model_id().as_deref(), Some("selected"));
793 assert!(outcome.response.is_none());
794 assert!(outcome.metadata.is_none());
795
796 let outcome = RoutingOutcome::route_to("only".into(), Vec::new(), request());
797 assert_eq!(outcome.selected_model_ids, target_set(&["only"]));
798
799 let outcome = RoutingOutcome::answered(
800 "answered".into(),
801 request(),
802 Response {
803 llm_response: LlmResponse::Agg(text_response(None, "existing")),
804 metadata: None,
805 upstream_headers: http::HeaderMap::new(),
806 },
807 );
808
809 assert_eq!(outcome.selected_model_ids, target_set(&["answered"]));
810 assert_eq!(outcome.request.model_id().as_deref(), Some("answered"));
811 assert_eq!(
812 outcome
813 .response
814 .as_ref()
815 .and_then(|response| response.llm_response.as_agg())
816 .map(completion_text),
817 Some("existing".to_string())
818 );
819 }
820
821 fn target_set(names: &[&str]) -> Vec<ModelId> {
822 names.iter().map(|name| ModelId::from(*name)).collect()
823 }
824
825 #[tokio::test]
826 async fn typed_driver_preserves_call_and_stream_boundaries() -> Result<()> {
827 tokio::time::timeout(std::time::Duration::from_secs(1), async {
828 let (driver, mut step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
831 let first_driver = driver.clone();
832 let mut first = tokio::spawn(async move {
833 first_driver
834 .call_model(request(), vec![ModelId::from("first")])
835 .await
836 });
837 let second = tokio::spawn(async move {
838 driver
839 .call_model(request(), vec![ModelId::from("second")])
840 .await
841 });
842
843 let mut calls = HashMap::new();
844 for _ in 0..2 {
845 let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
846 let Step::CallModel(call) = step else {
847 return Err(test_error("expected a CallModel step"));
848 };
849 let selected_model = call
850 .models
851 .first()
852 .ok_or_else(|| test_error("model call has no candidates"))?
853 .to_string();
854 calls.insert(selected_model, call);
855 }
856 assert!(
857 tokio::time::timeout(std::time::Duration::from_millis(20), &mut first)
858 .await
859 .is_err(),
860 "call completed before the host responded"
861 );
862 calls
863 .remove("second")
864 .ok_or_else(|| test_error("missing second call"))?
865 .respond(Ok(reply("second response")))?;
866 calls
867 .remove("first")
868 .ok_or_else(|| test_error("missing first call"))?
869 .respond(Ok(reply("first response")))?;
870
871 let first_response = first
872 .await
873 .map_err(|source| LibsyError::external("joining a test task", source))??;
874 let second_response = second
875 .await
876 .map_err(|source| LibsyError::external("joining a test task", source))??;
877 assert_eq!(
878 first_response.llm_response.as_agg().map(completion_text),
879 Some("first response".to_string())
880 );
881 assert_eq!(
882 second_response.llm_response.as_agg().map(completion_text),
883 Some("second response".to_string())
884 );
885
886 let (driver, mut step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
888 let producer = tokio::spawn(async move {
889 driver
890 .call_model(request(), vec![ModelId::from("dropped")])
891 .await
892 });
893 let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
894 let Step::CallModel(call) = step else {
895 return Err(test_error("expected a CallModel step"));
896 };
897 drop(call);
898 let result = producer
899 .await
900 .map_err(|source| LibsyError::external("joining a test task", source))?;
901 assert!(matches!(
902 result,
903 Err(LibsyError::Driver(DriverError::ResponseDropped))
904 ));
905
906 let (driver, step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
908 drop(step_rx);
909 let result = driver
910 .call_model(request(), vec![ModelId::from("closed")])
911 .await;
912 assert!(matches!(
913 result,
914 Err(LibsyError::Driver(DriverError::StreamClosed))
915 ));
916
917 fn decision_request() -> DecisionRequest {
918 DecisionRequest {
919 model: Some("overwritten".into()),
920 context: serde_json::json!({"task": "choose a route"}),
921 questions: Default::default(),
922 }
923 }
924
925 fn decision_response() -> DecisionResponse {
926 DecisionResponse {
927 id: Some("decision-1".to_string()),
928 model: Some("provider-model".into()),
929 answers: [(
930 "p_solve".to_string(),
931 switchyard_protocol::DecisionAnswer {
932 value: switchyard_protocol::DecisionValue::Boolean(
933 switchyard_protocol::BooleanEstimate::ProbabilityTrue(
934 switchyard_protocol::Probability(0.8),
935 ),
936 ),
937 provider_confidence: None,
938 },
939 )]
940 .into(),
941 usage: Default::default(),
942 }
943 }
944
945 struct MixedCalls(&'static str);
946
947 #[async_trait]
948 impl Algorithm for MixedCalls {
949 fn name(&self) -> &str {
950 "mixed"
951 }
952
953 async fn route(
954 self: Arc<Self>,
955 driver: Driver,
956 request: Request,
957 ) -> Result<RoutingOutcome> {
958 let (llm, decision) = tokio::join!(
959 driver.call_model(request.clone(), vec!["llm".into()]),
960 driver.call_decision(decision_request(), "decision".into()),
961 );
962 assert_eq!(
963 llm?.llm_response.as_agg().map(completion_text),
964 Some("llm reply".into())
965 );
966 match self.0 {
967 "mock" => assert_eq!(decision?, DecisionResponse {
968 id: None,
969 model: Some("decision".into()),
970 answers: Default::default(),
971 usage: Default::default(),
972 }),
973 "reply" => assert_eq!(decision?, decision_response()),
974 "error" => assert!(matches!(
975 decision,
976 Err(LibsyError::AlgorithmError { message }) if message == "provider failed"
977 )),
978 "drop" => assert!(matches!(
979 decision,
980 Err(LibsyError::Driver(DriverError::ResponseDropped))
981 )),
982 "abort" => return std::future::pending().await,
983 _ => unreachable!(),
984 }
985 Ok(RoutingOutcome::route_to("answer".into(), vec![], request))
986 }
987 }
988
989 for mode in ["mock", "reply", "error", "drop", "abort"] {
990 let barrier = Arc::new(tokio::sync::Barrier::new(2));
992 let outcome = drive(
993 Arc::new(MixedCalls(mode)),
994 request(),
995 Arc::new(RuntimeModels::default()),
996 move |call| {
997 let barrier = barrier.clone();
998 async move {
999 barrier.wait().await;
1000 let call = match call {
1001 Call::Model(call) => return call.respond(Ok(reply("llm reply"))),
1002 Call::Decision(call) => *call,
1003 };
1004 assert_eq!(call.algorithm, "mixed");
1005 assert_eq!(call.model, "decision");
1006 assert_eq!(call.request.model, Some("decision".into()));
1007 assert_eq!(call.request.context, decision_request().context);
1008 match mode {
1009 "mock" => serve_decision(call).await,
1010 "reply" => call.respond(Ok(decision_response())),
1011 "error" => call.respond(Err(LibsyError::AlgorithmError {
1012 message: "provider failed".into(),
1013 })),
1014 "drop" => {
1015 drop(call);
1016 Ok(())
1017 }
1018 "abort" => call.fail(test_error("host aborted")),
1019 _ => unreachable!(),
1020 }
1021 }
1022 },
1023 )
1024 .await;
1025 if mode == "abort" {
1026 assert!(matches!(outcome, Err(LibsyError::External { .. })));
1027 } else {
1028 assert_eq!(outcome?.selected_model_id()?, "answer");
1029 }
1030 }
1031
1032 let (driver, mut steps) = Driver::new("test", Arc::new(RuntimeModels::default()));
1033 let mut pending = Box::pin(driver.call_decision(decision_request(), "decision".into()));
1034 assert!(futures::poll!(&mut pending).is_pending());
1035 let Some(Ok(Step::CallDecision(call))) = steps.recv().await else {
1036 return Err(test_error("expected a decision call"));
1037 };
1038 drop(pending);
1039 assert!(matches!(
1040 call.respond(Ok(decision_response())),
1041 Err(LibsyError::Driver(DriverError::ResponseDropped))
1042 ));
1043 drop(steps);
1044 assert!(matches!(
1045 driver.call_decision(decision_request(), "decision".into()).await,
1046 Err(LibsyError::Driver(DriverError::StreamClosed))
1047 ));
1048 Ok(())
1049 })
1050 .await
1051 .map_err(|error| LibsyError::external("waiting for typed driver boundaries", error))?
1052 }
1053
1054 #[test]
1055 fn target_lookup_returns_the_missing_target() {
1056 let error = ensure_model_is_target(&target_set(&[]), &ModelId::from("missing")).err();
1057 assert!(matches!(
1058 error,
1059 Some(LibsyError::TargetNotFound { target }) if target == "missing"
1060 ));
1061 }
1062
1063 fn streaming_orch(chunks: Vec<LlmResponseChunk>) -> (Arc<dyn Algorithm>, impl Serve) {
1066 let algo = orch(target_set(&["stream/model"]));
1067 let serve = move |_target: ModelId, _request: Request| {
1068 let chunks = chunks.clone();
1069 async move {
1070 let stream =
1071 futures::stream::iter(chunks.into_iter().map(|chunk| Ok(chunk.into()))).boxed();
1072 Ok(Response {
1073 llm_response: LlmResponse::Stream(stream),
1074 metadata: None,
1075 upstream_headers: http::HeaderMap::new(),
1076 })
1077 }
1078 };
1079 (algo, serve)
1080 }
1081
1082 #[tokio::test]
1083 async fn run_returns_a_streamed_response_the_caller_aggregates() -> Result<()> {
1084 let (orch, serve) = streaming_orch(vec![
1087 LlmResponseChunk::MessageStart {
1088 id: Some("m1".to_string()),
1089 model: Some("stream/model".to_string()),
1090 },
1091 LlmResponseChunk::TextDelta {
1092 index: 0,
1093 text: "hel".to_string(),
1094 },
1095 LlmResponseChunk::TextDelta {
1096 index: 0,
1097 text: "lo".to_string(),
1098 },
1099 LlmResponseChunk::MessageStop {
1100 reason: Some("stop".to_string()),
1101 },
1102 ]);
1103 let (selected_model, response) = test_drive(orch, request(), serve).await?;
1104 let agg = response
1106 .llm_response
1107 .into_agg()
1108 .await
1109 .map_err(|error| LibsyError::external("aggregating response stream", error))?;
1110 assert_eq!(completion_text(&agg), "hello");
1111 assert_eq!(agg.model.as_deref(), Some("stream/model"));
1112 assert_eq!(selected_model, "stream/model");
1113 Ok(())
1114 }
1115
1116 #[tokio::test]
1117 async fn aggregating_a_streamed_response_propagates_a_mid_stream_error() -> Result<()> {
1118 let (orch, serve) = streaming_orch(vec![
1121 LlmResponseChunk::TextDelta {
1122 index: 0,
1123 text: "partial".to_string(),
1124 },
1125 LlmResponseChunk::StreamError {
1126 message: "upstream exploded".to_string(),
1127 },
1128 ]);
1129 let (_, response) = test_drive(orch, request(), serve).await?;
1130 match response.llm_response.into_agg().await {
1131 Ok(_) => panic!("expected a mid-stream error, got an aggregate"),
1132 Err(err) => {
1133 assert!(err.to_string().contains("upstream exploded"));
1134 Ok(())
1135 }
1136 }
1137 }
1138
1139 #[tokio::test]
1140 async fn run_offloads_via_promise_then_finishes() -> Result<()> {
1141 let stream = orch(target_set(&["offload/model"]))
1144 .run_stream(request(), Arc::new(RuntimeModels::default()));
1145 tokio::pin!(stream);
1146
1147 let mut saw_call = false;
1148 let mut final_completion = None;
1149 while let Some(step) = stream.next().await {
1150 match step? {
1151 Step::CallDecision(_) => return Err(test_error("unexpected decision call")),
1152 Step::CallModel(call) => {
1153 saw_call = true;
1154 assert_eq!(call.models, vec![ModelId::from("offload/model")]);
1155 call.respond(Ok(Response {
1157 llm_response: LlmResponse::Agg(text_response(
1158 None,
1159 "fulfilled".to_string(),
1160 )),
1161 metadata: None,
1162 upstream_headers: http::HeaderMap::new(),
1163 }))?;
1164 }
1165 Step::Done(outcome) => {
1166 let metadata = outcome
1167 .metadata
1168 .as_ref()
1169 .expect("run_stream should attach outcome metadata");
1170 assert_eq!(metadata.algorithm, "test");
1171 assert_eq!(
1172 uuid::Uuid::parse_str(metadata.outcome_id())
1173 .expect("outcome id should be a UUID")
1174 .get_version_num(),
1175 7
1176 );
1177 assert_eq!(
1178 metadata.evidence,
1179 Some(serde_json::json!({"source": "test"}))
1180 );
1181 let response = outcome
1182 .response
1183 .ok_or_else(|| test_error("expected an answered outcome"))?;
1184 final_completion = Some(
1185 response
1186 .llm_response
1187 .as_agg()
1188 .map(completion_text)
1189 .unwrap_or_default(),
1190 );
1191 }
1192 }
1193 }
1194
1195 assert!(saw_call, "expected a CallModel step before Done");
1196 assert_eq!(
1197 final_completion.ok_or_else(|| test_error("no Done step"))?,
1198 "fulfilled"
1199 );
1200 Ok(())
1201 }
1202
1203 #[tokio::test(flavor = "multi_thread", worker_threads = 12)]
1204 async fn requests_are_processed_in_parallel() -> Result<()> {
1205 use std::time::Duration;
1206 use tokio::sync::Barrier;
1207
1208 const N: usize = 12;
1209
1210 let barrier = Arc::new(Barrier::new(N));
1215 let algo = orch(target_set(&["m"]));
1217
1218 let mut handles = Vec::new();
1219 for _ in 0..N {
1220 let algo = algo.clone();
1221 let barrier = barrier.clone();
1222 let serve = move |target: ModelId, _request: Request| {
1223 let barrier = barrier.clone();
1224 async move {
1225 barrier.wait().await;
1226 Ok(reply(target))
1227 }
1228 };
1229 handles.push(tokio::spawn(async move {
1230 test_drive(algo, request(), serve)
1231 .await
1232 .map(|(_, response)| {
1233 response
1234 .llm_response
1235 .as_agg()
1236 .map(completion_text)
1237 .unwrap_or_default()
1238 })
1239 }));
1240 }
1241
1242 for handle in handles {
1243 let completion = tokio::time::timeout(Duration::from_secs(5), handle)
1245 .await
1246 .map_err(|error| LibsyError::external("waiting for test task", error))?
1247 .map_err(|source| LibsyError::external("joining a test task", source))??;
1248 assert_eq!(completion, "m");
1249 }
1250 Ok(())
1251 }
1252
1253 #[tokio::test]
1254 async fn offload_error_propagates_back_to_the_algorithm() -> Result<()> {
1255 let stream = orch(target_set(&["offload/model"]))
1259 .run_stream(request(), Arc::new(RuntimeModels::default()));
1260 tokio::pin!(stream);
1261
1262 let mut saw_error = false;
1263 while let Some(step) = stream.next().await {
1264 match step {
1265 Ok(Step::CallDecision(_)) => return Err(test_error("unexpected decision call")),
1266 Ok(Step::CallModel(call)) => {
1267 call.respond(Err(test_error("upstream model call failed")))?;
1268 }
1269 Ok(Step::Done(..)) => {
1270 return Err(test_error(
1271 "expected the offload error to propagate, got a response",
1272 ));
1273 }
1274 Err(err) => {
1275 assert!(err.to_string().contains("upstream model call failed"));
1277 saw_error = true;
1278 }
1279 }
1280 }
1281
1282 assert!(saw_error, "expected an error step");
1283 Ok(())
1284 }
1285
1286 #[tokio::test]
1287 async fn dropping_the_stream_cancels_the_algorithm_task() -> Result<()> {
1288 use std::sync::atomic::{AtomicBool, Ordering};
1289 use std::time::Duration;
1290 use tokio::sync::mpsc;
1291
1292 struct DropGuard(Arc<AtomicBool>);
1295 impl Drop for DropGuard {
1296 fn drop(&mut self) {
1297 self.0.store(true, Ordering::SeqCst);
1298 }
1299 }
1300
1301 struct StuckAlgo {
1302 started: mpsc::UnboundedSender<()>,
1303 dropped: Arc<AtomicBool>,
1304 }
1305
1306 #[async_trait]
1307 impl Algorithm for StuckAlgo {
1308 fn name(&self) -> &str {
1309 "stuck"
1310 }
1311
1312 async fn route(
1313 self: Arc<Self>,
1314 _driver: Driver,
1315 _request: Request,
1316 ) -> Result<RoutingOutcome> {
1317 let _guard = DropGuard(self.dropped.clone());
1318 let _ = self.started.send(());
1319 std::future::pending::<()>().await;
1321 unreachable!()
1322 }
1323 }
1324
1325 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1326 let dropped = Arc::new(AtomicBool::new(false));
1327 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1328 started: started_tx,
1329 dropped: dropped.clone(),
1330 });
1331
1332 let stream = algo.run_stream(request(), Arc::new(RuntimeModels::default()));
1333 started_rx
1334 .recv()
1335 .await
1336 .ok_or_else(|| test_error("task never started"))?;
1337 drop(stream);
1338 tokio::time::sleep(Duration::from_millis(100)).await;
1339
1340 assert!(
1341 dropped.load(Ordering::SeqCst),
1342 "algorithm task was NOT cancelled after dropping the stream"
1343 );
1344 Ok(())
1345 }
1346
1347 #[tokio::test]
1348 async fn route_panic_surfaces_as_a_stream_error() -> Result<()> {
1349 struct Panicky;
1352
1353 #[async_trait]
1354 impl Algorithm for Panicky {
1355 fn name(&self) -> &str {
1356 "panicky"
1357 }
1358
1359 async fn route(
1360 self: Arc<Self>,
1361 _driver: Driver,
1362 _request: Request,
1363 ) -> Result<RoutingOutcome> {
1364 panic!("boom");
1365 }
1366 }
1367
1368 let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
1369 let stream = algo.run_stream(request(), Arc::new(RuntimeModels::default()));
1370 tokio::pin!(stream);
1371
1372 let mut saw_error = false;
1373 while let Some(step) = stream.next().await {
1374 match step {
1375 Err(err) => {
1376 assert!(err.to_string().contains("algorithm task panicked: boom"));
1378 saw_error = true;
1379 }
1380 Ok(_) => return Err(test_error("expected the panic to surface as an error step")),
1381 }
1382 }
1383
1384 assert!(saw_error, "expected an error step from the panicked task");
1385 Ok(())
1386 }
1387
1388 #[tokio::test]
1392 async fn a_panic_with_a_leaked_driver_clone_still_terminates_the_run() -> Result<()> {
1393 struct LeakyPanic;
1394
1395 #[async_trait]
1396 impl Algorithm for LeakyPanic {
1397 fn name(&self) -> &str {
1398 "leaky_panic"
1399 }
1400
1401 async fn route(
1402 self: Arc<Self>,
1403 driver: Driver,
1404 _request: Request,
1405 ) -> Result<RoutingOutcome> {
1406 tokio::spawn(async move {
1407 let _keep_alive = driver;
1409 std::future::pending::<()>().await;
1410 });
1411 tokio::task::yield_now().await;
1412 panic!("boom");
1413 }
1414 }
1415
1416 let algo: Arc<dyn Algorithm> = Arc::new(LeakyPanic);
1417 let result = tokio::time::timeout(
1419 std::time::Duration::from_secs(1),
1420 test_drive(algo, request(), echo()),
1421 )
1422 .await
1423 .map_err(|error| LibsyError::external("waiting for the panicked run to end", error))?;
1424
1425 match result {
1426 Ok(_) => Err(test_error(
1427 "expected the panic to end the run with an error",
1428 )),
1429 Err(err) => {
1430 assert!(err.to_string().contains("algorithm task panicked: boom"));
1431 Ok(())
1432 }
1433 }
1434 }
1435
1436 #[tokio::test]
1437 async fn cancelling_run_cancels_the_algorithm_task() -> Result<()> {
1438 use std::sync::atomic::{AtomicBool, Ordering};
1439 use std::time::Duration;
1440 use tokio::sync::mpsc;
1441
1442 struct DropGuard(Arc<AtomicBool>);
1445 impl Drop for DropGuard {
1446 fn drop(&mut self) {
1447 self.0.store(true, Ordering::SeqCst);
1448 }
1449 }
1450
1451 struct StuckAlgo {
1452 started: mpsc::UnboundedSender<()>,
1453 dropped: Arc<AtomicBool>,
1454 }
1455
1456 #[async_trait]
1457 impl Algorithm for StuckAlgo {
1458 fn name(&self) -> &str {
1459 "stuck"
1460 }
1461
1462 async fn route(
1463 self: Arc<Self>,
1464 _driver: Driver,
1465 _request: Request,
1466 ) -> Result<RoutingOutcome> {
1467 let _guard = DropGuard(self.dropped.clone());
1468 let _ = self.started.send(());
1469 std::future::pending::<()>().await;
1472 unreachable!()
1473 }
1474 }
1475
1476 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1477 let dropped = Arc::new(AtomicBool::new(false));
1478 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1479 started: started_tx,
1480 dropped: dropped.clone(),
1481 });
1482
1483 let run_task = tokio::spawn(async move { test_drive(algo, request(), echo()).await });
1486 started_rx
1487 .recv()
1488 .await
1489 .ok_or_else(|| test_error("task never started"))?;
1490 run_task.abort();
1491 tokio::time::sleep(Duration::from_millis(100)).await;
1492
1493 assert!(
1494 dropped.load(Ordering::SeqCst),
1495 "algorithm task was NOT cancelled after cancelling run"
1496 );
1497 Ok(())
1498 }
1499
1500 struct Hedge {
1505 winner: String,
1506 loser: String,
1507 }
1508
1509 #[async_trait]
1510 impl Algorithm for Hedge {
1511 fn name(&self) -> &str {
1512 "hedge"
1513 }
1514
1515 async fn route(
1516 self: Arc<Self>,
1517 driver: Driver,
1518 request: Request,
1519 ) -> Result<RoutingOutcome> {
1520 let outcome_request = request.clone();
1521 let win = driver.call_model(request.clone(), vec![self.winner.clone().into()]);
1522 let lose = driver.call_model(request, vec![self.loser.clone().into()]);
1523 tokio::select! {
1525 res = win => Ok(RoutingOutcome::answered(
1526 self.winner.clone().into(),
1527 outcome_request,
1528 res?,
1529 )),
1530 res = lose => Ok(RoutingOutcome::answered(
1531 self.loser.clone().into(),
1532 outcome_request,
1533 res?,
1534 )),
1535 }
1536 }
1537 }
1538
1539 fn hedge(loser_delay: Option<std::time::Duration>) -> (Arc<dyn Algorithm>, impl Serve) {
1543 let started = Arc::new(tokio::sync::Notify::new());
1544 let algo = Arc::new(Hedge {
1545 winner: "winner".to_string(),
1546 loser: "loser".to_string(),
1547 });
1548 let serve = move |target: ModelId, _request: Request| {
1549 let started = started.clone();
1550 async move {
1551 if target == "loser" {
1552 started.notify_one();
1553 match loser_delay {
1554 Some(delay) => tokio::time::sleep(delay).await,
1555 None => std::future::pending::<()>().await,
1556 }
1557 } else {
1558 started.notified().await;
1559 }
1560 Ok(reply(target))
1561 }
1562 };
1563 (algo, serve)
1564 }
1565
1566 #[tokio::test]
1567 async fn run_returns_the_winner_without_a_late_loser_overwriting_it() -> Result<()> {
1568 let (algo, serve) = hedge(Some(std::time::Duration::from_millis(50)));
1571 let (_, response) = test_drive(algo, request(), serve).await?;
1572 assert_eq!(
1573 response
1574 .llm_response
1575 .as_agg()
1576 .map(completion_text)
1577 .unwrap_or_default(),
1578 "winner"
1579 );
1580 Ok(())
1581 }
1582
1583 #[tokio::test]
1584 async fn run_returns_the_winner_without_hanging_on_a_pending_loser() -> Result<()> {
1585 let (algo, serve) = hedge(None);
1588 let run = test_drive(algo, request(), serve);
1589 let (_, response) = tokio::time::timeout(std::time::Duration::from_secs(1), run)
1590 .await
1591 .map_err(|error| LibsyError::external("waiting for pending loser", error))??;
1592 assert_eq!(
1593 response
1594 .llm_response
1595 .as_agg()
1596 .map(completion_text)
1597 .unwrap_or_default(),
1598 "winner"
1599 );
1600 Ok(())
1601 }
1602
1603 #[tokio::test]
1604 async fn run_surfaces_a_terminal_error_with_many_calls_in_flight() -> Result<()> {
1605 use std::sync::atomic::{AtomicUsize, Ordering};
1606
1607 const N: usize = 10;
1610
1611 struct FanOutThenError {
1614 all_started: Arc<tokio::sync::Notify>,
1615 n: usize,
1616 }
1617
1618 #[async_trait]
1619 impl Algorithm for FanOutThenError {
1620 fn name(&self) -> &str {
1621 "fan_out_then_error"
1622 }
1623
1624 async fn route(
1625 self: Arc<Self>,
1626 driver: Driver,
1627 request: Request,
1628 ) -> Result<RoutingOutcome> {
1629 let offloads = futures::future::join_all(
1630 (0..self.n)
1631 .map(|i| driver.call_model(request.clone(), vec![format!("m{i}").into()])),
1632 );
1633 tokio::select! {
1634 _ = offloads => Err(test_error("offloads unexpectedly completed")),
1635 _ = self.all_started.notified() => {
1636 Err(test_error("terminal error while calls pending"))
1637 }
1638 }
1639 }
1640 }
1641
1642 let all_started = Arc::new(tokio::sync::Notify::new());
1643 let algo: Arc<dyn Algorithm> = Arc::new(FanOutThenError {
1644 all_started: all_started.clone(),
1645 n: N,
1646 });
1647
1648 let started = Arc::new(AtomicUsize::new(0));
1650 let serve = move |_target: ModelId, _request: Request| {
1651 let started = started.clone();
1652 let all_started = all_started.clone();
1653 async move {
1654 if started.fetch_add(1, Ordering::SeqCst) + 1 == N {
1655 all_started.notify_one();
1656 }
1657 std::future::pending::<ServeResult>().await
1658 }
1659 };
1660
1661 let run = test_drive(algo, request(), serve);
1664 let result = tokio::time::timeout(std::time::Duration::from_millis(500), run)
1665 .await
1666 .map_err(|error| {
1667 LibsyError::external("waiting for terminal error with full call cap", error)
1668 })?;
1669 match result {
1670 Ok(_) => Err(test_error("expected the terminal error, got a response")),
1671 Err(err) => {
1672 assert!(
1673 err.to_string()
1674 .contains("terminal error while calls pending")
1675 );
1676 Ok(())
1677 }
1678 }
1679 }
1680}