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 tracing::Instrument;
20
21use switchyard_protocol::{
29 Context, Decision, LlmClientError, Request, Response, RoutingFallbackReason, Signals,
30};
31
32use super::driver::{DriverRequest, DriverStep, TypeErasedDriver};
33use crate::{DriverError, LibsyError, Result, observability};
34
35pub type StepStream = Pin<Box<dyn Stream<Item = Result<Step>> + Send>>;
39
40#[derive(Clone)]
49pub struct RoutedRequest {
50 pub request: Request,
52 pub decision: Arc<dyn Decision>,
54 pub ctx: Context,
57}
58
59pub struct CallLlmRequest {
68 inner: DriverRequest,
69 routed: RoutedRequest,
70}
71
72impl CallLlmRequest {
73 fn new(inner: DriverRequest) -> Self {
76 let routed = match inner.request::<RoutedRequest>() {
79 Ok(routed) => routed.clone(),
80 Err(_) => unreachable!("CallLlmRequest payload is always a RoutedRequest"),
81 };
82 Self { inner, routed }
83 }
84
85 pub fn get_routed(&self) -> &RoutedRequest {
88 &self.routed
89 }
90
91 pub fn get_request(&self) -> &Request {
93 &self.get_routed().request
94 }
95
96 pub fn get_decision(&self) -> &dyn Decision {
98 self.get_routed().decision.as_ref()
99 }
100
101 pub fn respond(self, result: Result<Response>) -> Result<()> {
105 self.inner.respond::<Response>(result)
106 }
107}
108
109#[derive(Clone)]
116pub struct Driver {
117 driver: TypeErasedDriver,
118}
119
120impl Driver {
121 pub(crate) fn new() -> Self {
124 Self {
125 driver: TypeErasedDriver::new(),
126 }
127 }
128
129 #[tracing::instrument(
138 target = "libsy",
139 name = "libsy.llm_call",
140 skip_all,
141 fields(
142 algorithm = observability::algorithm_label(&routed.ctx),
143 selected_model = routed.decision.selected_model(),
144 openinference.span.kind = "CHAIN",
145 outcome = tracing::field::Empty,
146 error = tracing::field::Empty,
147 input_tokens = tracing::field::Empty,
148 output_tokens = tracing::field::Empty,
149 total_tokens = tracing::field::Empty,
150 reasoning_tokens = tracing::field::Empty,
151 )
152 )]
153 pub async fn call_llm(&self, routed: RoutedRequest) -> Result<Response> {
154 let algorithm = observability::algorithm_label(&routed.ctx).to_string();
155 let selected_model = routed.decision.selected_model().to_string();
156 let tier = routed.decision.routing_tier().map(str::to_string);
157 let is_routed = routed.decision.is_routed_call();
158 let started = Instant::now();
159 let result = self
160 .driver
161 .fulfill_request::<RoutedRequest, Response>(routed.ctx.clone(), routed)
162 .await;
163 let elapsed = started.elapsed();
164 observability::record_llm_call(
165 &algorithm,
166 &selected_model,
167 tier.as_deref(),
168 is_routed,
169 elapsed,
170 &result,
171 &tracing::Span::current(),
172 );
173 result
174 }
175
176 pub async fn info(&self, ctx: Context, decision: Arc<dyn Decision>) -> Result<()> {
180 self.driver.info(ctx.clone(), decision.clone()).await?;
181 observability::record_decision(&ctx, decision.as_ref());
182 Ok(())
183 }
184
185 pub(crate) async fn finish(&self, ctx: Context, result: Result<Response>) -> Result<()> {
189 match result {
190 Ok(response) => self.driver.done(ctx, response).await,
191 Err(err) => self.driver.fail(ctx, err).await,
192 }
193 }
194
195 pub(crate) fn stream(&self) -> impl Stream<Item = Result<Step>> + use<> {
199 self.driver.stream().map(|item| match item? {
200 DriverStep::Request(req) => Ok(Step::CallLlm(Box::new(CallLlmRequest::new(req)))),
201 DriverStep::Info(payload) => payload
202 .downcast::<Arc<dyn Decision>>()
203 .map(|decision| Step::Decision(*decision))
204 .map_err(|_| {
205 DriverError::TypeMismatch {
206 expected: "Arc<dyn Decision>",
207 }
208 .into()
209 }),
210 DriverStep::Done(payload) => payload
211 .downcast::<Response>()
212 .map(Step::ReturnToAgent)
213 .map_err(|_| {
214 DriverError::TypeMismatch {
215 expected: "Response",
216 }
217 .into()
218 }),
219 })
220 }
221}
222
223impl Default for Driver {
224 fn default() -> Self {
225 Self::new()
226 }
227}
228
229pub enum Step {
231 CallLlm(Box<CallLlmRequest>),
234 Decision(Arc<dyn Decision>),
237 ReturnToAgent(Box<Response>),
239}
240
241pub async fn drive<F, Fut>(
254 algorithm: Arc<dyn Algorithm>,
255 ctx: Context,
256 request: Request,
257 serve: F,
258) -> Result<(Vec<Arc<dyn Decision>>, Response)>
259where
260 F: Fn(CallLlmRequest) -> Fut,
261 Fut: Future<Output = Result<()>>,
262{
263 let stream = algorithm.run_stream(ctx, request);
264 tokio::pin!(stream);
265
266 let mut trace: Vec<Arc<dyn Decision>> = Vec::new();
267 let mut in_flight = futures::stream::FuturesUnordered::new();
268 let mut final_response: Option<Response> = None;
269
270 loop {
271 tokio::select! {
272 Some(result) = in_flight.next() => match result {
273 Ok(()) => {}, Err(err) => return Err(err), },
276 step = stream.next() => {
277 match step {
278 None => break, Some(item) => match item? {
280 Step::CallLlm(call) => in_flight.push(serve(*call)),
281 Step::Decision(decision) => trace.push(decision),
282 Step::ReturnToAgent(response) => {
283 final_response = Some(*response);
284 break;
285 }
286 }
287 }
288 },
289 }
290 }
291 final_response
292 .map(|response| (trace, response))
293 .ok_or(LibsyError::MissingFinalResponse)
294}
295
296struct AbortOnDrop(tokio::task::AbortHandle);
298
299impl Drop for AbortOnDrop {
300 fn drop(&mut self) {
301 self.0.abort();
302 }
303}
304
305#[derive(Clone)]
309pub struct LlmTarget {
310 pub semantic_name: String,
314}
315
316#[derive(Clone)]
320pub struct LlmTargetSet {
321 targets: Vec<LlmTarget>,
322}
323
324impl LlmTargetSet {
325 pub fn new(targets: Vec<LlmTarget>) -> Self {
327 Self { targets }
328 }
329
330 pub fn targets(&self) -> &[LlmTarget] {
332 &self.targets
333 }
334
335 pub fn get_target(&self, name: &str) -> Result<LlmTarget> {
337 self.targets
338 .iter()
339 .find(|t| t.semantic_name == name)
340 .cloned()
341 .ok_or_else(|| LibsyError::TargetNotFound {
342 target: name.to_string(),
343 })
344 }
345
346 pub fn resolve_target(&self, name: &str, ctx: &Context) -> Result<LlmTarget> {
349 let target = self.get_target(name)?;
350 if !ctx.is_excluded(&target.semantic_name) {
351 return Ok(target);
352 }
353 self.targets
354 .iter()
355 .find(|t| !ctx.is_excluded(&t.semantic_name))
356 .cloned()
357 .ok_or(LibsyError::AllTargetsExcluded)
358 }
359}
360
361#[derive(Clone, Hash, PartialEq, Eq)]
365pub(crate) enum RoutingIdentity {
366 Session(String),
368 Subagent { session: String, agent: String },
370}
371
372impl RoutingIdentity {
373 pub(crate) fn from_request(request: &Request) -> Option<Self> {
378 let metadata = request.metadata.as_ref()?;
379 let session = metadata.session_id.as_deref().filter(|id| !id.is_empty())?;
380 if metadata.is_subagent {
381 let agent = metadata.agent_id.as_deref().filter(|id| !id.is_empty())?;
382 Some(Self::Subagent {
383 session: session.to_string(),
384 agent: agent.to_string(),
385 })
386 } else {
387 Some(Self::Session(session.to_string()))
388 }
389 }
390
391 fn session(&self) -> &str {
393 match self {
394 Self::Session(session) | Self::Subagent { session, .. } => session,
395 }
396 }
397}
398
399const MAX_EVICTION_IDENTITIES: usize = 1_024;
402
403#[derive(Default)]
409pub(crate) struct SessionEvictions {
410 by_identity: Mutex<HashMap<RoutingIdentity, HashSet<String>>>,
411}
412
413impl SessionEvictions {
414 pub(crate) fn remove_session(&self, session: &str) {
416 self.by_identity
417 .lock()
418 .retain(|identity, _| identity.session() != session);
419 }
420
421 fn evicted_for(&self, identity: Option<&RoutingIdentity>) -> Vec<String> {
423 let Some(identity) = identity else {
424 return Vec::new();
425 };
426 self.by_identity
427 .lock()
428 .get(identity)
429 .map(|targets| targets.iter().cloned().collect())
430 .unwrap_or_default()
431 }
432
433 fn record(&self, identity: Option<&RoutingIdentity>, target: &str) {
436 let Some(identity) = identity else { return };
437 let mut histories = self.by_identity.lock();
438 if histories.len() >= MAX_EVICTION_IDENTITIES
439 && !histories.contains_key(identity)
440 && let Some(oldest) = histories.keys().next().cloned()
441 {
442 histories.remove(&oldest);
443 }
444 histories
445 .entry(identity.clone())
446 .or_default()
447 .insert(target.to_string());
448 }
449}
450
451fn eligible_targets(targets: &LlmTargetSet, ctx: &Context) -> usize {
453 targets
454 .targets()
455 .iter()
456 .filter(|t| !ctx.is_excluded(&t.semantic_name))
457 .count()
458}
459
460pub(crate) fn exclude_evicted(
463 ctx: &mut Context,
464 targets: &LlmTargetSet,
465 evictions: &SessionEvictions,
466 identity: Option<&RoutingIdentity>,
467) {
468 for target in evictions.evicted_for(identity) {
469 if eligible_targets(targets, ctx) <= 1 {
472 break;
473 }
474 ctx.exclude_target(target);
475 }
476}
477
478fn classify_fallback(error: &LibsyError) -> Option<(&str, RoutingFallbackReason)> {
480 let LibsyError::ClientCall { target, source } = error else {
481 return None;
482 };
483 let reason = match source {
484 LlmClientError::ContextWindowExceeded { .. } => RoutingFallbackReason::ContextWindow,
485 LlmClientError::Transport { .. } | LlmClientError::Timeout { .. } => {
486 RoutingFallbackReason::Unavailable
487 }
488 LlmClientError::UpstreamHttp { status, .. }
489 if matches!(*status, 403 | 408 | 429) || (500..=599).contains(status) =>
490 {
491 RoutingFallbackReason::Unavailable
492 }
493 _ => return None,
494 };
495 Some((target, reason))
496}
497
498#[allow(clippy::too_many_arguments)]
506pub(crate) async fn call_llm_with_fallback(
507 mut ctx: Context,
508 driver: &Driver,
509 targets: &LlmTargetSet,
510 mut target: LlmTarget,
511 mut decision: Arc<dyn Decision>,
512 request: Request,
513 identity: Option<&RoutingIdentity>,
514 evictions: &SessionEvictions,
515 target_unavailable: impl Fn(&Request, &str),
516 fallback_decision: impl Fn(&LlmTarget, &LlmTarget, RoutingFallbackReason) -> Arc<dyn Decision>,
517) -> Result<Response> {
518 loop {
519 let result = driver
520 .call_llm(RoutedRequest {
521 request: request.clone(),
522 decision: decision.clone(),
523 ctx: ctx.clone(),
524 })
525 .await;
526 let Err(error) = result else { return result };
527 let Some((failed, reason)) = classify_fallback(&error) else {
528 return Err(error);
529 };
530 if !ctx.exclude_target(failed) {
533 return Err(error);
534 }
535 match reason {
536 RoutingFallbackReason::ContextWindow => evictions.record(identity, failed),
537 RoutingFallbackReason::Unavailable => target_unavailable(&request, failed),
538 }
539 let Ok(next) = targets.resolve_target(&target.semantic_name, &ctx) else {
540 return Err(error);
541 };
542 decision = fallback_decision(&target, &next, reason);
543 target = next;
544 driver.info(ctx.clone(), decision.clone()).await?;
545 }
546}
547
548#[async_trait]
570pub trait Algorithm: Send + Sync + 'static {
571 fn name(&self) -> &str;
575
576 async fn create_run_task(
582 self: Arc<Self>,
583 ctx: Context,
584 driver: Driver,
585 request: Request,
586 ) -> Result<Response>;
587
588 #[allow(unused_variables)]
592 async fn process_signals(self: Arc<Self>, signals: Signals) -> Result<()> {
593 Ok(())
594 }
595
596 fn run_stream(self: Arc<Self>, ctx: Context, request: Request) -> StepStream {
605 let mut ctx = ctx;
608 ctx.values.insert(
609 observability::ALGORITHM_KEY.to_string(),
610 self.name().to_string(),
611 );
612 let driver = Driver::new();
613 let task_driver = driver.clone();
614 let task_ctx = ctx.clone();
615 let stream = task_driver.stream();
616 let span = observability::run_span(self.name(), &request);
620 let handle = tokio::spawn(
621 async move {
622 observability::observe_run(
623 task_ctx.clone(),
624 self.create_run_task(task_ctx, task_driver, request),
625 )
626 .await
627 }
628 .instrument(span),
629 );
630 let abort_guard = AbortOnDrop(handle.abort_handle());
632
633 let finish_driver = driver.clone();
634 let finish_ctx = ctx;
635 let tail: StepStream = Box::pin(
636 futures::stream::once(async move {
637 let result = match handle.await {
638 Ok(response) => response,
639 Err(source) => Err(LibsyError::AlgorithmTask { source }),
640 };
641 finish_driver.finish(finish_ctx, result).await
642 })
643 .filter_map(|finish_result| async move { finish_result.err().map(Err) }),
644 );
645
646 let stream: StepStream = Box::pin(stream);
647 Box::pin(futures::stream::select(stream, tail).map(move |step| {
648 let _keep_alive = &abort_guard;
650 step
651 }))
652 }
653}
654
655#[cfg(test)]
656mod tests {
657 use super::*;
658 use crate::core::testing::{Serve, ServeResult, echo, reply, test_drive};
659 use futures::StreamExt;
660 use switchyard_protocol::{
661 LlmResponse, LlmResponseChunk, completion_text, text_request, text_response,
662 };
663
664 #[derive(Debug, thiserror::Error)]
665 #[error("{0}")]
666 struct TestError(&'static str);
667
668 fn test_error(message: &'static str) -> LibsyError {
669 LibsyError::external("test", TestError(message))
670 }
671
672 fn classified_client_error(source: LlmClientError) -> Option<RoutingFallbackReason> {
673 classify_fallback(&LibsyError::client_call("target", source)).map(|(_, reason)| reason)
674 }
675
676 #[test]
677 fn route_fallback_only_accepts_context_and_unavailable_failures() {
678 assert_eq!(
679 classified_client_error(LlmClientError::ContextWindowExceeded {
680 model: "target".to_string(),
681 message: "too long".to_string(),
682 }),
683 Some(RoutingFallbackReason::ContextWindow)
684 );
685 for source in [
686 LlmClientError::Transport {
687 source: Box::new(std::io::Error::other("connection failed")),
688 },
689 LlmClientError::Timeout {
690 source: Box::new(std::io::Error::other("request timed out")),
691 },
692 ] {
693 assert_eq!(
694 classified_client_error(source),
695 Some(RoutingFallbackReason::Unavailable)
696 );
697 }
698 for (status, expected) in [
699 (400, None),
700 (401, None),
701 (403, Some(RoutingFallbackReason::Unavailable)),
702 (404, None),
703 (408, Some(RoutingFallbackReason::Unavailable)),
704 (409, None),
705 (429, Some(RoutingFallbackReason::Unavailable)),
706 (499, None),
707 (500, Some(RoutingFallbackReason::Unavailable)),
708 (599, Some(RoutingFallbackReason::Unavailable)),
709 (600, None),
710 ] {
711 assert_eq!(
712 classified_client_error(LlmClientError::UpstreamHttp {
713 status,
714 body: "failed".to_string(),
715 }),
716 expected
717 );
718 }
719 assert_eq!(
720 classified_client_error(LlmClientError::InvalidResponse {
721 source: Box::new(std::io::Error::other("invalid response")),
722 }),
723 None
724 );
725 }
726
727 struct TestDecision {
730 model: String,
731 }
732
733 impl Decision for TestDecision {
734 fn selected_model(&self) -> &str {
735 &self.model
736 }
737 fn reasoning(&self) -> Option<&str> {
738 None
739 }
740 fn as_any(&self) -> &dyn std::any::Any {
741 self
742 }
743 }
744
745 struct TestAlgo {
746 target_set: LlmTargetSet,
747 }
748
749 #[async_trait]
750 impl Algorithm for TestAlgo {
751 fn name(&self) -> &str {
752 "test"
753 }
754
755 async fn create_run_task(
756 self: Arc<Self>,
757 ctx: Context,
758 driver: Driver,
759 request: Request,
760 ) -> Result<Response> {
761 let target = self
762 .target_set
763 .targets()
764 .first()
765 .ok_or(LibsyError::NoTargets)?
766 .clone();
767 let decision: Arc<dyn Decision> = Arc::new(TestDecision {
768 model: target.semantic_name.clone(),
769 });
770 driver.info(ctx.clone(), decision.clone()).await?;
771 driver
772 .call_llm(RoutedRequest {
773 request,
774 decision,
775 ctx,
776 })
777 .await
778 }
779 }
780
781 fn orch(target_set: LlmTargetSet) -> Arc<dyn Algorithm> {
783 Arc::new(TestAlgo { target_set })
784 }
785
786 fn request() -> Request {
787 Request {
788 llm_request: text_request(Some("auto".to_string()), "hi".to_string()),
789 raw_request: None,
790 metadata: None,
791 }
792 }
793
794 fn target_set(names: &[&str]) -> LlmTargetSet {
795 let targets = names
796 .iter()
797 .map(|name| LlmTarget {
798 semantic_name: name.to_string(),
799 })
800 .collect();
801 LlmTargetSet::new(targets)
802 }
803
804 #[test]
805 fn target_lookup_returns_the_missing_target() {
806 let error = target_set(&[]).get_target("missing").err();
807 assert!(matches!(
808 error,
809 Some(LibsyError::TargetNotFound { target }) if target == "missing"
810 ));
811 }
812
813 fn streaming_orch(chunks: Vec<LlmResponseChunk>) -> (Arc<dyn Algorithm>, impl Serve) {
816 let algo = orch(target_set(&["stream/model"]));
817 let serve = move |_decision: Arc<dyn Decision>, _request: Request| {
818 let chunks = chunks.clone();
819 async move {
820 let stream =
821 futures::stream::iter(chunks.into_iter().map(|chunk| Ok(chunk.into()))).boxed();
822 Ok(Response {
823 llm_response: LlmResponse::Stream(stream),
824 metadata: None,
825 })
826 }
827 };
828 (algo, serve)
829 }
830
831 #[tokio::test]
832 async fn run_returns_a_streamed_response_the_caller_aggregates() -> Result<()> {
833 let (orch, serve) = streaming_orch(vec![
836 LlmResponseChunk::MessageStart {
837 id: Some("m1".to_string()),
838 model: Some("stream/model".to_string()),
839 },
840 LlmResponseChunk::TextDelta {
841 index: 0,
842 text: "hel".to_string(),
843 },
844 LlmResponseChunk::TextDelta {
845 index: 0,
846 text: "lo".to_string(),
847 },
848 LlmResponseChunk::MessageStop {
849 reason: Some("stop".to_string()),
850 },
851 ]);
852 let (trace, response) = test_drive(orch, Context::default(), request(), serve).await?;
853 let agg = response
855 .llm_response
856 .into_agg()
857 .await
858 .map_err(|error| LibsyError::external("aggregating response stream", error))?;
859 assert_eq!(completion_text(&agg), "hello");
860 assert_eq!(agg.model.as_deref(), Some("stream/model"));
861 assert_eq!(trace.len(), 1);
862 Ok(())
863 }
864
865 #[tokio::test]
866 async fn aggregating_a_streamed_response_propagates_a_mid_stream_error() -> Result<()> {
867 let (orch, serve) = streaming_orch(vec![
870 LlmResponseChunk::TextDelta {
871 index: 0,
872 text: "partial".to_string(),
873 },
874 LlmResponseChunk::StreamError {
875 message: "upstream exploded".to_string(),
876 },
877 ]);
878 let (_, response) = test_drive(orch, Context::default(), request(), serve).await?;
879 match response.llm_response.into_agg().await {
880 Ok(_) => panic!("expected a mid-stream error, got an aggregate"),
881 Err(err) => {
882 assert!(err.to_string().contains("upstream exploded"));
883 Ok(())
884 }
885 }
886 }
887
888 #[tokio::test]
889 async fn run_offloads_via_promise_then_returns_to_agent() -> Result<()> {
890 let stream = orch(target_set(&["offload/model"])).run_stream(Context::default(), request());
893 tokio::pin!(stream);
894
895 let mut saw_call = false;
896 let mut final_completion = None;
897 while let Some(step) = stream.next().await {
898 match step? {
899 Step::CallLlm(call) => {
900 saw_call = true;
901 assert_eq!(call.get_decision().selected_model(), "offload/model");
903 call.respond(Ok(Response {
905 llm_response: LlmResponse::Agg(text_response(
906 None,
907 "fulfilled".to_string(),
908 )),
909 metadata: None,
910 }))?;
911 }
912 Step::Decision(decision) => {
913 assert_eq!(decision.selected_model(), "offload/model");
914 }
915 Step::ReturnToAgent(response) => {
916 final_completion = Some(
917 response
918 .llm_response
919 .as_agg()
920 .map(completion_text)
921 .unwrap_or_default(),
922 );
923 }
924 }
925 }
926
927 assert!(saw_call, "expected a CallLlm step before ReturnToAgent");
928 assert_eq!(
929 final_completion.ok_or_else(|| test_error("no ReturnToAgent step"))?,
930 "fulfilled"
931 );
932 Ok(())
933 }
934
935 #[tokio::test]
936 async fn a_driven_run_returns_the_trace_and_the_final_response() -> Result<()> {
937 let (trace, response) = test_drive(
938 orch(target_set(&["direct/model"])),
939 Context::default(),
940 request(),
941 echo(),
942 )
943 .await?;
944 assert_eq!(
946 response
947 .llm_response
948 .as_agg()
949 .map(completion_text)
950 .unwrap_or_default(),
951 "direct/model"
952 );
953 assert_eq!(trace[0].selected_model(), "direct/model");
954 Ok(())
955 }
956
957 #[tokio::test(flavor = "multi_thread", worker_threads = 12)]
958 async fn requests_are_processed_in_parallel() -> Result<()> {
959 use std::time::Duration;
960 use tokio::sync::Barrier;
961
962 const N: usize = 12;
963
964 let barrier = Arc::new(Barrier::new(N));
969 let algo = orch(target_set(&["m"]));
971
972 let mut handles = Vec::new();
973 for _ in 0..N {
974 let algo = algo.clone();
975 let barrier = barrier.clone();
976 let serve = move |decision: Arc<dyn Decision>, _request: Request| {
977 let barrier = barrier.clone();
978 async move {
979 barrier.wait().await;
980 Ok(reply(decision.selected_model()))
981 }
982 };
983 handles.push(tokio::spawn(async move {
984 test_drive(algo, Context::default(), request(), serve)
985 .await
986 .map(|(_, response)| {
987 response
988 .llm_response
989 .as_agg()
990 .map(completion_text)
991 .unwrap_or_default()
992 })
993 }));
994 }
995
996 for handle in handles {
997 let completion = tokio::time::timeout(Duration::from_secs(5), handle)
999 .await
1000 .map_err(|error| LibsyError::external("waiting for test task", error))?
1001 .map_err(|source| LibsyError::AlgorithmTask { source })??;
1002 assert_eq!(completion, "m");
1003 }
1004 Ok(())
1005 }
1006
1007 #[tokio::test]
1008 async fn offload_error_propagates_back_to_the_algorithm() -> Result<()> {
1009 let stream = orch(target_set(&["offload/model"])).run_stream(Context::default(), request());
1013 tokio::pin!(stream);
1014
1015 let mut saw_error = false;
1016 while let Some(step) = stream.next().await {
1017 match step {
1018 Ok(Step::CallLlm(call)) => {
1019 call.respond(Err(test_error("upstream model call failed")))?;
1020 }
1021 Ok(Step::Decision(_)) => {}
1022 Ok(Step::ReturnToAgent(..)) => {
1023 return Err(test_error(
1024 "expected the offload error to propagate, got a response",
1025 ));
1026 }
1027 Err(err) => {
1028 assert!(err.to_string().contains("upstream model call failed"));
1030 saw_error = true;
1031 }
1032 }
1033 }
1034
1035 assert!(saw_error, "expected an error step");
1036 Ok(())
1037 }
1038
1039 #[tokio::test]
1040 async fn dropping_the_stream_cancels_the_algorithm_task() -> Result<()> {
1041 use std::sync::atomic::{AtomicBool, Ordering};
1042 use std::time::Duration;
1043 use tokio::sync::mpsc;
1044
1045 struct DropGuard(Arc<AtomicBool>);
1048 impl Drop for DropGuard {
1049 fn drop(&mut self) {
1050 self.0.store(true, Ordering::SeqCst);
1051 }
1052 }
1053
1054 struct StuckAlgo {
1055 started: mpsc::UnboundedSender<()>,
1056 dropped: Arc<AtomicBool>,
1057 }
1058
1059 #[async_trait]
1060 impl Algorithm for StuckAlgo {
1061 fn name(&self) -> &str {
1062 "stuck"
1063 }
1064
1065 async fn create_run_task(
1066 self: Arc<Self>,
1067 _ctx: Context,
1068 _driver: Driver,
1069 _request: Request,
1070 ) -> Result<Response> {
1071 let _guard = DropGuard(self.dropped.clone());
1072 let _ = self.started.send(());
1073 std::future::pending::<()>().await;
1075 unreachable!()
1076 }
1077 }
1078
1079 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1080 let dropped = Arc::new(AtomicBool::new(false));
1081 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1082 started: started_tx,
1083 dropped: dropped.clone(),
1084 });
1085
1086 let stream = algo.run_stream(Context::default(), request());
1087 started_rx
1088 .recv()
1089 .await
1090 .ok_or_else(|| test_error("task never started"))?;
1091 drop(stream);
1092 tokio::time::sleep(Duration::from_millis(100)).await;
1093
1094 assert!(
1095 dropped.load(Ordering::SeqCst),
1096 "algorithm task was NOT cancelled after dropping the stream"
1097 );
1098 Ok(())
1099 }
1100
1101 #[tokio::test]
1102 async fn create_run_task_panic_surfaces_as_a_stream_error() -> Result<()> {
1103 struct Panicky;
1106
1107 #[async_trait]
1108 impl Algorithm for Panicky {
1109 fn name(&self) -> &str {
1110 "panicky"
1111 }
1112
1113 async fn create_run_task(
1114 self: Arc<Self>,
1115 _ctx: Context,
1116 _driver: Driver,
1117 _request: Request,
1118 ) -> Result<Response> {
1119 panic!("boom");
1120 }
1121 }
1122
1123 let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
1124 let stream = algo.run_stream(Context::default(), request());
1125 tokio::pin!(stream);
1126
1127 let mut saw_error = false;
1128 while let Some(step) = stream.next().await {
1129 match step {
1130 Err(err) => {
1131 assert!(matches!(err, LibsyError::AlgorithmTask { .. }));
1132 saw_error = true;
1133 }
1134 Ok(_) => return Err(test_error("expected the panic to surface as an error step")),
1135 }
1136 }
1137
1138 assert!(saw_error, "expected an error step from the panicked task");
1139 Ok(())
1140 }
1141
1142 #[tokio::test]
1143 async fn run_returns_an_error_when_the_algorithm_task_panics() -> Result<()> {
1144 struct Panicky;
1147
1148 #[async_trait]
1149 impl Algorithm for Panicky {
1150 fn name(&self) -> &str {
1151 "panicky"
1152 }
1153
1154 async fn create_run_task(
1155 self: Arc<Self>,
1156 _ctx: Context,
1157 _driver: Driver,
1158 _request: Request,
1159 ) -> Result<Response> {
1160 panic!("boom");
1161 }
1162 }
1163
1164 let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
1165 match test_drive(algo, Context::default(), request(), echo()).await {
1166 Ok(_) => Err(test_error(
1167 "expected the run to surface the algorithm panic as an error",
1168 )),
1169 Err(err) => {
1170 assert!(matches!(err, LibsyError::AlgorithmTask { .. }));
1171 Ok(())
1172 }
1173 }
1174 }
1175
1176 #[tokio::test]
1177 async fn cancelling_run_cancels_the_algorithm_task() -> Result<()> {
1178 use std::sync::atomic::{AtomicBool, Ordering};
1179 use std::time::Duration;
1180 use tokio::sync::mpsc;
1181
1182 struct DropGuard(Arc<AtomicBool>);
1185 impl Drop for DropGuard {
1186 fn drop(&mut self) {
1187 self.0.store(true, Ordering::SeqCst);
1188 }
1189 }
1190
1191 struct StuckAlgo {
1192 started: mpsc::UnboundedSender<()>,
1193 dropped: Arc<AtomicBool>,
1194 }
1195
1196 #[async_trait]
1197 impl Algorithm for StuckAlgo {
1198 fn name(&self) -> &str {
1199 "stuck"
1200 }
1201
1202 async fn create_run_task(
1203 self: Arc<Self>,
1204 _ctx: Context,
1205 _driver: Driver,
1206 _request: Request,
1207 ) -> Result<Response> {
1208 let _guard = DropGuard(self.dropped.clone());
1209 let _ = self.started.send(());
1210 std::future::pending::<()>().await;
1213 unreachable!()
1214 }
1215 }
1216
1217 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1218 let dropped = Arc::new(AtomicBool::new(false));
1219 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1220 started: started_tx,
1221 dropped: dropped.clone(),
1222 });
1223
1224 let run_task =
1227 tokio::spawn(
1228 async move { test_drive(algo, Context::default(), request(), echo()).await },
1229 );
1230 started_rx
1231 .recv()
1232 .await
1233 .ok_or_else(|| test_error("task never started"))?;
1234 run_task.abort();
1235 tokio::time::sleep(Duration::from_millis(100)).await;
1236
1237 assert!(
1238 dropped.load(Ordering::SeqCst),
1239 "algorithm task was NOT cancelled after cancelling run"
1240 );
1241 Ok(())
1242 }
1243
1244 struct Hedge {
1249 winner: LlmTarget,
1250 loser: LlmTarget,
1251 }
1252
1253 #[async_trait]
1254 impl Algorithm for Hedge {
1255 fn name(&self) -> &str {
1256 "hedge"
1257 }
1258
1259 async fn create_run_task(
1260 self: Arc<Self>,
1261 ctx: Context,
1262 driver: Driver,
1263 request: Request,
1264 ) -> Result<Response> {
1265 let dec_w: Arc<dyn Decision> = Arc::new(TestDecision {
1266 model: self.winner.semantic_name.clone(),
1267 });
1268 let dec_l: Arc<dyn Decision> = Arc::new(TestDecision {
1269 model: self.loser.semantic_name.clone(),
1270 });
1271 let win = driver.call_llm(RoutedRequest {
1272 request: request.clone(),
1273 decision: dec_w,
1274 ctx: ctx.clone(),
1275 });
1276 let lose = driver.call_llm(RoutedRequest {
1277 request,
1278 decision: dec_l,
1279 ctx,
1280 });
1281 tokio::select! {
1283 res = win => res,
1284 res = lose => res,
1285 }
1286 }
1287 }
1288
1289 fn hedge(loser_delay: Option<std::time::Duration>) -> (Arc<dyn Algorithm>, impl Serve) {
1293 let started = Arc::new(tokio::sync::Notify::new());
1294 let algo = Arc::new(Hedge {
1295 winner: LlmTarget {
1296 semantic_name: "winner".to_string(),
1297 },
1298 loser: LlmTarget {
1299 semantic_name: "loser".to_string(),
1300 },
1301 });
1302 let serve = move |decision: Arc<dyn Decision>, _request: Request| {
1303 let started = started.clone();
1304 async move {
1305 if decision.selected_model() == "loser" {
1306 started.notify_one();
1307 match loser_delay {
1308 Some(delay) => tokio::time::sleep(delay).await,
1309 None => std::future::pending::<()>().await,
1310 }
1311 } else {
1312 started.notified().await;
1313 }
1314 Ok(reply(decision.selected_model()))
1315 }
1316 };
1317 (algo, serve)
1318 }
1319
1320 #[tokio::test]
1321 async fn run_returns_the_winner_without_a_late_loser_overwriting_it() -> Result<()> {
1322 let (algo, serve) = hedge(Some(std::time::Duration::from_millis(50)));
1325 let (_trace, response) = test_drive(algo, Context::default(), request(), serve).await?;
1326 assert_eq!(
1327 response
1328 .llm_response
1329 .as_agg()
1330 .map(completion_text)
1331 .unwrap_or_default(),
1332 "winner"
1333 );
1334 Ok(())
1335 }
1336
1337 #[tokio::test]
1338 async fn run_returns_the_winner_without_hanging_on_a_pending_loser() -> Result<()> {
1339 let (algo, serve) = hedge(None);
1342 let run = test_drive(algo, Context::default(), request(), serve);
1343 let (_trace, response) = tokio::time::timeout(std::time::Duration::from_secs(1), run)
1344 .await
1345 .map_err(|error| LibsyError::external("waiting for pending loser", error))??;
1346 assert_eq!(
1347 response
1348 .llm_response
1349 .as_agg()
1350 .map(completion_text)
1351 .unwrap_or_default(),
1352 "winner"
1353 );
1354 Ok(())
1355 }
1356
1357 #[tokio::test]
1358 async fn run_surfaces_a_terminal_error_with_many_calls_in_flight() -> Result<()> {
1359 use std::sync::atomic::{AtomicUsize, Ordering};
1360
1361 const N: usize = 10;
1364
1365 struct FanOutThenError {
1368 all_started: Arc<tokio::sync::Notify>,
1369 n: usize,
1370 }
1371
1372 #[async_trait]
1373 impl Algorithm for FanOutThenError {
1374 fn name(&self) -> &str {
1375 "fan_out_then_error"
1376 }
1377
1378 async fn create_run_task(
1379 self: Arc<Self>,
1380 ctx: Context,
1381 driver: Driver,
1382 request: Request,
1383 ) -> Result<Response> {
1384 let offloads = futures::future::join_all((0..self.n).map(|i| {
1385 let decision: Arc<dyn Decision> = Arc::new(TestDecision {
1386 model: format!("m{i}"),
1387 });
1388 driver.call_llm(RoutedRequest {
1389 request: request.clone(),
1390 decision,
1391 ctx: ctx.clone(),
1392 })
1393 }));
1394 tokio::select! {
1395 _ = offloads => Err(test_error("offloads unexpectedly completed")),
1396 _ = self.all_started.notified() => {
1397 Err(test_error("terminal error while calls pending"))
1398 }
1399 }
1400 }
1401 }
1402
1403 let all_started = Arc::new(tokio::sync::Notify::new());
1404 let algo: Arc<dyn Algorithm> = Arc::new(FanOutThenError {
1405 all_started: all_started.clone(),
1406 n: N,
1407 });
1408
1409 let started = Arc::new(AtomicUsize::new(0));
1411 let serve = move |_decision: Arc<dyn Decision>, _request: Request| {
1412 let started = started.clone();
1413 let all_started = all_started.clone();
1414 async move {
1415 if started.fetch_add(1, Ordering::SeqCst) + 1 == N {
1416 all_started.notify_one();
1417 }
1418 std::future::pending::<ServeResult>().await
1419 }
1420 };
1421
1422 let run = test_drive(algo, Context::default(), request(), serve);
1425 let result = tokio::time::timeout(std::time::Duration::from_millis(500), run)
1426 .await
1427 .map_err(|error| {
1428 LibsyError::external("waiting for terminal error with full call cap", error)
1429 })?;
1430 match result {
1431 Ok(_) => Err(test_error("expected the terminal error, got a response")),
1432 Err(err) => {
1433 assert!(
1434 err.to_string()
1435 .contains("terminal error while calls pending")
1436 );
1437 Ok(())
1438 }
1439 }
1440 }
1441}