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}
120
121impl CallModel {
122 pub fn respond(mut self, result: Result<Response>) -> Result<()> {
126 self.record(result.is_ok());
127 self.reply
128 .take()
129 .ok_or(DriverError::ResponseDropped)?
130 .send(result)
131 .map_err(|_| DriverError::ResponseDropped.into())
132 }
133
134 pub fn fail(mut self, error: LibsyError) -> Result<()> {
137 self.reply = None;
138 self.record(false);
139 Err(error)
140 }
141
142 fn record(&self, is_ok: bool) {
143 observability::record_llm_call(
144 &self.algorithm,
145 self.models
146 .first()
147 .map(ModelId::as_str)
148 .unwrap_or("NoTargets"),
149 self.started.elapsed(),
150 is_ok,
151 );
152 }
153}
154
155impl Drop for CallModel {
156 fn drop(&mut self) {
157 if self.reply.is_some() {
158 self.record(false);
159 }
160 }
161}
162
163pub struct CallDecision {
165 pub algorithm: String,
167 pub request: DecisionRequest,
169 pub model: ModelId,
171 reply: oneshot::Sender<Result<DecisionResponse>>,
172}
173
174impl CallDecision {
175 pub fn respond(self, result: Result<DecisionResponse>) -> Result<()> {
177 self.reply
178 .send(result)
179 .map_err(|_| DriverError::ResponseDropped.into())
180 }
181
182 pub fn fail(self, error: LibsyError) -> Result<()> {
184 Err(error)
185 }
186}
187
188pub struct RoutingOutcome {
190 pub selected_model_ids: Vec<ModelId>,
192 pub request: Request,
194 pub response: Option<Response>,
196 pub metadata: Option<crate::OutcomeMetadata>,
201}
202
203impl RoutingOutcome {
204 pub fn selected_model_id(&self) -> Result<&ModelId> {
207 self.selected_model_ids.first().ok_or(LibsyError::NoTargets)
208 }
209
210 pub fn route_to(
214 selected_model_id: ModelId,
215 fallback_models: Vec<ModelId>,
216 mut request: Request,
217 ) -> Self {
218 request.llm_request.model = Some(selected_model_id.to_string());
219 let mut selected_model_ids = Vec::with_capacity(1 + fallback_models.len());
220 selected_model_ids.push(selected_model_id);
221 selected_model_ids.extend(fallback_models);
222 Self {
223 selected_model_ids,
224 request,
225 response: None,
226 metadata: None,
227 }
228 }
229
230 pub fn answered(selected_model_id: ModelId, mut request: Request, response: Response) -> Self {
233 request.llm_request.model = Some(selected_model_id.to_string());
234 Self {
235 selected_model_ids: vec![selected_model_id],
236 request,
237 response: Some(response),
238 metadata: None,
239 }
240 }
241}
242
243#[derive(Clone)]
245pub struct Driver {
246 step_tx: mpsc::Sender<Result<Step>>,
247
248 algorithm: String,
250
251 evidence: Arc<Mutex<Option<Value>>>,
253
254 models: Arc<RuntimeModels>,
256
257 scope: Scope,
259}
260
261impl Driver {
262 pub(crate) fn new(
265 algorithm: &str,
266 models: Arc<RuntimeModels>,
267 ) -> (Self, mpsc::Receiver<Result<Step>>) {
268 let (step_tx, step_rx) = mpsc::channel(1);
273 (
274 Self {
275 step_tx,
276 algorithm: algorithm.to_string(),
277 evidence: Arc::new(Mutex::new(None)),
278 models,
279 scope: Scope::Parent,
280 },
281 step_rx,
282 )
283 }
284
285 pub(crate) fn set_evidence(&self, evidence: Value) {
287 *self.evidence.lock() = Some(evidence);
288 }
289
290 pub(crate) fn set_evidence_if_empty(&self, evidence: Value) {
292 let mut current = self.evidence.lock();
293 if current.is_none() {
294 *current = Some(evidence);
295 }
296 }
297
298 pub async fn call_model(&self, request: Request, models: Vec<ModelId>) -> Result<Response> {
308 self.call_model_with_error_recovery(request, models, false)
309 .await
310 }
311
312 #[tracing::instrument(
314 target = "libsy",
315 name = "libsy.llm_call",
316 skip_all,
317 fields(
318 algorithm = self.algorithm,
319 selected_model = %models.first().map(ModelId::as_str).unwrap_or("NoTargets"),
320 openinference.span.kind = "CHAIN",
321 outcome = tracing::field::Empty,
322 input_tokens = tracing::field::Empty,
323 output_tokens = tracing::field::Empty,
324 total_tokens = tracing::field::Empty,
325 reasoning_tokens = tracing::field::Empty,
326 )
327 )]
328 pub(crate) async fn call_model_with_error_recovery(
329 &self,
330 mut request: Request,
331 models: Vec<ModelId>,
332 recover_errors: bool,
333 ) -> Result<Response> {
334 let Some(selected_model_id) = models.first() else {
335 return Err(LibsyError::NoTargets);
336 };
337 request.llm_request.model = Some(selected_model_id.to_string());
338 let started = Instant::now();
339 let (reply, response) = oneshot::channel::<Result<Response>>();
340 let call = CallModel {
341 algorithm: self.algorithm.clone(),
342 request,
343 models,
344 recover_errors,
345 reply: Some(reply),
346 started,
347 };
348 let result = async {
349 self.step_tx
350 .send(Ok(Step::CallModel(Box::new(call))))
351 .await
352 .map_err(|_| DriverError::StreamClosed)?;
353 response
354 .await
355 .map_err(|_| LibsyError::from(DriverError::ResponseDropped))?
356 }
357 .await;
358 observability::record_llm_call_span(&result, &tracing::Span::current());
359 result
360 }
361
362 #[tracing::instrument(
364 target = "libsy",
365 name = "libsy.decision_call",
366 skip_all,
367 fields(algorithm = self.algorithm, selected_model = %model),
368 )]
369 pub async fn call_decision(
370 &self,
371 mut request: DecisionRequest,
372 model: ModelId,
373 ) -> Result<DecisionResponse> {
374 request.model = Some(model.clone());
375 let (reply, response) = oneshot::channel();
376 let call = CallDecision {
377 algorithm: self.algorithm.clone(),
378 request,
379 model,
380 reply,
381 };
382 self.step_tx
383 .send(Ok(Step::CallDecision(Box::new(call))))
384 .await
385 .map_err(|_| DriverError::StreamClosed)?;
386 response
387 .await
388 .map_err(|_| LibsyError::from(DriverError::ResponseDropped))?
389 }
390
391 pub fn models_for(&self, category: &Category) -> &[ModelId] {
393 match self.scope {
394 Scope::Parent => self.models.models_for(category),
395 Scope::Subagent => self.models.subagent_models_for(category),
396 }
397 }
398
399 pub fn first_model_for(&self, category: &Category) -> Result<&ModelId> {
401 self.models_for(category)
402 .first()
403 .ok_or_else(|| LibsyError::AlgorithmError {
404 message: format!("no models available for category {}", category.as_str()),
405 })
406 }
407
408 pub fn for_subagent(&self) -> Result<Self> {
411 if self.models.subagent.is_none() {
412 return Err(LibsyError::AlgorithmError {
413 message: "delegated work has no sub-agent models".to_string(),
414 });
415 }
416 Ok(Self {
417 scope: Scope::Subagent,
418 ..self.clone()
419 })
420 }
421
422 pub(crate) async fn finish(&self, result: Result<RoutingOutcome>) -> Result<()> {
426 let result = result.map(|mut outcome| {
427 let metadata = outcome.metadata.get_or_insert_with(|| {
428 crate::OutcomeMetadata::new(self.algorithm.clone(), self.evidence.lock().take())
429 });
430 observability::record_outcome(metadata, &outcome.selected_model_ids);
431 outcome
432 });
433 let selected_model = result
434 .as_ref()
435 .ok()
436 .and_then(|outcome| outcome.selected_model_id().ok().cloned());
437 let step = result.map(|outcome| Step::Done(Box::new(outcome)));
438 self.step_tx
439 .send(step)
440 .await
441 .map_err(|_| DriverError::StreamClosed)?;
442 if let Some(selected_model) = selected_model {
443 observability::record_decision(&self.algorithm, &selected_model);
444 }
445 Ok(())
446 }
447}
448
449pub enum Step {
452 CallModel(Box<CallModel>),
455 CallDecision(Box<CallDecision>),
457 Done(Box<RoutingOutcome>),
459}
460
461pub enum Call {
464 Model(Box<CallModel>),
466 Decision(Box<CallDecision>),
468}
469
470pub async fn drive<F, Fut>(
480 algorithm: Arc<dyn Algorithm>,
481 request: Request,
482 models: Arc<RuntimeModels>,
483 serve: F,
484) -> Result<RoutingOutcome>
485where
486 F: Fn(Call) -> Fut,
487 Fut: Future<Output = Result<()>>,
488{
489 let stream = algorithm.run_stream(request, models);
490 tokio::pin!(stream);
491
492 let mut in_flight = futures::stream::FuturesUnordered::new();
493 let mut final_outcome: Option<RoutingOutcome> = None;
494
495 loop {
496 tokio::select! {
497 Some(result) = in_flight.next() => match result {
498 Ok(()) => {}, Err(err) => return Err(err), },
501 step = stream.next() => {
502 match step {
503 None => break, Some(item) => match item? {
505 Step::CallModel(call) => in_flight.push(serve(Call::Model(call))),
506 Step::CallDecision(call) => {
507 in_flight.push(serve(Call::Decision(call)));
508 }
509 Step::Done(outcome) => {
510 final_outcome = Some(*outcome);
511 break;
512 }
513 }
514 }
515 },
516 }
517 }
518 final_outcome.ok_or(LibsyError::MissingFinalResponse)
519}
520
521fn panic_message(payload: &(dyn std::any::Any + Send)) -> String {
523 payload
524 .downcast_ref::<&'static str>()
525 .map(|message| (*message).to_string())
526 .or_else(|| payload.downcast_ref::<String>().cloned())
527 .unwrap_or_else(|| "unknown panic payload".to_string())
528}
529
530struct AbortOnDrop(tokio::task::AbortHandle);
532
533impl Drop for AbortOnDrop {
534 fn drop(&mut self) {
535 self.0.abort();
536 }
537}
538
539pub(crate) fn ensure_model_is_target(targets: &[ModelId], name: &ModelId) -> Result<()> {
544 targets
545 .iter()
546 .any(|target| target == name)
547 .then_some(())
548 .ok_or_else(|| LibsyError::TargetNotFound {
549 target: name.clone(),
550 })
551}
552
553#[derive(Clone, Hash, PartialEq, Eq)]
556pub(crate) enum RoutingIdentity {
557 Session(String),
559 Subagent { session: String, agent: String },
561}
562
563impl RoutingIdentity {
564 pub(crate) fn from_request(request: &Request) -> Option<Self> {
569 let metadata = request.metadata.as_ref()?;
570 let session = metadata.session_id.as_deref().filter(|id| !id.is_empty())?;
571 if metadata.is_subagent {
572 let agent = metadata.agent_id.as_deref().filter(|id| !id.is_empty())?;
573 Some(Self::Subagent {
574 session: session.to_string(),
575 agent: agent.to_string(),
576 })
577 } else {
578 Some(Self::Session(session.to_string()))
579 }
580 }
581}
582
583#[async_trait]
616pub trait Algorithm: Send + Sync + 'static {
617 fn name(&self) -> &str;
621
622 async fn route(self: Arc<Self>, driver: Driver, request: Request) -> Result<RoutingOutcome>;
625
626 fn run_stream(self: Arc<Self>, request: Request, models: Arc<RuntimeModels>) -> StepStream {
634 let (driver, step_rx) = Driver::new(self.name(), models);
635 let span = observability::run_span(self.name(), &request);
636 let handle = tokio::spawn(
637 async move {
638 let algorithm = self.name().to_string();
639 let route = AssertUnwindSafe(self.route(driver.clone(), request)).catch_unwind();
641 let result = observability::observe_run(&algorithm, async move {
642 route.await.unwrap_or_else(|payload| {
643 Err(LibsyError::AlgorithmError {
644 message: format!(
645 "algorithm task panicked: {}",
646 panic_message(payload.as_ref())
647 ),
648 })
649 })
650 })
651 .await;
652
653 let _ = driver.finish(result).await;
654 }
655 .instrument(span),
656 );
657 let abort_guard = AbortOnDrop(handle.abort_handle());
659 Box::pin(ReceiverStream::new(step_rx).map(move |step| {
660 let _keep_alive = &abort_guard;
662 step
663 }))
664 }
665}
666
667#[cfg(test)]
668mod tests {
669 use std::collections::HashMap;
670
671 use super::*;
672 use crate::core::testing::{Serve, ServeResult, echo, reply, serve_decision, test_drive};
673 use futures::StreamExt;
674 use switchyard_protocol::{
675 LlmResponse, LlmResponseChunk, completion_text, text_request, text_response,
676 };
677
678 #[derive(Debug, thiserror::Error)]
679 #[error("{0}")]
680 struct TestError(&'static str);
681
682 fn test_error(message: &'static str) -> LibsyError {
683 LibsyError::external("test", TestError(message))
684 }
685
686 struct TestAlgo {
689 target_set: Vec<ModelId>,
690 }
691
692 #[async_trait]
693 impl Algorithm for TestAlgo {
694 fn name(&self) -> &str {
695 "test"
696 }
697
698 async fn route(
699 self: Arc<Self>,
700 driver: Driver,
701 request: Request,
702 ) -> Result<RoutingOutcome> {
703 let target = self
704 .target_set
705 .first()
706 .ok_or(LibsyError::NoTargets)?
707 .clone();
708 let response = driver
709 .call_model(request.clone(), vec![target.clone()])
710 .await?;
711 driver.set_evidence(serde_json::json!({"source": "test"}));
712 driver.set_evidence_if_empty(serde_json::json!({"source": "ignored"}));
713 Ok(RoutingOutcome::answered(target, request, response))
714 }
715 }
716
717 fn orch(target_set: Vec<ModelId>) -> Arc<dyn Algorithm> {
719 Arc::new(TestAlgo { target_set })
720 }
721
722 fn request() -> Request {
723 Request {
724 llm_request: text_request(Some("auto".to_string()), "hi".to_string()),
725 raw_request: None,
726 metadata: None,
727 }
728 }
729
730 #[test]
731 fn routing_outcome_constructors_stamp_selection_and_preserve_payloads() {
732 let outcome = RoutingOutcome::route_to(
733 "selected".into(),
734 target_set(&["fallback-one", "fallback-two"]),
735 request(),
736 );
737
738 assert_eq!(
739 outcome.selected_model_ids,
740 target_set(&["selected", "fallback-one", "fallback-two"])
741 );
742 assert_eq!(outcome.request.model_id().as_deref(), Some("selected"));
743 assert!(outcome.response.is_none());
744 assert!(outcome.metadata.is_none());
745
746 let outcome = RoutingOutcome::route_to("only".into(), Vec::new(), request());
747 assert_eq!(outcome.selected_model_ids, target_set(&["only"]));
748
749 let outcome = RoutingOutcome::answered(
750 "answered".into(),
751 request(),
752 Response {
753 llm_response: LlmResponse::Agg(text_response(None, "existing")),
754 metadata: None,
755 upstream_headers: http::HeaderMap::new(),
756 },
757 );
758
759 assert_eq!(outcome.selected_model_ids, target_set(&["answered"]));
760 assert_eq!(outcome.request.model_id().as_deref(), Some("answered"));
761 assert_eq!(
762 outcome
763 .response
764 .as_ref()
765 .and_then(|response| response.llm_response.as_agg())
766 .map(completion_text),
767 Some("existing".to_string())
768 );
769 }
770
771 fn target_set(names: &[&str]) -> Vec<ModelId> {
772 names.iter().map(|name| ModelId::from(*name)).collect()
773 }
774
775 #[tokio::test]
776 async fn typed_driver_preserves_call_and_stream_boundaries() -> Result<()> {
777 tokio::time::timeout(std::time::Duration::from_secs(1), async {
778 let (driver, mut step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
781 let first_driver = driver.clone();
782 let mut first = tokio::spawn(async move {
783 first_driver
784 .call_model(request(), vec![ModelId::from("first")])
785 .await
786 });
787 let second = tokio::spawn(async move {
788 driver
789 .call_model(request(), vec![ModelId::from("second")])
790 .await
791 });
792
793 let mut calls = HashMap::new();
794 for _ in 0..2 {
795 let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
796 let Step::CallModel(call) = step else {
797 return Err(test_error("expected a CallModel step"));
798 };
799 let selected_model = call
800 .models
801 .first()
802 .ok_or_else(|| test_error("model call has no candidates"))?
803 .to_string();
804 calls.insert(selected_model, call);
805 }
806 assert!(
807 tokio::time::timeout(std::time::Duration::from_millis(20), &mut first)
808 .await
809 .is_err(),
810 "call completed before the host responded"
811 );
812 calls
813 .remove("second")
814 .ok_or_else(|| test_error("missing second call"))?
815 .respond(Ok(reply("second response")))?;
816 calls
817 .remove("first")
818 .ok_or_else(|| test_error("missing first call"))?
819 .respond(Ok(reply("first response")))?;
820
821 let first_response = first
822 .await
823 .map_err(|source| LibsyError::external("joining a test task", source))??;
824 let second_response = second
825 .await
826 .map_err(|source| LibsyError::external("joining a test task", source))??;
827 assert_eq!(
828 first_response.llm_response.as_agg().map(completion_text),
829 Some("first response".to_string())
830 );
831 assert_eq!(
832 second_response.llm_response.as_agg().map(completion_text),
833 Some("second response".to_string())
834 );
835
836 let (driver, mut step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
838 let producer = tokio::spawn(async move {
839 driver
840 .call_model(request(), vec![ModelId::from("dropped")])
841 .await
842 });
843 let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
844 let Step::CallModel(call) = step else {
845 return Err(test_error("expected a CallModel step"));
846 };
847 drop(call);
848 let result = producer
849 .await
850 .map_err(|source| LibsyError::external("joining a test task", source))?;
851 assert!(matches!(
852 result,
853 Err(LibsyError::Driver(DriverError::ResponseDropped))
854 ));
855
856 let (driver, step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
858 drop(step_rx);
859 let result = driver
860 .call_model(request(), vec![ModelId::from("closed")])
861 .await;
862 assert!(matches!(
863 result,
864 Err(LibsyError::Driver(DriverError::StreamClosed))
865 ));
866
867 fn decision_request() -> DecisionRequest {
868 DecisionRequest {
869 model: Some("overwritten".into()),
870 context: serde_json::json!({"task": "choose a route"}),
871 questions: Default::default(),
872 }
873 }
874
875 fn decision_response() -> DecisionResponse {
876 DecisionResponse {
877 id: Some("decision-1".to_string()),
878 model: Some("provider-model".into()),
879 answers: [(
880 "p_solve".to_string(),
881 switchyard_protocol::DecisionAnswer {
882 value: switchyard_protocol::DecisionValue::Boolean(
883 switchyard_protocol::BooleanEstimate::ProbabilityTrue(
884 switchyard_protocol::Probability(0.8),
885 ),
886 ),
887 provider_confidence: None,
888 },
889 )]
890 .into(),
891 usage: Default::default(),
892 }
893 }
894
895 struct MixedCalls(&'static str);
896
897 #[async_trait]
898 impl Algorithm for MixedCalls {
899 fn name(&self) -> &str {
900 "mixed"
901 }
902
903 async fn route(
904 self: Arc<Self>,
905 driver: Driver,
906 request: Request,
907 ) -> Result<RoutingOutcome> {
908 let (llm, decision) = tokio::join!(
909 driver.call_model(request.clone(), vec!["llm".into()]),
910 driver.call_decision(decision_request(), "decision".into()),
911 );
912 assert_eq!(
913 llm?.llm_response.as_agg().map(completion_text),
914 Some("llm reply".into())
915 );
916 match self.0 {
917 "mock" => assert_eq!(decision?, DecisionResponse {
918 id: None,
919 model: Some("decision".into()),
920 answers: Default::default(),
921 usage: Default::default(),
922 }),
923 "reply" => assert_eq!(decision?, decision_response()),
924 "error" => assert!(matches!(
925 decision,
926 Err(LibsyError::AlgorithmError { message }) if message == "provider failed"
927 )),
928 "drop" => assert!(matches!(
929 decision,
930 Err(LibsyError::Driver(DriverError::ResponseDropped))
931 )),
932 "abort" => return std::future::pending().await,
933 _ => unreachable!(),
934 }
935 Ok(RoutingOutcome::route_to("answer".into(), vec![], request))
936 }
937 }
938
939 for mode in ["mock", "reply", "error", "drop", "abort"] {
940 let barrier = Arc::new(tokio::sync::Barrier::new(2));
942 let outcome = drive(
943 Arc::new(MixedCalls(mode)),
944 request(),
945 Arc::new(RuntimeModels::default()),
946 move |call| {
947 let barrier = barrier.clone();
948 async move {
949 barrier.wait().await;
950 let call = match call {
951 Call::Model(call) => return call.respond(Ok(reply("llm reply"))),
952 Call::Decision(call) => *call,
953 };
954 assert_eq!(call.algorithm, "mixed");
955 assert_eq!(call.model, "decision");
956 assert_eq!(call.request.model, Some("decision".into()));
957 assert_eq!(call.request.context, decision_request().context);
958 match mode {
959 "mock" => serve_decision(call).await,
960 "reply" => call.respond(Ok(decision_response())),
961 "error" => call.respond(Err(LibsyError::AlgorithmError {
962 message: "provider failed".into(),
963 })),
964 "drop" => {
965 drop(call);
966 Ok(())
967 }
968 "abort" => call.fail(test_error("host aborted")),
969 _ => unreachable!(),
970 }
971 }
972 },
973 )
974 .await;
975 if mode == "abort" {
976 assert!(matches!(outcome, Err(LibsyError::External { .. })));
977 } else {
978 assert_eq!(outcome?.selected_model_id()?, "answer");
979 }
980 }
981
982 let (driver, mut steps) = Driver::new("test", Arc::new(RuntimeModels::default()));
983 let mut pending = Box::pin(driver.call_decision(decision_request(), "decision".into()));
984 assert!(futures::poll!(&mut pending).is_pending());
985 let Some(Ok(Step::CallDecision(call))) = steps.recv().await else {
986 return Err(test_error("expected a decision call"));
987 };
988 drop(pending);
989 assert!(matches!(
990 call.respond(Ok(decision_response())),
991 Err(LibsyError::Driver(DriverError::ResponseDropped))
992 ));
993 drop(steps);
994 assert!(matches!(
995 driver.call_decision(decision_request(), "decision".into()).await,
996 Err(LibsyError::Driver(DriverError::StreamClosed))
997 ));
998 Ok(())
999 })
1000 .await
1001 .map_err(|error| LibsyError::external("waiting for typed driver boundaries", error))?
1002 }
1003
1004 #[test]
1005 fn target_lookup_returns_the_missing_target() {
1006 let error = ensure_model_is_target(&target_set(&[]), &ModelId::from("missing")).err();
1007 assert!(matches!(
1008 error,
1009 Some(LibsyError::TargetNotFound { target }) if target == "missing"
1010 ));
1011 }
1012
1013 fn streaming_orch(chunks: Vec<LlmResponseChunk>) -> (Arc<dyn Algorithm>, impl Serve) {
1016 let algo = orch(target_set(&["stream/model"]));
1017 let serve = move |_target: ModelId, _request: Request| {
1018 let chunks = chunks.clone();
1019 async move {
1020 let stream =
1021 futures::stream::iter(chunks.into_iter().map(|chunk| Ok(chunk.into()))).boxed();
1022 Ok(Response {
1023 llm_response: LlmResponse::Stream(stream),
1024 metadata: None,
1025 upstream_headers: http::HeaderMap::new(),
1026 })
1027 }
1028 };
1029 (algo, serve)
1030 }
1031
1032 #[tokio::test]
1033 async fn run_returns_a_streamed_response_the_caller_aggregates() -> Result<()> {
1034 let (orch, serve) = streaming_orch(vec![
1037 LlmResponseChunk::MessageStart {
1038 id: Some("m1".to_string()),
1039 model: Some("stream/model".to_string()),
1040 },
1041 LlmResponseChunk::TextDelta {
1042 index: 0,
1043 text: "hel".to_string(),
1044 },
1045 LlmResponseChunk::TextDelta {
1046 index: 0,
1047 text: "lo".to_string(),
1048 },
1049 LlmResponseChunk::MessageStop {
1050 reason: Some("stop".to_string()),
1051 },
1052 ]);
1053 let (selected_model, response) = test_drive(orch, request(), serve).await?;
1054 let agg = response
1056 .llm_response
1057 .into_agg()
1058 .await
1059 .map_err(|error| LibsyError::external("aggregating response stream", error))?;
1060 assert_eq!(completion_text(&agg), "hello");
1061 assert_eq!(agg.model.as_deref(), Some("stream/model"));
1062 assert_eq!(selected_model, "stream/model");
1063 Ok(())
1064 }
1065
1066 #[tokio::test]
1067 async fn aggregating_a_streamed_response_propagates_a_mid_stream_error() -> Result<()> {
1068 let (orch, serve) = streaming_orch(vec![
1071 LlmResponseChunk::TextDelta {
1072 index: 0,
1073 text: "partial".to_string(),
1074 },
1075 LlmResponseChunk::StreamError {
1076 message: "upstream exploded".to_string(),
1077 },
1078 ]);
1079 let (_, response) = test_drive(orch, request(), serve).await?;
1080 match response.llm_response.into_agg().await {
1081 Ok(_) => panic!("expected a mid-stream error, got an aggregate"),
1082 Err(err) => {
1083 assert!(err.to_string().contains("upstream exploded"));
1084 Ok(())
1085 }
1086 }
1087 }
1088
1089 #[tokio::test]
1090 async fn run_offloads_via_promise_then_finishes() -> Result<()> {
1091 let stream = orch(target_set(&["offload/model"]))
1094 .run_stream(request(), Arc::new(RuntimeModels::default()));
1095 tokio::pin!(stream);
1096
1097 let mut saw_call = false;
1098 let mut final_completion = None;
1099 while let Some(step) = stream.next().await {
1100 match step? {
1101 Step::CallDecision(_) => return Err(test_error("unexpected decision call")),
1102 Step::CallModel(call) => {
1103 saw_call = true;
1104 assert_eq!(call.models, vec![ModelId::from("offload/model")]);
1105 call.respond(Ok(Response {
1107 llm_response: LlmResponse::Agg(text_response(
1108 None,
1109 "fulfilled".to_string(),
1110 )),
1111 metadata: None,
1112 upstream_headers: http::HeaderMap::new(),
1113 }))?;
1114 }
1115 Step::Done(outcome) => {
1116 let metadata = outcome
1117 .metadata
1118 .as_ref()
1119 .expect("run_stream should attach outcome metadata");
1120 assert_eq!(metadata.algorithm, "test");
1121 assert_eq!(
1122 uuid::Uuid::parse_str(metadata.outcome_id())
1123 .expect("outcome id should be a UUID")
1124 .get_version_num(),
1125 7
1126 );
1127 assert_eq!(
1128 metadata.evidence,
1129 Some(serde_json::json!({"source": "test"}))
1130 );
1131 let response = outcome
1132 .response
1133 .ok_or_else(|| test_error("expected an answered outcome"))?;
1134 final_completion = Some(
1135 response
1136 .llm_response
1137 .as_agg()
1138 .map(completion_text)
1139 .unwrap_or_default(),
1140 );
1141 }
1142 }
1143 }
1144
1145 assert!(saw_call, "expected a CallModel step before Done");
1146 assert_eq!(
1147 final_completion.ok_or_else(|| test_error("no Done step"))?,
1148 "fulfilled"
1149 );
1150 Ok(())
1151 }
1152
1153 #[tokio::test(flavor = "multi_thread", worker_threads = 12)]
1154 async fn requests_are_processed_in_parallel() -> Result<()> {
1155 use std::time::Duration;
1156 use tokio::sync::Barrier;
1157
1158 const N: usize = 12;
1159
1160 let barrier = Arc::new(Barrier::new(N));
1165 let algo = orch(target_set(&["m"]));
1167
1168 let mut handles = Vec::new();
1169 for _ in 0..N {
1170 let algo = algo.clone();
1171 let barrier = barrier.clone();
1172 let serve = move |target: ModelId, _request: Request| {
1173 let barrier = barrier.clone();
1174 async move {
1175 barrier.wait().await;
1176 Ok(reply(target))
1177 }
1178 };
1179 handles.push(tokio::spawn(async move {
1180 test_drive(algo, request(), serve)
1181 .await
1182 .map(|(_, response)| {
1183 response
1184 .llm_response
1185 .as_agg()
1186 .map(completion_text)
1187 .unwrap_or_default()
1188 })
1189 }));
1190 }
1191
1192 for handle in handles {
1193 let completion = tokio::time::timeout(Duration::from_secs(5), handle)
1195 .await
1196 .map_err(|error| LibsyError::external("waiting for test task", error))?
1197 .map_err(|source| LibsyError::external("joining a test task", source))??;
1198 assert_eq!(completion, "m");
1199 }
1200 Ok(())
1201 }
1202
1203 #[tokio::test]
1204 async fn offload_error_propagates_back_to_the_algorithm() -> Result<()> {
1205 let stream = orch(target_set(&["offload/model"]))
1209 .run_stream(request(), Arc::new(RuntimeModels::default()));
1210 tokio::pin!(stream);
1211
1212 let mut saw_error = false;
1213 while let Some(step) = stream.next().await {
1214 match step {
1215 Ok(Step::CallDecision(_)) => return Err(test_error("unexpected decision call")),
1216 Ok(Step::CallModel(call)) => {
1217 call.respond(Err(test_error("upstream model call failed")))?;
1218 }
1219 Ok(Step::Done(..)) => {
1220 return Err(test_error(
1221 "expected the offload error to propagate, got a response",
1222 ));
1223 }
1224 Err(err) => {
1225 assert!(err.to_string().contains("upstream model call failed"));
1227 saw_error = true;
1228 }
1229 }
1230 }
1231
1232 assert!(saw_error, "expected an error step");
1233 Ok(())
1234 }
1235
1236 #[tokio::test]
1237 async fn dropping_the_stream_cancels_the_algorithm_task() -> Result<()> {
1238 use std::sync::atomic::{AtomicBool, Ordering};
1239 use std::time::Duration;
1240 use tokio::sync::mpsc;
1241
1242 struct DropGuard(Arc<AtomicBool>);
1245 impl Drop for DropGuard {
1246 fn drop(&mut self) {
1247 self.0.store(true, Ordering::SeqCst);
1248 }
1249 }
1250
1251 struct StuckAlgo {
1252 started: mpsc::UnboundedSender<()>,
1253 dropped: Arc<AtomicBool>,
1254 }
1255
1256 #[async_trait]
1257 impl Algorithm for StuckAlgo {
1258 fn name(&self) -> &str {
1259 "stuck"
1260 }
1261
1262 async fn route(
1263 self: Arc<Self>,
1264 _driver: Driver,
1265 _request: Request,
1266 ) -> Result<RoutingOutcome> {
1267 let _guard = DropGuard(self.dropped.clone());
1268 let _ = self.started.send(());
1269 std::future::pending::<()>().await;
1271 unreachable!()
1272 }
1273 }
1274
1275 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1276 let dropped = Arc::new(AtomicBool::new(false));
1277 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1278 started: started_tx,
1279 dropped: dropped.clone(),
1280 });
1281
1282 let stream = algo.run_stream(request(), Arc::new(RuntimeModels::default()));
1283 started_rx
1284 .recv()
1285 .await
1286 .ok_or_else(|| test_error("task never started"))?;
1287 drop(stream);
1288 tokio::time::sleep(Duration::from_millis(100)).await;
1289
1290 assert!(
1291 dropped.load(Ordering::SeqCst),
1292 "algorithm task was NOT cancelled after dropping the stream"
1293 );
1294 Ok(())
1295 }
1296
1297 #[tokio::test]
1298 async fn route_panic_surfaces_as_a_stream_error() -> Result<()> {
1299 struct Panicky;
1302
1303 #[async_trait]
1304 impl Algorithm for Panicky {
1305 fn name(&self) -> &str {
1306 "panicky"
1307 }
1308
1309 async fn route(
1310 self: Arc<Self>,
1311 _driver: Driver,
1312 _request: Request,
1313 ) -> Result<RoutingOutcome> {
1314 panic!("boom");
1315 }
1316 }
1317
1318 let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
1319 let stream = algo.run_stream(request(), Arc::new(RuntimeModels::default()));
1320 tokio::pin!(stream);
1321
1322 let mut saw_error = false;
1323 while let Some(step) = stream.next().await {
1324 match step {
1325 Err(err) => {
1326 assert!(err.to_string().contains("algorithm task panicked: boom"));
1328 saw_error = true;
1329 }
1330 Ok(_) => return Err(test_error("expected the panic to surface as an error step")),
1331 }
1332 }
1333
1334 assert!(saw_error, "expected an error step from the panicked task");
1335 Ok(())
1336 }
1337
1338 #[tokio::test]
1342 async fn a_panic_with_a_leaked_driver_clone_still_terminates_the_run() -> Result<()> {
1343 struct LeakyPanic;
1344
1345 #[async_trait]
1346 impl Algorithm for LeakyPanic {
1347 fn name(&self) -> &str {
1348 "leaky_panic"
1349 }
1350
1351 async fn route(
1352 self: Arc<Self>,
1353 driver: Driver,
1354 _request: Request,
1355 ) -> Result<RoutingOutcome> {
1356 tokio::spawn(async move {
1357 let _keep_alive = driver;
1359 std::future::pending::<()>().await;
1360 });
1361 tokio::task::yield_now().await;
1362 panic!("boom");
1363 }
1364 }
1365
1366 let algo: Arc<dyn Algorithm> = Arc::new(LeakyPanic);
1367 let result = tokio::time::timeout(
1369 std::time::Duration::from_secs(1),
1370 test_drive(algo, request(), echo()),
1371 )
1372 .await
1373 .map_err(|error| LibsyError::external("waiting for the panicked run to end", error))?;
1374
1375 match result {
1376 Ok(_) => Err(test_error(
1377 "expected the panic to end the run with an error",
1378 )),
1379 Err(err) => {
1380 assert!(err.to_string().contains("algorithm task panicked: boom"));
1381 Ok(())
1382 }
1383 }
1384 }
1385
1386 #[tokio::test]
1387 async fn cancelling_run_cancels_the_algorithm_task() -> Result<()> {
1388 use std::sync::atomic::{AtomicBool, Ordering};
1389 use std::time::Duration;
1390 use tokio::sync::mpsc;
1391
1392 struct DropGuard(Arc<AtomicBool>);
1395 impl Drop for DropGuard {
1396 fn drop(&mut self) {
1397 self.0.store(true, Ordering::SeqCst);
1398 }
1399 }
1400
1401 struct StuckAlgo {
1402 started: mpsc::UnboundedSender<()>,
1403 dropped: Arc<AtomicBool>,
1404 }
1405
1406 #[async_trait]
1407 impl Algorithm for StuckAlgo {
1408 fn name(&self) -> &str {
1409 "stuck"
1410 }
1411
1412 async fn route(
1413 self: Arc<Self>,
1414 _driver: Driver,
1415 _request: Request,
1416 ) -> Result<RoutingOutcome> {
1417 let _guard = DropGuard(self.dropped.clone());
1418 let _ = self.started.send(());
1419 std::future::pending::<()>().await;
1422 unreachable!()
1423 }
1424 }
1425
1426 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1427 let dropped = Arc::new(AtomicBool::new(false));
1428 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1429 started: started_tx,
1430 dropped: dropped.clone(),
1431 });
1432
1433 let run_task = tokio::spawn(async move { test_drive(algo, request(), echo()).await });
1436 started_rx
1437 .recv()
1438 .await
1439 .ok_or_else(|| test_error("task never started"))?;
1440 run_task.abort();
1441 tokio::time::sleep(Duration::from_millis(100)).await;
1442
1443 assert!(
1444 dropped.load(Ordering::SeqCst),
1445 "algorithm task was NOT cancelled after cancelling run"
1446 );
1447 Ok(())
1448 }
1449
1450 struct Hedge {
1455 winner: String,
1456 loser: String,
1457 }
1458
1459 #[async_trait]
1460 impl Algorithm for Hedge {
1461 fn name(&self) -> &str {
1462 "hedge"
1463 }
1464
1465 async fn route(
1466 self: Arc<Self>,
1467 driver: Driver,
1468 request: Request,
1469 ) -> Result<RoutingOutcome> {
1470 let outcome_request = request.clone();
1471 let win = driver.call_model(request.clone(), vec![self.winner.clone().into()]);
1472 let lose = driver.call_model(request, vec![self.loser.clone().into()]);
1473 tokio::select! {
1475 res = win => Ok(RoutingOutcome::answered(
1476 self.winner.clone().into(),
1477 outcome_request,
1478 res?,
1479 )),
1480 res = lose => Ok(RoutingOutcome::answered(
1481 self.loser.clone().into(),
1482 outcome_request,
1483 res?,
1484 )),
1485 }
1486 }
1487 }
1488
1489 fn hedge(loser_delay: Option<std::time::Duration>) -> (Arc<dyn Algorithm>, impl Serve) {
1493 let started = Arc::new(tokio::sync::Notify::new());
1494 let algo = Arc::new(Hedge {
1495 winner: "winner".to_string(),
1496 loser: "loser".to_string(),
1497 });
1498 let serve = move |target: ModelId, _request: Request| {
1499 let started = started.clone();
1500 async move {
1501 if target == "loser" {
1502 started.notify_one();
1503 match loser_delay {
1504 Some(delay) => tokio::time::sleep(delay).await,
1505 None => std::future::pending::<()>().await,
1506 }
1507 } else {
1508 started.notified().await;
1509 }
1510 Ok(reply(target))
1511 }
1512 };
1513 (algo, serve)
1514 }
1515
1516 #[tokio::test]
1517 async fn run_returns_the_winner_without_a_late_loser_overwriting_it() -> Result<()> {
1518 let (algo, serve) = hedge(Some(std::time::Duration::from_millis(50)));
1521 let (_, response) = test_drive(algo, request(), serve).await?;
1522 assert_eq!(
1523 response
1524 .llm_response
1525 .as_agg()
1526 .map(completion_text)
1527 .unwrap_or_default(),
1528 "winner"
1529 );
1530 Ok(())
1531 }
1532
1533 #[tokio::test]
1534 async fn run_returns_the_winner_without_hanging_on_a_pending_loser() -> Result<()> {
1535 let (algo, serve) = hedge(None);
1538 let run = test_drive(algo, request(), serve);
1539 let (_, response) = tokio::time::timeout(std::time::Duration::from_secs(1), run)
1540 .await
1541 .map_err(|error| LibsyError::external("waiting for pending loser", error))??;
1542 assert_eq!(
1543 response
1544 .llm_response
1545 .as_agg()
1546 .map(completion_text)
1547 .unwrap_or_default(),
1548 "winner"
1549 );
1550 Ok(())
1551 }
1552
1553 #[tokio::test]
1554 async fn run_surfaces_a_terminal_error_with_many_calls_in_flight() -> Result<()> {
1555 use std::sync::atomic::{AtomicUsize, Ordering};
1556
1557 const N: usize = 10;
1560
1561 struct FanOutThenError {
1564 all_started: Arc<tokio::sync::Notify>,
1565 n: usize,
1566 }
1567
1568 #[async_trait]
1569 impl Algorithm for FanOutThenError {
1570 fn name(&self) -> &str {
1571 "fan_out_then_error"
1572 }
1573
1574 async fn route(
1575 self: Arc<Self>,
1576 driver: Driver,
1577 request: Request,
1578 ) -> Result<RoutingOutcome> {
1579 let offloads = futures::future::join_all(
1580 (0..self.n)
1581 .map(|i| driver.call_model(request.clone(), vec![format!("m{i}").into()])),
1582 );
1583 tokio::select! {
1584 _ = offloads => Err(test_error("offloads unexpectedly completed")),
1585 _ = self.all_started.notified() => {
1586 Err(test_error("terminal error while calls pending"))
1587 }
1588 }
1589 }
1590 }
1591
1592 let all_started = Arc::new(tokio::sync::Notify::new());
1593 let algo: Arc<dyn Algorithm> = Arc::new(FanOutThenError {
1594 all_started: all_started.clone(),
1595 n: N,
1596 });
1597
1598 let started = Arc::new(AtomicUsize::new(0));
1600 let serve = move |_target: ModelId, _request: Request| {
1601 let started = started.clone();
1602 let all_started = all_started.clone();
1603 async move {
1604 if started.fetch_add(1, Ordering::SeqCst) + 1 == N {
1605 all_started.notify_one();
1606 }
1607 std::future::pending::<ServeResult>().await
1608 }
1609 };
1610
1611 let run = test_drive(algo, request(), serve);
1614 let result = tokio::time::timeout(std::time::Duration::from_millis(500), run)
1615 .await
1616 .map_err(|error| {
1617 LibsyError::external("waiting for terminal error with full call cap", error)
1618 })?;
1619 match result {
1620 Ok(_) => Err(test_error("expected the terminal error, got a response")),
1621 Err(err) => {
1622 assert!(
1623 err.to_string()
1624 .contains("terminal error while calls pending")
1625 );
1626 Ok(())
1627 }
1628 }
1629 }
1630}