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 fn needs_history_replay(&self, _request: &Request) -> bool {
624 false
625 }
626
627 async fn route(self: Arc<Self>, driver: Driver, request: Request) -> Result<RoutingOutcome>;
630
631 fn run_stream(self: Arc<Self>, request: Request, models: Arc<RuntimeModels>) -> StepStream {
639 let (driver, step_rx) = Driver::new(self.name(), models);
640 let span = observability::run_span(self.name(), &request);
641 let handle = tokio::spawn(
642 async move {
643 let algorithm = self.name().to_string();
644 let route = AssertUnwindSafe(self.route(driver.clone(), request)).catch_unwind();
646 let result = observability::observe_run(&algorithm, async move {
647 route.await.unwrap_or_else(|payload| {
648 Err(LibsyError::AlgorithmError {
649 message: format!(
650 "algorithm task panicked: {}",
651 panic_message(payload.as_ref())
652 ),
653 })
654 })
655 })
656 .await;
657
658 let _ = driver.finish(result).await;
659 }
660 .instrument(span),
661 );
662 let abort_guard = AbortOnDrop(handle.abort_handle());
664 Box::pin(ReceiverStream::new(step_rx).map(move |step| {
665 let _keep_alive = &abort_guard;
667 step
668 }))
669 }
670}
671
672#[cfg(test)]
673mod tests {
674 use std::collections::HashMap;
675
676 use super::*;
677 use crate::core::testing::{Serve, ServeResult, echo, reply, serve_decision, test_drive};
678 use futures::StreamExt;
679 use switchyard_protocol::{
680 LlmResponse, LlmResponseChunk, completion_text, text_request, text_response,
681 };
682
683 #[derive(Debug, thiserror::Error)]
684 #[error("{0}")]
685 struct TestError(&'static str);
686
687 fn test_error(message: &'static str) -> LibsyError {
688 LibsyError::external("test", TestError(message))
689 }
690
691 struct TestAlgo {
694 target_set: Vec<ModelId>,
695 }
696
697 #[async_trait]
698 impl Algorithm for TestAlgo {
699 fn name(&self) -> &str {
700 "test"
701 }
702
703 async fn route(
704 self: Arc<Self>,
705 driver: Driver,
706 request: Request,
707 ) -> Result<RoutingOutcome> {
708 let target = self
709 .target_set
710 .first()
711 .ok_or(LibsyError::NoTargets)?
712 .clone();
713 let response = driver
714 .call_model(request.clone(), vec![target.clone()])
715 .await?;
716 driver.set_evidence(serde_json::json!({"source": "test"}));
717 driver.set_evidence_if_empty(serde_json::json!({"source": "ignored"}));
718 Ok(RoutingOutcome::answered(target, request, response))
719 }
720 }
721
722 fn orch(target_set: Vec<ModelId>) -> Arc<dyn Algorithm> {
724 Arc::new(TestAlgo { target_set })
725 }
726
727 fn request() -> Request {
728 Request {
729 llm_request: text_request(Some("auto".to_string()), "hi".to_string()),
730 raw_request: None,
731 metadata: None,
732 }
733 }
734
735 #[test]
736 fn routing_outcome_constructors_stamp_selection_and_preserve_payloads() {
737 let outcome = RoutingOutcome::route_to(
738 "selected".into(),
739 target_set(&["fallback-one", "fallback-two"]),
740 request(),
741 );
742
743 assert_eq!(
744 outcome.selected_model_ids,
745 target_set(&["selected", "fallback-one", "fallback-two"])
746 );
747 assert_eq!(outcome.request.model_id().as_deref(), Some("selected"));
748 assert!(outcome.response.is_none());
749 assert!(outcome.metadata.is_none());
750
751 let outcome = RoutingOutcome::route_to("only".into(), Vec::new(), request());
752 assert_eq!(outcome.selected_model_ids, target_set(&["only"]));
753
754 let outcome = RoutingOutcome::answered(
755 "answered".into(),
756 request(),
757 Response {
758 llm_response: LlmResponse::Agg(text_response(None, "existing")),
759 metadata: None,
760 upstream_headers: http::HeaderMap::new(),
761 },
762 );
763
764 assert_eq!(outcome.selected_model_ids, target_set(&["answered"]));
765 assert_eq!(outcome.request.model_id().as_deref(), Some("answered"));
766 assert_eq!(
767 outcome
768 .response
769 .as_ref()
770 .and_then(|response| response.llm_response.as_agg())
771 .map(completion_text),
772 Some("existing".to_string())
773 );
774 }
775
776 fn target_set(names: &[&str]) -> Vec<ModelId> {
777 names.iter().map(|name| ModelId::from(*name)).collect()
778 }
779
780 #[tokio::test]
781 async fn typed_driver_preserves_call_and_stream_boundaries() -> Result<()> {
782 tokio::time::timeout(std::time::Duration::from_secs(1), async {
783 let (driver, mut step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
786 let first_driver = driver.clone();
787 let mut first = tokio::spawn(async move {
788 first_driver
789 .call_model(request(), vec![ModelId::from("first")])
790 .await
791 });
792 let second = tokio::spawn(async move {
793 driver
794 .call_model(request(), vec![ModelId::from("second")])
795 .await
796 });
797
798 let mut calls = HashMap::new();
799 for _ in 0..2 {
800 let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
801 let Step::CallModel(call) = step else {
802 return Err(test_error("expected a CallModel step"));
803 };
804 let selected_model = call
805 .models
806 .first()
807 .ok_or_else(|| test_error("model call has no candidates"))?
808 .to_string();
809 calls.insert(selected_model, call);
810 }
811 assert!(
812 tokio::time::timeout(std::time::Duration::from_millis(20), &mut first)
813 .await
814 .is_err(),
815 "call completed before the host responded"
816 );
817 calls
818 .remove("second")
819 .ok_or_else(|| test_error("missing second call"))?
820 .respond(Ok(reply("second response")))?;
821 calls
822 .remove("first")
823 .ok_or_else(|| test_error("missing first call"))?
824 .respond(Ok(reply("first response")))?;
825
826 let first_response = first
827 .await
828 .map_err(|source| LibsyError::external("joining a test task", source))??;
829 let second_response = second
830 .await
831 .map_err(|source| LibsyError::external("joining a test task", source))??;
832 assert_eq!(
833 first_response.llm_response.as_agg().map(completion_text),
834 Some("first response".to_string())
835 );
836 assert_eq!(
837 second_response.llm_response.as_agg().map(completion_text),
838 Some("second response".to_string())
839 );
840
841 let (driver, mut step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
843 let producer = tokio::spawn(async move {
844 driver
845 .call_model(request(), vec![ModelId::from("dropped")])
846 .await
847 });
848 let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
849 let Step::CallModel(call) = step else {
850 return Err(test_error("expected a CallModel step"));
851 };
852 drop(call);
853 let result = producer
854 .await
855 .map_err(|source| LibsyError::external("joining a test task", source))?;
856 assert!(matches!(
857 result,
858 Err(LibsyError::Driver(DriverError::ResponseDropped))
859 ));
860
861 let (driver, step_rx) = Driver::new("test", Arc::new(RuntimeModels::default()));
863 drop(step_rx);
864 let result = driver
865 .call_model(request(), vec![ModelId::from("closed")])
866 .await;
867 assert!(matches!(
868 result,
869 Err(LibsyError::Driver(DriverError::StreamClosed))
870 ));
871
872 fn decision_request() -> DecisionRequest {
873 DecisionRequest {
874 model: Some("overwritten".into()),
875 context: serde_json::json!({"task": "choose a route"}),
876 questions: Default::default(),
877 }
878 }
879
880 fn decision_response() -> DecisionResponse {
881 DecisionResponse {
882 id: Some("decision-1".to_string()),
883 model: Some("provider-model".into()),
884 answers: [(
885 "p_solve".to_string(),
886 switchyard_protocol::DecisionAnswer {
887 value: switchyard_protocol::DecisionValue::Boolean(
888 switchyard_protocol::BooleanEstimate::ProbabilityTrue(
889 switchyard_protocol::Probability(0.8),
890 ),
891 ),
892 provider_confidence: None,
893 },
894 )]
895 .into(),
896 usage: Default::default(),
897 }
898 }
899
900 struct MixedCalls(&'static str);
901
902 #[async_trait]
903 impl Algorithm for MixedCalls {
904 fn name(&self) -> &str {
905 "mixed"
906 }
907
908 async fn route(
909 self: Arc<Self>,
910 driver: Driver,
911 request: Request,
912 ) -> Result<RoutingOutcome> {
913 let (llm, decision) = tokio::join!(
914 driver.call_model(request.clone(), vec!["llm".into()]),
915 driver.call_decision(decision_request(), "decision".into()),
916 );
917 assert_eq!(
918 llm?.llm_response.as_agg().map(completion_text),
919 Some("llm reply".into())
920 );
921 match self.0 {
922 "mock" => assert_eq!(decision?, DecisionResponse {
923 id: None,
924 model: Some("decision".into()),
925 answers: Default::default(),
926 usage: Default::default(),
927 }),
928 "reply" => assert_eq!(decision?, decision_response()),
929 "error" => assert!(matches!(
930 decision,
931 Err(LibsyError::AlgorithmError { message }) if message == "provider failed"
932 )),
933 "drop" => assert!(matches!(
934 decision,
935 Err(LibsyError::Driver(DriverError::ResponseDropped))
936 )),
937 "abort" => return std::future::pending().await,
938 _ => unreachable!(),
939 }
940 Ok(RoutingOutcome::route_to("answer".into(), vec![], request))
941 }
942 }
943
944 for mode in ["mock", "reply", "error", "drop", "abort"] {
945 let barrier = Arc::new(tokio::sync::Barrier::new(2));
947 let outcome = drive(
948 Arc::new(MixedCalls(mode)),
949 request(),
950 Arc::new(RuntimeModels::default()),
951 move |call| {
952 let barrier = barrier.clone();
953 async move {
954 barrier.wait().await;
955 let call = match call {
956 Call::Model(call) => return call.respond(Ok(reply("llm reply"))),
957 Call::Decision(call) => *call,
958 };
959 assert_eq!(call.algorithm, "mixed");
960 assert_eq!(call.model, "decision");
961 assert_eq!(call.request.model, Some("decision".into()));
962 assert_eq!(call.request.context, decision_request().context);
963 match mode {
964 "mock" => serve_decision(call).await,
965 "reply" => call.respond(Ok(decision_response())),
966 "error" => call.respond(Err(LibsyError::AlgorithmError {
967 message: "provider failed".into(),
968 })),
969 "drop" => {
970 drop(call);
971 Ok(())
972 }
973 "abort" => call.fail(test_error("host aborted")),
974 _ => unreachable!(),
975 }
976 }
977 },
978 )
979 .await;
980 if mode == "abort" {
981 assert!(matches!(outcome, Err(LibsyError::External { .. })));
982 } else {
983 assert_eq!(outcome?.selected_model_id()?, "answer");
984 }
985 }
986
987 let (driver, mut steps) = Driver::new("test", Arc::new(RuntimeModels::default()));
988 let mut pending = Box::pin(driver.call_decision(decision_request(), "decision".into()));
989 assert!(futures::poll!(&mut pending).is_pending());
990 let Some(Ok(Step::CallDecision(call))) = steps.recv().await else {
991 return Err(test_error("expected a decision call"));
992 };
993 drop(pending);
994 assert!(matches!(
995 call.respond(Ok(decision_response())),
996 Err(LibsyError::Driver(DriverError::ResponseDropped))
997 ));
998 drop(steps);
999 assert!(matches!(
1000 driver.call_decision(decision_request(), "decision".into()).await,
1001 Err(LibsyError::Driver(DriverError::StreamClosed))
1002 ));
1003 Ok(())
1004 })
1005 .await
1006 .map_err(|error| LibsyError::external("waiting for typed driver boundaries", error))?
1007 }
1008
1009 #[test]
1010 fn target_lookup_returns_the_missing_target() {
1011 let error = ensure_model_is_target(&target_set(&[]), &ModelId::from("missing")).err();
1012 assert!(matches!(
1013 error,
1014 Some(LibsyError::TargetNotFound { target }) if target == "missing"
1015 ));
1016 }
1017
1018 fn streaming_orch(chunks: Vec<LlmResponseChunk>) -> (Arc<dyn Algorithm>, impl Serve) {
1021 let algo = orch(target_set(&["stream/model"]));
1022 let serve = move |_target: ModelId, _request: Request| {
1023 let chunks = chunks.clone();
1024 async move {
1025 let stream =
1026 futures::stream::iter(chunks.into_iter().map(|chunk| Ok(chunk.into()))).boxed();
1027 Ok(Response {
1028 llm_response: LlmResponse::Stream(stream),
1029 metadata: None,
1030 upstream_headers: http::HeaderMap::new(),
1031 })
1032 }
1033 };
1034 (algo, serve)
1035 }
1036
1037 #[tokio::test]
1038 async fn run_returns_a_streamed_response_the_caller_aggregates() -> Result<()> {
1039 let (orch, serve) = streaming_orch(vec![
1042 LlmResponseChunk::MessageStart {
1043 id: Some("m1".to_string()),
1044 model: Some("stream/model".to_string()),
1045 },
1046 LlmResponseChunk::TextDelta {
1047 index: 0,
1048 text: "hel".to_string(),
1049 },
1050 LlmResponseChunk::TextDelta {
1051 index: 0,
1052 text: "lo".to_string(),
1053 },
1054 LlmResponseChunk::MessageStop {
1055 reason: Some("stop".to_string()),
1056 },
1057 ]);
1058 let (selected_model, response) = test_drive(orch, request(), serve).await?;
1059 let agg = response
1061 .llm_response
1062 .into_agg()
1063 .await
1064 .map_err(|error| LibsyError::external("aggregating response stream", error))?;
1065 assert_eq!(completion_text(&agg), "hello");
1066 assert_eq!(agg.model.as_deref(), Some("stream/model"));
1067 assert_eq!(selected_model, "stream/model");
1068 Ok(())
1069 }
1070
1071 #[tokio::test]
1072 async fn aggregating_a_streamed_response_propagates_a_mid_stream_error() -> Result<()> {
1073 let (orch, serve) = streaming_orch(vec![
1076 LlmResponseChunk::TextDelta {
1077 index: 0,
1078 text: "partial".to_string(),
1079 },
1080 LlmResponseChunk::StreamError {
1081 message: "upstream exploded".to_string(),
1082 },
1083 ]);
1084 let (_, response) = test_drive(orch, request(), serve).await?;
1085 match response.llm_response.into_agg().await {
1086 Ok(_) => panic!("expected a mid-stream error, got an aggregate"),
1087 Err(err) => {
1088 assert!(err.to_string().contains("upstream exploded"));
1089 Ok(())
1090 }
1091 }
1092 }
1093
1094 #[tokio::test]
1095 async fn run_offloads_via_promise_then_finishes() -> Result<()> {
1096 let stream = orch(target_set(&["offload/model"]))
1099 .run_stream(request(), Arc::new(RuntimeModels::default()));
1100 tokio::pin!(stream);
1101
1102 let mut saw_call = false;
1103 let mut final_completion = None;
1104 while let Some(step) = stream.next().await {
1105 match step? {
1106 Step::CallDecision(_) => return Err(test_error("unexpected decision call")),
1107 Step::CallModel(call) => {
1108 saw_call = true;
1109 assert_eq!(call.models, vec![ModelId::from("offload/model")]);
1110 call.respond(Ok(Response {
1112 llm_response: LlmResponse::Agg(text_response(
1113 None,
1114 "fulfilled".to_string(),
1115 )),
1116 metadata: None,
1117 upstream_headers: http::HeaderMap::new(),
1118 }))?;
1119 }
1120 Step::Done(outcome) => {
1121 let metadata = outcome
1122 .metadata
1123 .as_ref()
1124 .expect("run_stream should attach outcome metadata");
1125 assert_eq!(metadata.algorithm, "test");
1126 assert_eq!(
1127 uuid::Uuid::parse_str(metadata.outcome_id())
1128 .expect("outcome id should be a UUID")
1129 .get_version_num(),
1130 7
1131 );
1132 assert_eq!(
1133 metadata.evidence,
1134 Some(serde_json::json!({"source": "test"}))
1135 );
1136 let response = outcome
1137 .response
1138 .ok_or_else(|| test_error("expected an answered outcome"))?;
1139 final_completion = Some(
1140 response
1141 .llm_response
1142 .as_agg()
1143 .map(completion_text)
1144 .unwrap_or_default(),
1145 );
1146 }
1147 }
1148 }
1149
1150 assert!(saw_call, "expected a CallModel step before Done");
1151 assert_eq!(
1152 final_completion.ok_or_else(|| test_error("no Done step"))?,
1153 "fulfilled"
1154 );
1155 Ok(())
1156 }
1157
1158 #[tokio::test(flavor = "multi_thread", worker_threads = 12)]
1159 async fn requests_are_processed_in_parallel() -> Result<()> {
1160 use std::time::Duration;
1161 use tokio::sync::Barrier;
1162
1163 const N: usize = 12;
1164
1165 let barrier = Arc::new(Barrier::new(N));
1170 let algo = orch(target_set(&["m"]));
1172
1173 let mut handles = Vec::new();
1174 for _ in 0..N {
1175 let algo = algo.clone();
1176 let barrier = barrier.clone();
1177 let serve = move |target: ModelId, _request: Request| {
1178 let barrier = barrier.clone();
1179 async move {
1180 barrier.wait().await;
1181 Ok(reply(target))
1182 }
1183 };
1184 handles.push(tokio::spawn(async move {
1185 test_drive(algo, request(), serve)
1186 .await
1187 .map(|(_, response)| {
1188 response
1189 .llm_response
1190 .as_agg()
1191 .map(completion_text)
1192 .unwrap_or_default()
1193 })
1194 }));
1195 }
1196
1197 for handle in handles {
1198 let completion = tokio::time::timeout(Duration::from_secs(5), handle)
1200 .await
1201 .map_err(|error| LibsyError::external("waiting for test task", error))?
1202 .map_err(|source| LibsyError::external("joining a test task", source))??;
1203 assert_eq!(completion, "m");
1204 }
1205 Ok(())
1206 }
1207
1208 #[tokio::test]
1209 async fn offload_error_propagates_back_to_the_algorithm() -> Result<()> {
1210 let stream = orch(target_set(&["offload/model"]))
1214 .run_stream(request(), Arc::new(RuntimeModels::default()));
1215 tokio::pin!(stream);
1216
1217 let mut saw_error = false;
1218 while let Some(step) = stream.next().await {
1219 match step {
1220 Ok(Step::CallDecision(_)) => return Err(test_error("unexpected decision call")),
1221 Ok(Step::CallModel(call)) => {
1222 call.respond(Err(test_error("upstream model call failed")))?;
1223 }
1224 Ok(Step::Done(..)) => {
1225 return Err(test_error(
1226 "expected the offload error to propagate, got a response",
1227 ));
1228 }
1229 Err(err) => {
1230 assert!(err.to_string().contains("upstream model call failed"));
1232 saw_error = true;
1233 }
1234 }
1235 }
1236
1237 assert!(saw_error, "expected an error step");
1238 Ok(())
1239 }
1240
1241 #[tokio::test]
1242 async fn dropping_the_stream_cancels_the_algorithm_task() -> Result<()> {
1243 use std::sync::atomic::{AtomicBool, Ordering};
1244 use std::time::Duration;
1245 use tokio::sync::mpsc;
1246
1247 struct DropGuard(Arc<AtomicBool>);
1250 impl Drop for DropGuard {
1251 fn drop(&mut self) {
1252 self.0.store(true, Ordering::SeqCst);
1253 }
1254 }
1255
1256 struct StuckAlgo {
1257 started: mpsc::UnboundedSender<()>,
1258 dropped: Arc<AtomicBool>,
1259 }
1260
1261 #[async_trait]
1262 impl Algorithm for StuckAlgo {
1263 fn name(&self) -> &str {
1264 "stuck"
1265 }
1266
1267 async fn route(
1268 self: Arc<Self>,
1269 _driver: Driver,
1270 _request: Request,
1271 ) -> Result<RoutingOutcome> {
1272 let _guard = DropGuard(self.dropped.clone());
1273 let _ = self.started.send(());
1274 std::future::pending::<()>().await;
1276 unreachable!()
1277 }
1278 }
1279
1280 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1281 let dropped = Arc::new(AtomicBool::new(false));
1282 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1283 started: started_tx,
1284 dropped: dropped.clone(),
1285 });
1286
1287 let stream = algo.run_stream(request(), Arc::new(RuntimeModels::default()));
1288 started_rx
1289 .recv()
1290 .await
1291 .ok_or_else(|| test_error("task never started"))?;
1292 drop(stream);
1293 tokio::time::sleep(Duration::from_millis(100)).await;
1294
1295 assert!(
1296 dropped.load(Ordering::SeqCst),
1297 "algorithm task was NOT cancelled after dropping the stream"
1298 );
1299 Ok(())
1300 }
1301
1302 #[tokio::test]
1303 async fn route_panic_surfaces_as_a_stream_error() -> Result<()> {
1304 struct Panicky;
1307
1308 #[async_trait]
1309 impl Algorithm for Panicky {
1310 fn name(&self) -> &str {
1311 "panicky"
1312 }
1313
1314 async fn route(
1315 self: Arc<Self>,
1316 _driver: Driver,
1317 _request: Request,
1318 ) -> Result<RoutingOutcome> {
1319 panic!("boom");
1320 }
1321 }
1322
1323 let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
1324 let stream = algo.run_stream(request(), Arc::new(RuntimeModels::default()));
1325 tokio::pin!(stream);
1326
1327 let mut saw_error = false;
1328 while let Some(step) = stream.next().await {
1329 match step {
1330 Err(err) => {
1331 assert!(err.to_string().contains("algorithm task panicked: boom"));
1333 saw_error = true;
1334 }
1335 Ok(_) => return Err(test_error("expected the panic to surface as an error step")),
1336 }
1337 }
1338
1339 assert!(saw_error, "expected an error step from the panicked task");
1340 Ok(())
1341 }
1342
1343 #[tokio::test]
1347 async fn a_panic_with_a_leaked_driver_clone_still_terminates_the_run() -> Result<()> {
1348 struct LeakyPanic;
1349
1350 #[async_trait]
1351 impl Algorithm for LeakyPanic {
1352 fn name(&self) -> &str {
1353 "leaky_panic"
1354 }
1355
1356 async fn route(
1357 self: Arc<Self>,
1358 driver: Driver,
1359 _request: Request,
1360 ) -> Result<RoutingOutcome> {
1361 tokio::spawn(async move {
1362 let _keep_alive = driver;
1364 std::future::pending::<()>().await;
1365 });
1366 tokio::task::yield_now().await;
1367 panic!("boom");
1368 }
1369 }
1370
1371 let algo: Arc<dyn Algorithm> = Arc::new(LeakyPanic);
1372 let result = tokio::time::timeout(
1374 std::time::Duration::from_secs(1),
1375 test_drive(algo, request(), echo()),
1376 )
1377 .await
1378 .map_err(|error| LibsyError::external("waiting for the panicked run to end", error))?;
1379
1380 match result {
1381 Ok(_) => Err(test_error(
1382 "expected the panic to end the run with an error",
1383 )),
1384 Err(err) => {
1385 assert!(err.to_string().contains("algorithm task panicked: boom"));
1386 Ok(())
1387 }
1388 }
1389 }
1390
1391 #[tokio::test]
1392 async fn cancelling_run_cancels_the_algorithm_task() -> Result<()> {
1393 use std::sync::atomic::{AtomicBool, Ordering};
1394 use std::time::Duration;
1395 use tokio::sync::mpsc;
1396
1397 struct DropGuard(Arc<AtomicBool>);
1400 impl Drop for DropGuard {
1401 fn drop(&mut self) {
1402 self.0.store(true, Ordering::SeqCst);
1403 }
1404 }
1405
1406 struct StuckAlgo {
1407 started: mpsc::UnboundedSender<()>,
1408 dropped: Arc<AtomicBool>,
1409 }
1410
1411 #[async_trait]
1412 impl Algorithm for StuckAlgo {
1413 fn name(&self) -> &str {
1414 "stuck"
1415 }
1416
1417 async fn route(
1418 self: Arc<Self>,
1419 _driver: Driver,
1420 _request: Request,
1421 ) -> Result<RoutingOutcome> {
1422 let _guard = DropGuard(self.dropped.clone());
1423 let _ = self.started.send(());
1424 std::future::pending::<()>().await;
1427 unreachable!()
1428 }
1429 }
1430
1431 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1432 let dropped = Arc::new(AtomicBool::new(false));
1433 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1434 started: started_tx,
1435 dropped: dropped.clone(),
1436 });
1437
1438 let run_task = tokio::spawn(async move { test_drive(algo, request(), echo()).await });
1441 started_rx
1442 .recv()
1443 .await
1444 .ok_or_else(|| test_error("task never started"))?;
1445 run_task.abort();
1446 tokio::time::sleep(Duration::from_millis(100)).await;
1447
1448 assert!(
1449 dropped.load(Ordering::SeqCst),
1450 "algorithm task was NOT cancelled after cancelling run"
1451 );
1452 Ok(())
1453 }
1454
1455 struct Hedge {
1460 winner: String,
1461 loser: String,
1462 }
1463
1464 #[async_trait]
1465 impl Algorithm for Hedge {
1466 fn name(&self) -> &str {
1467 "hedge"
1468 }
1469
1470 async fn route(
1471 self: Arc<Self>,
1472 driver: Driver,
1473 request: Request,
1474 ) -> Result<RoutingOutcome> {
1475 let outcome_request = request.clone();
1476 let win = driver.call_model(request.clone(), vec![self.winner.clone().into()]);
1477 let lose = driver.call_model(request, vec![self.loser.clone().into()]);
1478 tokio::select! {
1480 res = win => Ok(RoutingOutcome::answered(
1481 self.winner.clone().into(),
1482 outcome_request,
1483 res?,
1484 )),
1485 res = lose => Ok(RoutingOutcome::answered(
1486 self.loser.clone().into(),
1487 outcome_request,
1488 res?,
1489 )),
1490 }
1491 }
1492 }
1493
1494 fn hedge(loser_delay: Option<std::time::Duration>) -> (Arc<dyn Algorithm>, impl Serve) {
1498 let started = Arc::new(tokio::sync::Notify::new());
1499 let algo = Arc::new(Hedge {
1500 winner: "winner".to_string(),
1501 loser: "loser".to_string(),
1502 });
1503 let serve = move |target: ModelId, _request: Request| {
1504 let started = started.clone();
1505 async move {
1506 if target == "loser" {
1507 started.notify_one();
1508 match loser_delay {
1509 Some(delay) => tokio::time::sleep(delay).await,
1510 None => std::future::pending::<()>().await,
1511 }
1512 } else {
1513 started.notified().await;
1514 }
1515 Ok(reply(target))
1516 }
1517 };
1518 (algo, serve)
1519 }
1520
1521 #[tokio::test]
1522 async fn run_returns_the_winner_without_a_late_loser_overwriting_it() -> Result<()> {
1523 let (algo, serve) = hedge(Some(std::time::Duration::from_millis(50)));
1526 let (_, response) = test_drive(algo, request(), serve).await?;
1527 assert_eq!(
1528 response
1529 .llm_response
1530 .as_agg()
1531 .map(completion_text)
1532 .unwrap_or_default(),
1533 "winner"
1534 );
1535 Ok(())
1536 }
1537
1538 #[tokio::test]
1539 async fn run_returns_the_winner_without_hanging_on_a_pending_loser() -> Result<()> {
1540 let (algo, serve) = hedge(None);
1543 let run = test_drive(algo, request(), serve);
1544 let (_, response) = tokio::time::timeout(std::time::Duration::from_secs(1), run)
1545 .await
1546 .map_err(|error| LibsyError::external("waiting for pending loser", error))??;
1547 assert_eq!(
1548 response
1549 .llm_response
1550 .as_agg()
1551 .map(completion_text)
1552 .unwrap_or_default(),
1553 "winner"
1554 );
1555 Ok(())
1556 }
1557
1558 #[tokio::test]
1559 async fn run_surfaces_a_terminal_error_with_many_calls_in_flight() -> Result<()> {
1560 use std::sync::atomic::{AtomicUsize, Ordering};
1561
1562 const N: usize = 10;
1565
1566 struct FanOutThenError {
1569 all_started: Arc<tokio::sync::Notify>,
1570 n: usize,
1571 }
1572
1573 #[async_trait]
1574 impl Algorithm for FanOutThenError {
1575 fn name(&self) -> &str {
1576 "fan_out_then_error"
1577 }
1578
1579 async fn route(
1580 self: Arc<Self>,
1581 driver: Driver,
1582 request: Request,
1583 ) -> Result<RoutingOutcome> {
1584 let offloads = futures::future::join_all(
1585 (0..self.n)
1586 .map(|i| driver.call_model(request.clone(), vec![format!("m{i}").into()])),
1587 );
1588 tokio::select! {
1589 _ = offloads => Err(test_error("offloads unexpectedly completed")),
1590 _ = self.all_started.notified() => {
1591 Err(test_error("terminal error while calls pending"))
1592 }
1593 }
1594 }
1595 }
1596
1597 let all_started = Arc::new(tokio::sync::Notify::new());
1598 let algo: Arc<dyn Algorithm> = Arc::new(FanOutThenError {
1599 all_started: all_started.clone(),
1600 n: N,
1601 });
1602
1603 let started = Arc::new(AtomicUsize::new(0));
1605 let serve = move |_target: ModelId, _request: Request| {
1606 let started = started.clone();
1607 let all_started = all_started.clone();
1608 async move {
1609 if started.fetch_add(1, Ordering::SeqCst) + 1 == N {
1610 all_started.notify_one();
1611 }
1612 std::future::pending::<ServeResult>().await
1613 }
1614 };
1615
1616 let run = test_drive(algo, request(), serve);
1619 let result = tokio::time::timeout(std::time::Duration::from_millis(500), run)
1620 .await
1621 .map_err(|error| {
1622 LibsyError::external("waiting for terminal error with full call cap", error)
1623 })?;
1624 match result {
1625 Ok(_) => Err(test_error("expected the terminal error, got a response")),
1626 Err(err) => {
1627 assert!(
1628 err.to_string()
1629 .contains("terminal error while calls pending")
1630 );
1631 Ok(())
1632 }
1633 }
1634 }
1635}