1use std::{
9 collections::{HashMap, HashSet},
10 future::Future,
11 pin::Pin,
12 sync::Arc,
13 time::Instant,
14};
15
16use async_trait::async_trait;
17use futures::{Stream, StreamExt};
18use parking_lot::Mutex;
19use tokio::sync::{mpsc, oneshot};
20use tokio_stream::wrappers::ReceiverStream;
21use tracing::Instrument;
22
23use switchyard_protocol::{
31 Decision, LlmClientError, Request, Response, RoutingFallbackReason, Signals,
32};
33
34use crate::{DriverError, LibsyError, Result, observability};
35
36pub type StepStream = Pin<Box<dyn Stream<Item = Result<Step>> + Send>>;
40
41pub struct CallModel {
55 pub algorithm: String,
58 pub request: Request,
60 pub decision: Decision,
62 reply: oneshot::Sender<Result<Response>>,
63}
64
65impl CallModel {
66 pub fn respond(self, result: Result<Response>) -> Result<()> {
70 self.reply
71 .send(result)
72 .map_err(|_| DriverError::ResponseDropped.into())
73 }
74
75 pub fn into_parts(self) -> (Request, Decision) {
86 let Self {
87 request,
88 decision,
89 reply,
90 ..
91 } = self;
92 let _ = reply.send(Err(DriverError::Abandoned.into()));
96 (request, decision)
97 }
98}
99
100#[derive(Clone)]
107pub struct Driver {
108 step_tx: mpsc::Sender<Result<Step>>,
109 algorithm: String,
112}
113
114impl Driver {
115 pub(crate) fn new(algorithm: &str) -> (Self, mpsc::Receiver<Result<Step>>) {
118 let (step_tx, step_rx) = mpsc::channel(1);
123 (
124 Self {
125 step_tx,
126 algorithm: algorithm.to_string(),
127 },
128 step_rx,
129 )
130 }
131
132 #[tracing::instrument(
140 target = "libsy",
141 name = "libsy.llm_call",
142 skip_all,
143 fields(
144 algorithm = self.algorithm,
145 selected_model = decision.selected_model_id(),
146 openinference.span.kind = "CHAIN",
147 outcome = tracing::field::Empty,
148 error = tracing::field::Empty,
149 input_tokens = tracing::field::Empty,
150 output_tokens = tracing::field::Empty,
151 total_tokens = tracing::field::Empty,
152 reasoning_tokens = tracing::field::Empty,
153 )
154 )]
155 pub async fn call_model(&self, request: Request, decision: Decision) -> Result<Response> {
156 let selected_model_id = decision.selected_model_id().to_string();
157 let is_answer_call = decision.is_answer_call();
158 let started = Instant::now();
159 let (reply, response) = oneshot::channel::<Result<Response>>();
160 let call = CallModel {
161 algorithm: self.algorithm.clone(),
162 request,
163 decision,
164 reply,
165 };
166 let result = async {
167 self.step_tx
168 .send(Ok(Step::CallModel(Box::new(call))))
169 .await
170 .map_err(|_| DriverError::StreamClosed)?;
171 response
172 .await
173 .map_err(|_| LibsyError::from(DriverError::ResponseDropped))?
174 }
175 .await;
176 let elapsed = started.elapsed();
177 observability::record_llm_call(
178 &self.algorithm,
179 &selected_model_id,
180 is_answer_call,
181 elapsed,
182 &result,
183 &tracing::Span::current(),
184 );
185 result
186 }
187
188 pub async fn info(&self, decision: Decision) -> Result<()> {
192 self.step_tx
193 .send(Ok(Step::Decision(decision.clone())))
194 .await
195 .map_err(|_| DriverError::StreamClosed)?;
196 observability::record_decision(&self.algorithm, &decision);
197 Ok(())
198 }
199
200 pub(crate) async fn finish(&self, result: Result<Response>) -> Result<()> {
204 let step = result.map(|response| Step::ReturnToAgent(Box::new(response)));
205 self.step_tx
206 .send(step)
207 .await
208 .map_err(|_| DriverError::StreamClosed.into())
209 }
210}
211
212pub enum Step {
214 CallModel(Box<CallModel>),
217 Decision(Decision),
220 ReturnToAgent(Box<Response>),
222}
223
224pub async fn drive<F, Fut>(
237 algorithm: Arc<dyn Algorithm>,
238 request: Request,
239 serve: F,
240) -> Result<(Vec<Decision>, Response)>
241where
242 F: Fn(CallModel) -> Fut,
243 Fut: Future<Output = Result<()>>,
244{
245 let stream = algorithm.run_stream(request);
246 tokio::pin!(stream);
247
248 let mut trace: Vec<Decision> = Vec::new();
249 let mut in_flight = futures::stream::FuturesUnordered::new();
250 let mut final_response: Option<Response> = None;
251
252 loop {
253 tokio::select! {
254 Some(result) = in_flight.next() => match result {
255 Ok(()) => {}, Err(err) => return Err(err), },
258 step = stream.next() => {
259 match step {
260 None => break, Some(item) => match item? {
262 Step::CallModel(call) => in_flight.push(serve(*call)),
263 Step::Decision(decision) => trace.push(decision),
264 Step::ReturnToAgent(response) => {
265 final_response = Some(*response);
266 break;
267 }
268 }
269 }
270 },
271 }
272 }
273 final_response
274 .map(|response| (trace, response))
275 .ok_or(LibsyError::MissingFinalResponse)
276}
277
278struct AbortOnDrop(tokio::task::AbortHandle);
280
281impl Drop for AbortOnDrop {
282 fn drop(&mut self) {
283 self.0.abort();
284 }
285}
286
287#[derive(Clone)]
291pub struct LlmTarget {
292 pub semantic_name: String,
296}
297
298#[derive(Clone)]
302pub struct LlmTargetSet {
303 targets: Vec<LlmTarget>,
304}
305
306impl LlmTargetSet {
307 pub fn new(targets: Vec<LlmTarget>) -> Self {
309 Self { targets }
310 }
311
312 pub fn targets(&self) -> &[LlmTarget] {
314 &self.targets
315 }
316
317 pub fn get_target(&self, name: &str) -> Result<LlmTarget> {
319 self.targets
320 .iter()
321 .find(|t| t.semantic_name == name)
322 .cloned()
323 .ok_or_else(|| LibsyError::TargetNotFound {
324 target: name.to_string(),
325 })
326 }
327
328 pub fn resolve_target(&self, name: &str, excluded: &HashSet<String>) -> Result<LlmTarget> {
331 let target = self.get_target(name)?;
332 if !excluded.contains(&target.semantic_name) {
333 return Ok(target);
334 }
335 self.targets
336 .iter()
337 .find(|t| !excluded.contains(&t.semantic_name))
338 .cloned()
339 .ok_or(LibsyError::AllTargetsExcluded)
340 }
341}
342
343#[derive(Clone, Hash, PartialEq, Eq)]
347pub(crate) enum RoutingIdentity {
348 Session(String),
350 Subagent { session: String, agent: String },
352}
353
354impl RoutingIdentity {
355 pub(crate) fn from_request(request: &Request) -> Option<Self> {
360 let metadata = request.metadata.as_ref()?;
361 let session = metadata.session_id.as_deref().filter(|id| !id.is_empty())?;
362 if metadata.is_subagent {
363 let agent = metadata.agent_id.as_deref().filter(|id| !id.is_empty())?;
364 Some(Self::Subagent {
365 session: session.to_string(),
366 agent: agent.to_string(),
367 })
368 } else {
369 Some(Self::Session(session.to_string()))
370 }
371 }
372
373 fn session(&self) -> &str {
375 match self {
376 Self::Session(session) | Self::Subagent { session, .. } => session,
377 }
378 }
379}
380
381const MAX_EVICTION_IDENTITIES: usize = 1_024;
384
385#[derive(Default)]
391pub(crate) struct SessionEvictions {
392 by_identity: Mutex<HashMap<RoutingIdentity, HashSet<String>>>,
393}
394
395impl SessionEvictions {
396 pub(crate) fn remove_session(&self, session: &str) {
398 self.by_identity
399 .lock()
400 .retain(|identity, _| identity.session() != session);
401 }
402
403 fn evicted_for(&self, identity: Option<&RoutingIdentity>) -> Vec<String> {
405 let Some(identity) = identity else {
406 return Vec::new();
407 };
408 self.by_identity
409 .lock()
410 .get(identity)
411 .map(|targets| targets.iter().cloned().collect())
412 .unwrap_or_default()
413 }
414
415 fn record(&self, identity: Option<&RoutingIdentity>, target: &str) {
418 let Some(identity) = identity else { return };
419 let mut histories = self.by_identity.lock();
420 if histories.len() >= MAX_EVICTION_IDENTITIES
421 && !histories.contains_key(identity)
422 && let Some(oldest) = histories.keys().next().cloned()
423 {
424 histories.remove(&oldest);
425 }
426 histories
427 .entry(identity.clone())
428 .or_default()
429 .insert(target.to_string());
430 }
431}
432
433fn eligible_targets(targets: &LlmTargetSet, excluded: &HashSet<String>) -> usize {
435 targets
436 .targets()
437 .iter()
438 .filter(|t| !excluded.contains(&t.semantic_name))
439 .count()
440}
441
442pub(crate) fn exclude_evicted(
445 excluded: &mut HashSet<String>,
446 targets: &LlmTargetSet,
447 evictions: &SessionEvictions,
448 identity: Option<&RoutingIdentity>,
449) {
450 for target in evictions.evicted_for(identity) {
451 if eligible_targets(targets, excluded) <= 1 {
454 break;
455 }
456 excluded.insert(target);
457 }
458}
459
460fn classify_fallback(error: &LibsyError) -> Option<(&str, RoutingFallbackReason)> {
462 let LibsyError::ClientCall { target, source } = error else {
463 return None;
464 };
465 let reason = match source {
466 LlmClientError::ContextWindowExceeded { .. } => RoutingFallbackReason::ContextWindow,
467 LlmClientError::Transport { .. } | LlmClientError::Timeout { .. } => {
468 RoutingFallbackReason::Unavailable
469 }
470 LlmClientError::UpstreamHttp { status, .. }
471 if matches!(*status, 403 | 408 | 429) || (500..=599).contains(status) =>
472 {
473 RoutingFallbackReason::Unavailable
474 }
475 _ => return None,
476 };
477 Some((target, reason))
478}
479
480#[allow(clippy::too_many_arguments)]
488pub(crate) async fn call_model_with_fallback(
489 excluded: &mut HashSet<String>,
490 driver: &Driver,
491 targets: &LlmTargetSet,
492 mut target: LlmTarget,
493 mut decision: Decision,
494 request: Request,
495 identity: Option<&RoutingIdentity>,
496 evictions: &SessionEvictions,
497 target_unavailable: impl Fn(&Request, &str),
498 fallback_decision: impl Fn(&LlmTarget, &LlmTarget, RoutingFallbackReason) -> Decision,
499) -> Result<Response> {
500 loop {
501 let result = driver.call_model(request.clone(), decision.clone()).await;
502 let Err(error) = result else { return result };
503 let Some((failed, reason)) = classify_fallback(&error) else {
504 return Err(error);
505 };
506 if !excluded.insert(failed.to_string()) {
509 return Err(error);
510 }
511 match reason {
512 RoutingFallbackReason::ContextWindow => evictions.record(identity, failed),
513 RoutingFallbackReason::Unavailable => target_unavailable(&request, failed),
514 }
515 let Ok(next) = targets.resolve_target(&target.semantic_name, excluded) else {
516 return Err(error);
517 };
518 decision = fallback_decision(&target, &next, reason);
519 target = next;
520 driver.info(decision.clone()).await?;
521 }
522}
523
524#[async_trait]
546pub trait Algorithm: Send + Sync + 'static {
547 fn name(&self) -> &str;
551
552 async fn create_run_task(self: Arc<Self>, driver: Driver, request: Request)
556 -> Result<Response>;
557
558 #[allow(unused_variables)]
562 async fn process_signals(self: Arc<Self>, signals: Signals) -> Result<()> {
563 Ok(())
564 }
565
566 fn run_stream(self: Arc<Self>, request: Request) -> StepStream {
575 let (driver, step_rx) = Driver::new(self.name());
576 let task_driver = driver.clone();
577 let stream = ReceiverStream::new(step_rx);
578 let span = observability::run_span(self.name(), &request);
582 let handle = tokio::spawn(
583 async move {
584 let algorithm = self.name().to_string();
585 observability::observe_run(&algorithm, self.create_run_task(task_driver, request))
586 .await
587 }
588 .instrument(span),
589 );
590 let abort_guard = AbortOnDrop(handle.abort_handle());
592
593 let finish_driver = driver.clone();
594 let tail: StepStream = Box::pin(
595 futures::stream::once(async move {
596 let result = match handle.await {
597 Ok(response) => response,
598 Err(source) => Err(LibsyError::AlgorithmTask { source }),
599 };
600 finish_driver.finish(result).await
601 })
602 .filter_map(|finish_result| async move { finish_result.err().map(Err) }),
603 );
604
605 let stream: StepStream = Box::pin(stream);
606 Box::pin(futures::stream::select(stream, tail).map(move |step| {
607 let _keep_alive = &abort_guard;
609 step
610 }))
611 }
612}
613
614#[cfg(test)]
615mod tests {
616 use super::*;
617 use crate::core::testing::{Serve, ServeResult, echo, reply, test_drive};
618 use futures::StreamExt;
619 use switchyard_protocol::{
620 LlmResponse, LlmResponseChunk, completion_text, text_request, text_response,
621 };
622
623 #[derive(Debug, thiserror::Error)]
624 #[error("{0}")]
625 struct TestError(&'static str);
626
627 fn test_error(message: &'static str) -> LibsyError {
628 LibsyError::external("test", TestError(message))
629 }
630
631 fn classified_client_error(source: LlmClientError) -> Option<RoutingFallbackReason> {
632 classify_fallback(&LibsyError::client_call("target", source)).map(|(_, reason)| reason)
633 }
634
635 #[test]
636 fn route_fallback_only_accepts_context_and_unavailable_failures() {
637 assert_eq!(
638 classified_client_error(LlmClientError::ContextWindowExceeded {
639 model: "target".to_string(),
640 message: "too long".to_string(),
641 }),
642 Some(RoutingFallbackReason::ContextWindow)
643 );
644 for source in [
645 LlmClientError::Transport {
646 source: Box::new(std::io::Error::other("connection failed")),
647 },
648 LlmClientError::Timeout {
649 source: Box::new(std::io::Error::other("request timed out")),
650 },
651 ] {
652 assert_eq!(
653 classified_client_error(source),
654 Some(RoutingFallbackReason::Unavailable)
655 );
656 }
657 for (status, expected) in [
658 (400, None),
659 (401, None),
660 (403, Some(RoutingFallbackReason::Unavailable)),
661 (404, None),
662 (408, Some(RoutingFallbackReason::Unavailable)),
663 (409, None),
664 (429, Some(RoutingFallbackReason::Unavailable)),
665 (499, None),
666 (500, Some(RoutingFallbackReason::Unavailable)),
667 (599, Some(RoutingFallbackReason::Unavailable)),
668 (600, None),
669 ] {
670 assert_eq!(
671 classified_client_error(LlmClientError::UpstreamHttp {
672 status,
673 body: "failed".to_string(),
674 }),
675 expected
676 );
677 }
678 assert_eq!(
679 classified_client_error(LlmClientError::InvalidResponse {
680 source: Box::new(std::io::Error::other("invalid response")),
681 }),
682 None
683 );
684 }
685
686 fn test_decision(selected_model_id: String) -> Decision {
688 Decision::new(selected_model_id, None, true)
689 }
690
691 struct TestAlgo {
694 target_set: LlmTargetSet,
695 }
696
697 #[async_trait]
698 impl Algorithm for TestAlgo {
699 fn name(&self) -> &str {
700 "test"
701 }
702
703 async fn create_run_task(
704 self: Arc<Self>,
705 driver: Driver,
706 request: Request,
707 ) -> Result<Response> {
708 let target = self
709 .target_set
710 .targets()
711 .first()
712 .ok_or(LibsyError::NoTargets)?
713 .clone();
714 let decision = test_decision(target.semantic_name.clone());
715 driver.info(decision.clone()).await?;
716 driver.call_model(request, decision).await
717 }
718 }
719
720 fn orch(target_set: LlmTargetSet) -> Arc<dyn Algorithm> {
722 Arc::new(TestAlgo { target_set })
723 }
724
725 fn request() -> Request {
726 Request {
727 llm_request: text_request(Some("auto".to_string()), "hi".to_string()),
728 raw_request: None,
729 metadata: None,
730 }
731 }
732
733 fn target_set(names: &[&str]) -> LlmTargetSet {
734 let targets = names
735 .iter()
736 .map(|name| LlmTarget {
737 semantic_name: name.to_string(),
738 })
739 .collect();
740 LlmTargetSet::new(targets)
741 }
742
743 #[tokio::test]
744 async fn typed_driver_preserves_call_and_stream_boundaries() -> Result<()> {
745 tokio::time::timeout(std::time::Duration::from_secs(1), async {
746 let (driver, mut step_rx) = Driver::new("test");
749 let first_driver = driver.clone();
750 let mut first = tokio::spawn(async move {
751 first_driver
752 .call_model(request(), test_decision("first".to_string()))
753 .await
754 });
755 let second = tokio::spawn(async move {
756 driver
757 .call_model(request(), test_decision("second".to_string()))
758 .await
759 });
760
761 let mut calls = HashMap::new();
762 for _ in 0..2 {
763 let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
764 let Step::CallModel(call) = step else {
765 return Err(test_error("expected a CallModel step"));
766 };
767 calls.insert(call.decision.selected_model_id().to_string(), call);
768 }
769 assert!(
770 tokio::time::timeout(std::time::Duration::from_millis(20), &mut first)
771 .await
772 .is_err(),
773 "call completed before the host responded"
774 );
775 calls
776 .remove("second")
777 .ok_or_else(|| test_error("missing second call"))?
778 .respond(Ok(reply("second response")))?;
779 calls
780 .remove("first")
781 .ok_or_else(|| test_error("missing first call"))?
782 .respond(Ok(reply("first response")))?;
783
784 let first_response = first
785 .await
786 .map_err(|source| LibsyError::AlgorithmTask { source })??;
787 let second_response = second
788 .await
789 .map_err(|source| LibsyError::AlgorithmTask { source })??;
790 assert_eq!(
791 first_response.llm_response.as_agg().map(completion_text),
792 Some("first response".to_string())
793 );
794 assert_eq!(
795 second_response.llm_response.as_agg().map(completion_text),
796 Some("second response".to_string())
797 );
798
799 let (driver, mut step_rx) = Driver::new("test");
801 let producer = tokio::spawn(async move {
802 driver
803 .call_model(request(), test_decision("dropped".to_string()))
804 .await
805 });
806 let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
807 let Step::CallModel(call) = step else {
808 return Err(test_error("expected a CallModel step"));
809 };
810 drop(call);
811 let result = producer
812 .await
813 .map_err(|source| LibsyError::AlgorithmTask { source })?;
814 assert!(matches!(
815 result,
816 Err(LibsyError::Driver(DriverError::ResponseDropped))
817 ));
818
819 let (driver, step_rx) = Driver::new("test");
821 drop(step_rx);
822 let decision = test_decision("closed".to_string());
823 let result = driver.info(decision).await;
824 assert!(matches!(
825 result,
826 Err(LibsyError::Driver(DriverError::StreamClosed))
827 ));
828 Ok(())
829 })
830 .await
831 .map_err(|error| LibsyError::external("waiting for typed driver boundaries", error))?
832 }
833
834 #[tokio::test]
839 async fn into_parts_yields_the_call_without_answering_it() -> Result<()> {
840 let (driver, mut step_rx) = Driver::new("test");
841 let decision = test_decision("answer/model".to_string());
842 let producer = tokio::spawn({
843 let decision = decision.clone();
844 async move { driver.call_model(request(), decision).await }
845 });
846
847 let step = step_rx.recv().await.ok_or(DriverError::StreamClosed)??;
848 let Step::CallModel(call) = step else {
849 return Err(test_error("expected a CallModel step"));
850 };
851 let (taken_request, taken_decision) = call.into_parts();
852 assert_eq!(taken_decision.selected_model_id(), "answer/model");
853 assert!(taken_decision.is_answer_call());
854 assert_eq!(taken_request.llm_request, request().llm_request);
855
856 let result = producer
857 .await
858 .map_err(|source| LibsyError::AlgorithmTask { source })?;
859 assert!(matches!(
860 result,
861 Err(LibsyError::Driver(DriverError::Abandoned))
862 ));
863 Ok(())
864 }
865
866 #[test]
869 fn abandoning_and_dropping_a_call_report_different_outcomes() {
870 let abandoned: Result<Response> = Err(DriverError::Abandoned.into());
871 let dropped: Result<Response> = Err(DriverError::ResponseDropped.into());
872 assert_eq!(observability::outcome_value(&abandoned), "abandoned");
873 assert_eq!(observability::outcome_value(&dropped), "error");
874 assert!(observability::is_abandoned(&abandoned));
875 assert!(!observability::is_abandoned(&dropped));
876 }
877
878 #[test]
879 fn target_lookup_returns_the_missing_target() {
880 let error = target_set(&[]).get_target("missing").err();
881 assert!(matches!(
882 error,
883 Some(LibsyError::TargetNotFound { target }) if target == "missing"
884 ));
885 }
886
887 fn streaming_orch(chunks: Vec<LlmResponseChunk>) -> (Arc<dyn Algorithm>, impl Serve) {
890 let algo = orch(target_set(&["stream/model"]));
891 let serve = move |_decision: Decision, _request: Request| {
892 let chunks = chunks.clone();
893 async move {
894 let stream =
895 futures::stream::iter(chunks.into_iter().map(|chunk| Ok(chunk.into()))).boxed();
896 Ok(Response {
897 llm_response: LlmResponse::Stream(stream),
898 metadata: None,
899 })
900 }
901 };
902 (algo, serve)
903 }
904
905 #[tokio::test]
906 async fn run_returns_a_streamed_response_the_caller_aggregates() -> Result<()> {
907 let (orch, serve) = streaming_orch(vec![
910 LlmResponseChunk::MessageStart {
911 id: Some("m1".to_string()),
912 model: Some("stream/model".to_string()),
913 },
914 LlmResponseChunk::TextDelta {
915 index: 0,
916 text: "hel".to_string(),
917 },
918 LlmResponseChunk::TextDelta {
919 index: 0,
920 text: "lo".to_string(),
921 },
922 LlmResponseChunk::MessageStop {
923 reason: Some("stop".to_string()),
924 },
925 ]);
926 let (trace, response) = test_drive(orch, request(), serve).await?;
927 let agg = response
929 .llm_response
930 .into_agg()
931 .await
932 .map_err(|error| LibsyError::external("aggregating response stream", error))?;
933 assert_eq!(completion_text(&agg), "hello");
934 assert_eq!(agg.model.as_deref(), Some("stream/model"));
935 assert_eq!(trace.len(), 1);
936 Ok(())
937 }
938
939 #[tokio::test]
940 async fn aggregating_a_streamed_response_propagates_a_mid_stream_error() -> Result<()> {
941 let (orch, serve) = streaming_orch(vec![
944 LlmResponseChunk::TextDelta {
945 index: 0,
946 text: "partial".to_string(),
947 },
948 LlmResponseChunk::StreamError {
949 message: "upstream exploded".to_string(),
950 },
951 ]);
952 let (_, response) = test_drive(orch, request(), serve).await?;
953 match response.llm_response.into_agg().await {
954 Ok(_) => panic!("expected a mid-stream error, got an aggregate"),
955 Err(err) => {
956 assert!(err.to_string().contains("upstream exploded"));
957 Ok(())
958 }
959 }
960 }
961
962 #[tokio::test]
963 async fn run_offloads_via_promise_then_returns_to_agent() -> Result<()> {
964 let stream = orch(target_set(&["offload/model"])).run_stream(request());
967 tokio::pin!(stream);
968
969 let mut saw_call = false;
970 let mut final_completion = None;
971 while let Some(step) = stream.next().await {
972 match step? {
973 Step::CallModel(call) => {
974 saw_call = true;
975 assert_eq!(call.decision.selected_model_id(), "offload/model");
977 call.respond(Ok(Response {
979 llm_response: LlmResponse::Agg(text_response(
980 None,
981 "fulfilled".to_string(),
982 )),
983 metadata: None,
984 }))?;
985 }
986 Step::Decision(decision) => {
987 assert_eq!(decision.selected_model_id(), "offload/model");
988 }
989 Step::ReturnToAgent(response) => {
990 final_completion = Some(
991 response
992 .llm_response
993 .as_agg()
994 .map(completion_text)
995 .unwrap_or_default(),
996 );
997 }
998 }
999 }
1000
1001 assert!(saw_call, "expected a CallModel step before ReturnToAgent");
1002 assert_eq!(
1003 final_completion.ok_or_else(|| test_error("no ReturnToAgent step"))?,
1004 "fulfilled"
1005 );
1006 Ok(())
1007 }
1008
1009 #[tokio::test]
1010 async fn a_driven_run_returns_the_trace_and_the_final_response() -> Result<()> {
1011 let (trace, response) =
1012 test_drive(orch(target_set(&["direct/model"])), request(), echo()).await?;
1013 assert_eq!(
1015 response
1016 .llm_response
1017 .as_agg()
1018 .map(completion_text)
1019 .unwrap_or_default(),
1020 "direct/model"
1021 );
1022 assert_eq!(trace[0].selected_model_id(), "direct/model");
1023 Ok(())
1024 }
1025
1026 #[tokio::test(flavor = "multi_thread", worker_threads = 12)]
1027 async fn requests_are_processed_in_parallel() -> Result<()> {
1028 use std::time::Duration;
1029 use tokio::sync::Barrier;
1030
1031 const N: usize = 12;
1032
1033 let barrier = Arc::new(Barrier::new(N));
1038 let algo = orch(target_set(&["m"]));
1040
1041 let mut handles = Vec::new();
1042 for _ in 0..N {
1043 let algo = algo.clone();
1044 let barrier = barrier.clone();
1045 let serve = move |decision: Decision, _request: Request| {
1046 let barrier = barrier.clone();
1047 async move {
1048 barrier.wait().await;
1049 Ok(reply(decision.selected_model_id()))
1050 }
1051 };
1052 handles.push(tokio::spawn(async move {
1053 test_drive(algo, request(), serve)
1054 .await
1055 .map(|(_, response)| {
1056 response
1057 .llm_response
1058 .as_agg()
1059 .map(completion_text)
1060 .unwrap_or_default()
1061 })
1062 }));
1063 }
1064
1065 for handle in handles {
1066 let completion = tokio::time::timeout(Duration::from_secs(5), handle)
1068 .await
1069 .map_err(|error| LibsyError::external("waiting for test task", error))?
1070 .map_err(|source| LibsyError::AlgorithmTask { source })??;
1071 assert_eq!(completion, "m");
1072 }
1073 Ok(())
1074 }
1075
1076 #[tokio::test]
1077 async fn offload_error_propagates_back_to_the_algorithm() -> Result<()> {
1078 let stream = orch(target_set(&["offload/model"])).run_stream(request());
1082 tokio::pin!(stream);
1083
1084 let mut saw_error = false;
1085 while let Some(step) = stream.next().await {
1086 match step {
1087 Ok(Step::CallModel(call)) => {
1088 call.respond(Err(test_error("upstream model call failed")))?;
1089 }
1090 Ok(Step::Decision(_)) => {}
1091 Ok(Step::ReturnToAgent(..)) => {
1092 return Err(test_error(
1093 "expected the offload error to propagate, got a response",
1094 ));
1095 }
1096 Err(err) => {
1097 assert!(err.to_string().contains("upstream model call failed"));
1099 saw_error = true;
1100 }
1101 }
1102 }
1103
1104 assert!(saw_error, "expected an error step");
1105 Ok(())
1106 }
1107
1108 #[tokio::test]
1109 async fn dropping_the_stream_cancels_the_algorithm_task() -> Result<()> {
1110 use std::sync::atomic::{AtomicBool, Ordering};
1111 use std::time::Duration;
1112 use tokio::sync::mpsc;
1113
1114 struct DropGuard(Arc<AtomicBool>);
1117 impl Drop for DropGuard {
1118 fn drop(&mut self) {
1119 self.0.store(true, Ordering::SeqCst);
1120 }
1121 }
1122
1123 struct StuckAlgo {
1124 started: mpsc::UnboundedSender<()>,
1125 dropped: Arc<AtomicBool>,
1126 }
1127
1128 #[async_trait]
1129 impl Algorithm for StuckAlgo {
1130 fn name(&self) -> &str {
1131 "stuck"
1132 }
1133
1134 async fn create_run_task(
1135 self: Arc<Self>,
1136 _driver: Driver,
1137 _request: Request,
1138 ) -> Result<Response> {
1139 let _guard = DropGuard(self.dropped.clone());
1140 let _ = self.started.send(());
1141 std::future::pending::<()>().await;
1143 unreachable!()
1144 }
1145 }
1146
1147 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1148 let dropped = Arc::new(AtomicBool::new(false));
1149 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1150 started: started_tx,
1151 dropped: dropped.clone(),
1152 });
1153
1154 let stream = algo.run_stream(request());
1155 started_rx
1156 .recv()
1157 .await
1158 .ok_or_else(|| test_error("task never started"))?;
1159 drop(stream);
1160 tokio::time::sleep(Duration::from_millis(100)).await;
1161
1162 assert!(
1163 dropped.load(Ordering::SeqCst),
1164 "algorithm task was NOT cancelled after dropping the stream"
1165 );
1166 Ok(())
1167 }
1168
1169 #[tokio::test]
1170 async fn create_run_task_panic_surfaces_as_a_stream_error() -> Result<()> {
1171 struct Panicky;
1174
1175 #[async_trait]
1176 impl Algorithm for Panicky {
1177 fn name(&self) -> &str {
1178 "panicky"
1179 }
1180
1181 async fn create_run_task(
1182 self: Arc<Self>,
1183 _driver: Driver,
1184 _request: Request,
1185 ) -> Result<Response> {
1186 panic!("boom");
1187 }
1188 }
1189
1190 let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
1191 let stream = algo.run_stream(request());
1192 tokio::pin!(stream);
1193
1194 let mut saw_error = false;
1195 while let Some(step) = stream.next().await {
1196 match step {
1197 Err(err) => {
1198 assert!(matches!(err, LibsyError::AlgorithmTask { .. }));
1199 saw_error = true;
1200 }
1201 Ok(_) => return Err(test_error("expected the panic to surface as an error step")),
1202 }
1203 }
1204
1205 assert!(saw_error, "expected an error step from the panicked task");
1206 Ok(())
1207 }
1208
1209 #[tokio::test]
1210 async fn run_returns_an_error_when_the_algorithm_task_panics() -> Result<()> {
1211 struct Panicky;
1214
1215 #[async_trait]
1216 impl Algorithm for Panicky {
1217 fn name(&self) -> &str {
1218 "panicky"
1219 }
1220
1221 async fn create_run_task(
1222 self: Arc<Self>,
1223 _driver: Driver,
1224 _request: Request,
1225 ) -> Result<Response> {
1226 panic!("boom");
1227 }
1228 }
1229
1230 let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
1231 match test_drive(algo, request(), echo()).await {
1232 Ok(_) => Err(test_error(
1233 "expected the run to surface the algorithm panic as an error",
1234 )),
1235 Err(err) => {
1236 assert!(matches!(err, LibsyError::AlgorithmTask { .. }));
1237 Ok(())
1238 }
1239 }
1240 }
1241
1242 #[tokio::test]
1243 async fn cancelling_run_cancels_the_algorithm_task() -> Result<()> {
1244 use std::sync::atomic::{AtomicBool, Ordering};
1245 use std::time::Duration;
1246 use tokio::sync::mpsc;
1247
1248 struct DropGuard(Arc<AtomicBool>);
1251 impl Drop for DropGuard {
1252 fn drop(&mut self) {
1253 self.0.store(true, Ordering::SeqCst);
1254 }
1255 }
1256
1257 struct StuckAlgo {
1258 started: mpsc::UnboundedSender<()>,
1259 dropped: Arc<AtomicBool>,
1260 }
1261
1262 #[async_trait]
1263 impl Algorithm for StuckAlgo {
1264 fn name(&self) -> &str {
1265 "stuck"
1266 }
1267
1268 async fn create_run_task(
1269 self: Arc<Self>,
1270 _driver: Driver,
1271 _request: Request,
1272 ) -> Result<Response> {
1273 let _guard = DropGuard(self.dropped.clone());
1274 let _ = self.started.send(());
1275 std::future::pending::<()>().await;
1278 unreachable!()
1279 }
1280 }
1281
1282 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1283 let dropped = Arc::new(AtomicBool::new(false));
1284 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1285 started: started_tx,
1286 dropped: dropped.clone(),
1287 });
1288
1289 let run_task = tokio::spawn(async move { test_drive(algo, request(), echo()).await });
1292 started_rx
1293 .recv()
1294 .await
1295 .ok_or_else(|| test_error("task never started"))?;
1296 run_task.abort();
1297 tokio::time::sleep(Duration::from_millis(100)).await;
1298
1299 assert!(
1300 dropped.load(Ordering::SeqCst),
1301 "algorithm task was NOT cancelled after cancelling run"
1302 );
1303 Ok(())
1304 }
1305
1306 struct Hedge {
1311 winner: LlmTarget,
1312 loser: LlmTarget,
1313 }
1314
1315 #[async_trait]
1316 impl Algorithm for Hedge {
1317 fn name(&self) -> &str {
1318 "hedge"
1319 }
1320
1321 async fn create_run_task(
1322 self: Arc<Self>,
1323 driver: Driver,
1324 request: Request,
1325 ) -> Result<Response> {
1326 let dec_w = test_decision(self.winner.semantic_name.clone());
1327 let dec_l = test_decision(self.loser.semantic_name.clone());
1328 let win = driver.call_model(request.clone(), dec_w);
1329 let lose = driver.call_model(request, dec_l);
1330 tokio::select! {
1332 res = win => res,
1333 res = lose => res,
1334 }
1335 }
1336 }
1337
1338 fn hedge(loser_delay: Option<std::time::Duration>) -> (Arc<dyn Algorithm>, impl Serve) {
1342 let started = Arc::new(tokio::sync::Notify::new());
1343 let algo = Arc::new(Hedge {
1344 winner: LlmTarget {
1345 semantic_name: "winner".to_string(),
1346 },
1347 loser: LlmTarget {
1348 semantic_name: "loser".to_string(),
1349 },
1350 });
1351 let serve = move |decision: Decision, _request: Request| {
1352 let started = started.clone();
1353 async move {
1354 if decision.selected_model_id() == "loser" {
1355 started.notify_one();
1356 match loser_delay {
1357 Some(delay) => tokio::time::sleep(delay).await,
1358 None => std::future::pending::<()>().await,
1359 }
1360 } else {
1361 started.notified().await;
1362 }
1363 Ok(reply(decision.selected_model_id()))
1364 }
1365 };
1366 (algo, serve)
1367 }
1368
1369 #[tokio::test]
1370 async fn run_returns_the_winner_without_a_late_loser_overwriting_it() -> Result<()> {
1371 let (algo, serve) = hedge(Some(std::time::Duration::from_millis(50)));
1374 let (_trace, response) = test_drive(algo, request(), serve).await?;
1375 assert_eq!(
1376 response
1377 .llm_response
1378 .as_agg()
1379 .map(completion_text)
1380 .unwrap_or_default(),
1381 "winner"
1382 );
1383 Ok(())
1384 }
1385
1386 #[tokio::test]
1387 async fn run_returns_the_winner_without_hanging_on_a_pending_loser() -> Result<()> {
1388 let (algo, serve) = hedge(None);
1391 let run = test_drive(algo, request(), serve);
1392 let (_trace, response) = tokio::time::timeout(std::time::Duration::from_secs(1), run)
1393 .await
1394 .map_err(|error| LibsyError::external("waiting for pending loser", error))??;
1395 assert_eq!(
1396 response
1397 .llm_response
1398 .as_agg()
1399 .map(completion_text)
1400 .unwrap_or_default(),
1401 "winner"
1402 );
1403 Ok(())
1404 }
1405
1406 #[tokio::test]
1407 async fn run_surfaces_a_terminal_error_with_many_calls_in_flight() -> Result<()> {
1408 use std::sync::atomic::{AtomicUsize, Ordering};
1409
1410 const N: usize = 10;
1413
1414 struct FanOutThenError {
1417 all_started: Arc<tokio::sync::Notify>,
1418 n: usize,
1419 }
1420
1421 #[async_trait]
1422 impl Algorithm for FanOutThenError {
1423 fn name(&self) -> &str {
1424 "fan_out_then_error"
1425 }
1426
1427 async fn create_run_task(
1428 self: Arc<Self>,
1429 driver: Driver,
1430 request: Request,
1431 ) -> Result<Response> {
1432 let offloads = futures::future::join_all((0..self.n).map(|i| {
1433 let decision = test_decision(format!("m{i}"));
1434 driver.call_model(request.clone(), decision)
1435 }));
1436 tokio::select! {
1437 _ = offloads => Err(test_error("offloads unexpectedly completed")),
1438 _ = self.all_started.notified() => {
1439 Err(test_error("terminal error while calls pending"))
1440 }
1441 }
1442 }
1443 }
1444
1445 let all_started = Arc::new(tokio::sync::Notify::new());
1446 let algo: Arc<dyn Algorithm> = Arc::new(FanOutThenError {
1447 all_started: all_started.clone(),
1448 n: N,
1449 });
1450
1451 let started = Arc::new(AtomicUsize::new(0));
1453 let serve = move |_decision: Decision, _request: Request| {
1454 let started = started.clone();
1455 let all_started = all_started.clone();
1456 async move {
1457 if started.fetch_add(1, Ordering::SeqCst) + 1 == N {
1458 all_started.notify_one();
1459 }
1460 std::future::pending::<ServeResult>().await
1461 }
1462 };
1463
1464 let run = test_drive(algo, request(), serve);
1467 let result = tokio::time::timeout(std::time::Duration::from_millis(500), run)
1468 .await
1469 .map_err(|error| {
1470 LibsyError::external("waiting for terminal error with full call cap", error)
1471 })?;
1472 match result {
1473 Ok(_) => Err(test_error("expected the terminal error, got a response")),
1474 Err(err) => {
1475 assert!(
1476 err.to_string()
1477 .contains("terminal error while calls pending")
1478 );
1479 Ok(())
1480 }
1481 }
1482 }
1483}