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