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