1use std::{
9 collections::{HashMap, HashSet},
10 pin::Pin,
11 sync::Arc,
12 time::{Duration, Instant},
13};
14
15use async_trait::async_trait;
16use futures::{Stream, StreamExt};
17use parking_lot::Mutex;
18use tracing::Instrument;
19
20use switchyard_protocol::{
28 Context, Decision, LlmClientError, Request, Response, RoutedLlmClient, Signals, Usage,
29};
30
31use super::driver::{DriverRequest, DriverStep, TypeErasedDriver};
32use crate::{DriverError, LibsyError, Result, observability};
33
34pub type StepStream = Pin<Box<dyn Stream<Item = Result<Step>> + Send>>;
38
39#[derive(Clone, Debug)]
41pub struct LlmCallObservation {
42 pub selected_model: String,
44 pub tier: Option<String>,
46 pub is_routed: bool,
48 pub is_success: bool,
50 pub duration: Duration,
52 pub usage: Option<Usage>,
54}
55
56#[derive(Clone, Debug)]
58pub enum RunObservation {
59 LlmCall(LlmCallObservation),
61 RoutingOverhead(Duration),
63}
64
65pub type RunObserver = Arc<dyn Fn(RunObservation) + Send + Sync>;
71
72#[derive(Clone)]
81pub struct RoutedRequest {
82 pub request: Request,
84 pub decision: Arc<dyn Decision>,
86 pub default_client: Option<Arc<dyn RoutedLlmClient>>,
90 pub ctx: Context,
94}
95
96pub struct CallLlmRequest {
104 inner: DriverRequest,
105 routed: RoutedRequest,
106}
107
108impl CallLlmRequest {
109 fn new(inner: DriverRequest) -> Self {
112 let routed = match inner.request::<RoutedRequest>() {
115 Ok(routed) => routed.clone(),
116 Err(_) => unreachable!("CallLlmRequest payload is always a RoutedRequest"),
117 };
118 Self { inner, routed }
119 }
120
121 pub fn get_routed(&self) -> &RoutedRequest {
125 &self.routed
126 }
127
128 pub fn get_request(&self) -> &Request {
130 &self.get_routed().request
131 }
132
133 pub fn get_decision(&self) -> &dyn Decision {
135 self.get_routed().decision.as_ref()
136 }
137
138 pub fn respond(self, result: Result<Response>) -> Result<()> {
142 self.inner.respond::<Response>(result)
143 }
144}
145
146#[derive(Clone)]
153pub struct Driver {
154 driver: TypeErasedDriver,
155 routed_call: Arc<Mutex<Option<Duration>>>,
157 observer: Option<RunObserver>,
158}
159
160impl Driver {
161 pub(crate) fn new() -> Self {
164 Self::with_observer(None)
165 }
166
167 fn with_observer(observer: Option<RunObserver>) -> Self {
168 Self {
169 driver: TypeErasedDriver::new(),
170 routed_call: Arc::new(Mutex::new(None)),
171 observer,
172 }
173 }
174
175 pub(crate) fn routed_call_duration(&self) -> Option<Duration> {
177 *self.routed_call.lock()
178 }
179
180 pub(crate) fn observe_routing_overhead(&self, duration: Duration) {
183 if let Some(observer) = &self.observer {
184 observer(RunObservation::RoutingOverhead(duration));
185 }
186 }
187
188 #[tracing::instrument(
197 target = "libsy",
198 name = "libsy.llm_call",
199 skip_all,
200 fields(
201 algorithm = observability::algorithm_label(&routed.ctx),
202 selected_model = routed.decision.selected_model(),
203 openinference.span.kind = "CHAIN",
204 outcome = tracing::field::Empty,
205 error = tracing::field::Empty,
206 input_tokens = tracing::field::Empty,
207 output_tokens = tracing::field::Empty,
208 total_tokens = tracing::field::Empty,
209 reasoning_tokens = tracing::field::Empty,
210 )
211 )]
212 pub async fn call_llm(&self, routed: RoutedRequest) -> Result<Response> {
213 let algorithm = observability::algorithm_label(&routed.ctx).to_string();
214 let selected_model = routed.decision.selected_model().to_string();
215 let tier = routed.decision.routing_tier().map(str::to_string);
216 let is_routed = routed.decision.is_routed_call();
217 let started = Instant::now();
218 let result = self
219 .driver
220 .fulfill_request::<RoutedRequest, Response>(routed.ctx.clone(), routed)
221 .await;
222 let elapsed = started.elapsed();
223 observability::record_llm_call(
224 &algorithm,
225 &selected_model,
226 tier.as_deref(),
227 is_routed,
228 elapsed,
229 &result,
230 &tracing::Span::current(),
231 );
232 if let Some(observer) = &self.observer {
233 observer(RunObservation::LlmCall(LlmCallObservation {
234 selected_model,
235 tier,
236 is_routed,
237 is_success: result.is_ok(),
238 duration: elapsed,
239 usage: result
240 .as_ref()
241 .ok()
242 .and_then(|response| response.llm_response.as_agg())
243 .map(|response| response.usage.clone()),
244 }));
245 }
246 if is_routed && result.is_ok() {
249 *self.routed_call.lock() = Some(elapsed);
250 }
251 result
252 }
253
254 pub async fn call_llm_target(
260 &self,
261 ctx: Context,
262 target: &LlmTarget,
263 request: Request,
264 decision: Arc<dyn Decision>,
265 ) -> Result<Response> {
266 self.call_llm(RoutedRequest {
267 request,
268 decision,
269 default_client: target.llm_client.clone(),
270 ctx,
271 })
272 .await
273 }
274
275 pub async fn info(&self, ctx: Context, decision: Arc<dyn Decision>) -> Result<()> {
279 self.driver.info(ctx.clone(), decision.clone()).await?;
280 observability::record_decision(&ctx, decision.as_ref());
281 Ok(())
282 }
283
284 pub(crate) async fn finish(&self, ctx: Context, result: Result<Response>) -> Result<()> {
288 match result {
289 Ok(response) => self.driver.done(ctx, response).await,
290 Err(err) => self.driver.fail(ctx, err).await,
291 }
292 }
293
294 pub(crate) fn stream(&self) -> impl Stream<Item = Result<Step>> + use<> {
298 self.driver.stream().map(|item| match item? {
299 DriverStep::Request(req) => Ok(Step::CallLlm(Box::new(CallLlmRequest::new(req)))),
300 DriverStep::Info(payload) => payload
301 .downcast::<Arc<dyn Decision>>()
302 .map(|decision| Step::Decision(*decision))
303 .map_err(|_| {
304 DriverError::TypeMismatch {
305 expected: "Arc<dyn Decision>",
306 }
307 .into()
308 }),
309 DriverStep::Done(payload) => payload
310 .downcast::<Response>()
311 .map(Step::ReturnToAgent)
312 .map_err(|_| {
313 DriverError::TypeMismatch {
314 expected: "Response",
315 }
316 .into()
317 }),
318 })
319 }
320}
321
322impl Default for Driver {
323 fn default() -> Self {
324 Self::new()
325 }
326}
327
328pub enum Step {
330 CallLlm(Box<CallLlmRequest>),
334 Decision(Arc<dyn Decision>),
337 ReturnToAgent(Box<Response>),
339}
340
341struct AbortOnDrop(tokio::task::AbortHandle);
343
344impl Drop for AbortOnDrop {
345 fn drop(&mut self) {
346 self.0.abort();
347 }
348}
349
350#[derive(Clone)]
355pub struct LlmTarget {
356 pub semantic_name: String,
360 pub llm_client: Option<Arc<dyn RoutedLlmClient>>,
363}
364
365#[derive(Clone)]
369pub struct LlmTargetSet {
370 targets: Vec<LlmTarget>,
371}
372
373impl LlmTargetSet {
374 pub fn new(targets: Vec<LlmTarget>) -> Self {
376 Self { targets }
377 }
378
379 pub fn targets(&self) -> &[LlmTarget] {
381 &self.targets
382 }
383
384 pub fn get_target(&self, name: &str) -> Result<LlmTarget> {
386 self.targets
387 .iter()
388 .find(|t| t.semantic_name == name)
389 .cloned()
390 .ok_or_else(|| LibsyError::TargetNotFound {
391 target: name.to_string(),
392 })
393 }
394
395 pub fn resolve_target(&self, name: &str, ctx: &Context) -> Result<LlmTarget> {
398 let target = self.get_target(name)?;
399 if !ctx.is_excluded(&target.semantic_name) {
400 return Ok(target);
401 }
402 self.targets
403 .iter()
404 .find(|t| !ctx.is_excluded(&t.semantic_name))
405 .cloned()
406 .ok_or(LibsyError::AllTargetsExcluded)
407 }
408}
409
410const MAX_EVICTION_SESSIONS: usize = 1_024;
413
414#[derive(Default)]
420pub(crate) struct SessionEvictions {
421 sessions: Mutex<HashMap<String, HashSet<String>>>,
422}
423
424impl SessionEvictions {
425 pub(crate) fn remove(&self, session: &str) {
427 self.sessions.lock().remove(session);
428 }
429
430 fn evicted_in(&self, session: Option<&str>) -> Vec<String> {
432 let Some(session) = session else {
433 return Vec::new();
434 };
435 self.sessions
436 .lock()
437 .get(session)
438 .map(|targets| targets.iter().cloned().collect())
439 .unwrap_or_default()
440 }
441
442 fn record(&self, session: Option<&str>, target: &str) {
445 let Some(session) = session else { return };
446 let mut sessions = self.sessions.lock();
447 if sessions.len() >= MAX_EVICTION_SESSIONS
448 && !sessions.contains_key(session)
449 && let Some(oldest) = sessions.keys().next().cloned()
450 {
451 sessions.remove(&oldest);
452 }
453 sessions
454 .entry(session.to_string())
455 .or_default()
456 .insert(target.to_string());
457 }
458}
459
460fn eligible_targets(targets: &LlmTargetSet, ctx: &Context) -> usize {
462 targets
463 .targets()
464 .iter()
465 .filter(|t| !ctx.is_excluded(&t.semantic_name))
466 .count()
467}
468
469pub(crate) fn exclude_evicted(
472 ctx: &mut Context,
473 targets: &LlmTargetSet,
474 evictions: &SessionEvictions,
475 session: Option<&str>,
476) {
477 for target in evictions.evicted_in(session) {
478 if eligible_targets(targets, ctx) <= 1 {
481 break;
482 }
483 ctx.exclude_target(target);
484 }
485}
486
487#[allow(clippy::too_many_arguments)]
495pub(crate) async fn call_llm_with_overflow_fallback(
496 mut ctx: Context,
497 driver: &Driver,
498 targets: &LlmTargetSet,
499 mut target: LlmTarget,
500 mut decision: Arc<dyn Decision>,
501 request: Request,
502 session: Option<&str>,
503 evictions: &SessionEvictions,
504 fallback_decision: impl Fn(&LlmTarget, &LlmTarget) -> Arc<dyn Decision>,
505) -> Result<Response> {
506 loop {
507 let result = driver
508 .call_llm_target(ctx.clone(), &target, request.clone(), decision.clone())
509 .await;
510 let Err(error) = result else { return result };
511 let LibsyError::ClientCall {
512 target: failed,
513 source: LlmClientError::ContextWindowExceeded { .. },
514 } = &error
515 else {
516 return Err(error);
517 };
518 if !ctx.exclude_target(failed) {
521 return Err(error);
522 }
523 evictions.record(session, failed);
524 let Ok(next) = targets.resolve_target(&target.semantic_name, &ctx) else {
525 return Err(error);
526 };
527 decision = fallback_decision(&target, &next);
528 target = next;
529 driver.info(ctx.clone(), decision.clone()).await?;
530 }
531}
532
533#[async_trait]
554pub trait Algorithm: Send + Sync + 'static {
555 fn name(&self) -> &str;
559
560 async fn create_run_task(
566 self: Arc<Self>,
567 ctx: Context,
568 driver: Driver,
569 request: Request,
570 ) -> Result<Response>;
571
572 #[allow(unused_variables)]
576 async fn process_signals(self: Arc<Self>, signals: Signals) -> Result<()> {
577 Ok(())
578 }
579
580 fn run_stream(
590 self: Arc<Self>,
591 ctx: Context,
592 request: Request,
593 observer: Option<RunObserver>,
594 ) -> StepStream {
595 let mut ctx = ctx;
598 ctx.values.insert(
599 observability::ALGORITHM_KEY.to_string(),
600 self.name().to_string(),
601 );
602 let driver = Driver::with_observer(observer);
603 let task_driver = driver.clone();
604 let task_ctx = ctx.clone();
605 let stream = task_driver.stream();
606 let span = observability::run_span(self.name(), &request);
610 let observed_driver = task_driver.clone();
611 let handle = tokio::spawn(
612 async move {
613 observability::observe_run(
614 task_ctx.clone(),
615 observed_driver,
616 self.create_run_task(task_ctx, task_driver, request),
617 )
618 .await
619 }
620 .instrument(span),
621 );
622 let abort_guard = AbortOnDrop(handle.abort_handle());
624
625 let finish_driver = driver.clone();
626 let finish_ctx = ctx;
627 let tail: StepStream = Box::pin(
628 futures::stream::once(async move {
629 let result = match handle.await {
630 Ok(response) => response,
631 Err(source) => Err(LibsyError::AlgorithmTask { source }),
632 };
633 finish_driver.finish(finish_ctx, result).await
634 })
635 .filter_map(|finish_result| async move { finish_result.err().map(Err) }),
636 );
637
638 let stream: StepStream = Box::pin(stream);
639 Box::pin(futures::stream::select(stream, tail).map(move |step| {
640 let _keep_alive = &abort_guard;
642 step
643 }))
644 }
645
646 async fn run(
650 self: Arc<Self>,
651 ctx: Context,
652 request: Request,
653 ) -> Result<(Vec<Arc<dyn Decision>>, Response)> {
654 self.run_observed(ctx, request, None).await
655 }
656
657 async fn run_observed(
659 self: Arc<Self>,
660 ctx: Context,
661 request: Request,
662 observer: Option<RunObserver>,
663 ) -> Result<(Vec<Arc<dyn Decision>>, Response)> {
664 #[tracing::instrument(
670 target = "libsy",
671 name = "libsy.client_call",
672 skip_all,
673 fields(
674 algorithm = observability::algorithm_label(&call.get_routed().ctx),
675 switchyard.algorithm = observability::algorithm_label(&call.get_routed().ctx),
676 switchyard.routing.tier = tracing::field::Empty,
677 selected_model = call.get_decision().selected_model(),
678 otel.kind = "client",
679 otel.name = %format_args!("chat {}", call.get_decision().selected_model()),
680 openinference.span.kind = "LLM",
681 gen_ai.operation.name = "chat",
682 gen_ai.request.model = call.get_decision().selected_model(),
683 gen_ai.request.stream = tracing::field::Empty,
684 gen_ai.request.temperature = tracing::field::Empty,
685 gen_ai.request.top_p = tracing::field::Empty,
686 gen_ai.request.top_k = tracing::field::Empty,
687 gen_ai.request.max_tokens = tracing::field::Empty,
688 gen_ai.request.reasoning.level = tracing::field::Empty,
689 gen_ai.output.type = tracing::field::Empty,
690 gen_ai.conversation.id = tracing::field::Empty,
691 server.address = tracing::field::Empty,
692 server.port = tracing::field::Empty,
693 gen_ai.response.id = tracing::field::Empty,
694 gen_ai.response.model = tracing::field::Empty,
695 gen_ai.usage.input_tokens = tracing::field::Empty,
696 gen_ai.usage.output_tokens = tracing::field::Empty,
697 gen_ai.usage.cache_read.input_tokens = tracing::field::Empty,
698 gen_ai.usage.cache_creation.input_tokens = tracing::field::Empty,
699 gen_ai.usage.reasoning.output_tokens = tracing::field::Empty,
700 outcome = tracing::field::Empty,
701 otel.status_code = tracing::field::Empty,
702 error.type = tracing::field::Empty,
703 error = tracing::field::Empty,
704 )
705 )]
706 async fn serve(call: CallLlmRequest) -> Result<()> {
707 let span = tracing::Span::current();
708 observability::record_gen_ai_request(&span, &call.get_routed().request.llm_request);
709 if let Some(tier) = call.get_decision().routing_tier() {
710 span.record("switchyard.routing.tier", tier);
711 }
712 if let Some(session_id) = call
713 .get_routed()
714 .request
715 .metadata
716 .as_ref()
717 .and_then(|metadata| metadata.session_id.as_deref())
718 {
719 span.record("gen_ai.conversation.id", session_id);
720 }
721 let routed = call.get_routed().clone();
722 let target = routed.decision.selected_model().to_string();
723 let client =
724 routed
725 .default_client
726 .clone()
727 .ok_or_else(|| LibsyError::MissingClient {
728 target: target.clone(),
729 })?;
730 let result = client
731 .call(routed.ctx, routed.request, routed.decision)
732 .await
733 .map_err(|source| LibsyError::client_call(target, source));
734 let result = observability::observe_client_call(result);
735 call.respond(result)
736 }
737
738 let stream = self.run_stream(ctx, request, observer);
739 tokio::pin!(stream);
740
741 let mut trace: Vec<Arc<dyn Decision>> = Vec::new();
742 let mut in_flight = futures::stream::FuturesUnordered::new();
743 let mut final_response: Option<Response> = None;
744
745 loop {
746 tokio::select! {
747 Some(result) = in_flight.next() => match result {
748 Ok(()) => {}, Err(err) => return Err(err), },
751 step = stream.next() => {
752 match step {
753 None => break, Some(item) => match item? {
755 Step::CallLlm(call) => in_flight.push(serve(*call)),
756 Step::Decision(decision) => trace.push(decision),
757 Step::ReturnToAgent(response) => {
758 final_response = Some(*response);
759 break;
760 }
761 }
762 }
763 },
764 }
765 }
766 final_response
767 .map(|response| (trace, response))
768 .ok_or(LibsyError::MissingFinalResponse)
769 }
770}
771
772#[cfg(test)]
773mod tests {
774 use super::*;
775 use futures::StreamExt;
776 use switchyard_protocol::{
777 LlmResponse, LlmResponseChunk, completion_text, text_request, text_response,
778 };
779
780 #[derive(Debug, thiserror::Error)]
781 #[error("{0}")]
782 struct TestError(&'static str);
783
784 fn test_error(message: &'static str) -> LibsyError {
785 LibsyError::external("test", TestError(message))
786 }
787
788 struct EchoClient;
790
791 #[async_trait]
792 impl RoutedLlmClient for EchoClient {
793 async fn call(
794 &self,
795 _ctx: Context,
796 _request: Request,
797 decision: Arc<dyn Decision>,
798 ) -> std::result::Result<Response, LlmClientError> {
799 Ok(Response {
801 llm_response: LlmResponse::Agg(text_response(
802 None,
803 decision.selected_model().to_string(),
804 )),
805 metadata: None,
806 })
807 }
808 }
809
810 struct TestDecision {
813 model: String,
814 }
815
816 impl Decision for TestDecision {
817 fn selected_model(&self) -> &str {
818 &self.model
819 }
820 fn reasoning(&self) -> Option<&str> {
821 None
822 }
823 fn as_any(&self) -> &dyn std::any::Any {
824 self
825 }
826 }
827
828 struct TestAlgo {
829 target_set: LlmTargetSet,
830 }
831
832 #[async_trait]
833 impl Algorithm for TestAlgo {
834 fn name(&self) -> &str {
835 "test"
836 }
837
838 async fn create_run_task(
839 self: Arc<Self>,
840 ctx: Context,
841 driver: Driver,
842 request: Request,
843 ) -> Result<Response> {
844 let target = self
845 .target_set
846 .targets()
847 .first()
848 .ok_or(LibsyError::NoTargets)?
849 .clone();
850 let decision: Arc<dyn Decision> = Arc::new(TestDecision {
851 model: target.semantic_name.clone(),
852 });
853 driver.info(ctx.clone(), decision.clone()).await?;
854 driver
855 .call_llm_target(ctx, &target, request, decision)
856 .await
857 }
858 }
859
860 fn orch(target_set: LlmTargetSet) -> Arc<dyn Algorithm> {
862 Arc::new(TestAlgo { target_set })
863 }
864
865 fn request() -> Request {
866 Request {
867 llm_request: text_request(Some("auto".to_string()), "hi".to_string()),
868 raw_request: None,
869 metadata: None,
870 }
871 }
872
873 fn target_set(names: &[(&str, bool)]) -> LlmTargetSet {
875 let targets = names
876 .iter()
877 .map(|(name, has_client)| LlmTarget {
878 semantic_name: name.to_string(),
879 llm_client: has_client.then(|| Arc::new(EchoClient) as Arc<dyn RoutedLlmClient>),
880 })
881 .collect();
882 LlmTargetSet::new(targets)
883 }
884
885 #[tokio::test]
886 async fn observed_run_reports_one_successful_routed_call() -> Result<()> {
887 let observations = Arc::new(Mutex::new(Vec::new()));
888 let observed = observations.clone();
889 let observer: RunObserver = Arc::new(move |observation| observed.lock().push(observation));
890 let (_, response) = orch(target_set(&[("direct/model", true)]))
891 .run_observed(Context::default(), request(), Some(observer))
892 .await?;
893 assert_eq!(
894 response.llm_response.as_agg().map(completion_text),
895 Some("direct/model".to_string())
896 );
897 let observations = observations.lock();
898 assert_eq!(observations.len(), 2);
899 let RunObservation::LlmCall(observation) = &observations[0] else {
900 return Err(test_error("expected an LLM call observation"));
901 };
902 assert_eq!(observation.selected_model, "direct/model");
903 assert!(observation.is_routed);
904 assert!(observation.is_success);
905 assert!(observation.usage.is_some());
906 assert!(matches!(
907 observations[1],
908 RunObservation::RoutingOverhead(_)
909 ));
910 Ok(())
911 }
912
913 #[test]
914 fn target_lookup_returns_the_missing_target() {
915 let error = target_set(&[]).get_target("missing").err();
916 assert!(matches!(
917 error,
918 Some(LibsyError::TargetNotFound { target }) if target == "missing"
919 ));
920 }
921
922 struct StreamingClient {
925 chunks: Vec<LlmResponseChunk>,
926 }
927
928 #[async_trait]
929 impl RoutedLlmClient for StreamingClient {
930 async fn call(
931 &self,
932 _ctx: Context,
933 _request: Request,
934 _decision: Arc<dyn Decision>,
935 ) -> std::result::Result<Response, LlmClientError> {
936 let stream = futures::stream::iter(
937 self.chunks
938 .clone()
939 .into_iter()
940 .map(|chunk| Ok(chunk.into())),
941 )
942 .boxed();
943 Ok(Response {
944 llm_response: LlmResponse::Stream(stream),
945 metadata: None,
946 })
947 }
948 }
949
950 fn streaming_orch(chunks: Vec<LlmResponseChunk>) -> Arc<dyn Algorithm> {
952 let target = LlmTarget {
953 semantic_name: "stream/model".to_string(),
954 llm_client: Some(Arc::new(StreamingClient { chunks }) as Arc<dyn RoutedLlmClient>),
955 };
956 orch(LlmTargetSet::new(vec![target]))
957 }
958
959 #[tokio::test]
960 async fn run_returns_a_streamed_response_the_caller_aggregates() -> Result<()> {
961 let orch = streaming_orch(vec![
964 LlmResponseChunk::MessageStart {
965 id: Some("m1".to_string()),
966 model: Some("stream/model".to_string()),
967 },
968 LlmResponseChunk::TextDelta {
969 index: 0,
970 text: "hel".to_string(),
971 },
972 LlmResponseChunk::TextDelta {
973 index: 0,
974 text: "lo".to_string(),
975 },
976 LlmResponseChunk::MessageStop {
977 reason: Some("stop".to_string()),
978 },
979 ]);
980 let (trace, response) = orch.run(Context::default(), request()).await?;
981 let agg = response
983 .llm_response
984 .into_agg()
985 .await
986 .map_err(|error| LibsyError::external("aggregating response stream", error))?;
987 assert_eq!(completion_text(&agg), "hello");
988 assert_eq!(agg.model.as_deref(), Some("stream/model"));
989 assert_eq!(trace.len(), 1);
990 Ok(())
991 }
992
993 #[tokio::test]
994 async fn aggregating_a_streamed_response_propagates_a_mid_stream_error() -> Result<()> {
995 let orch = streaming_orch(vec![
998 LlmResponseChunk::TextDelta {
999 index: 0,
1000 text: "partial".to_string(),
1001 },
1002 LlmResponseChunk::StreamError {
1003 message: "upstream exploded".to_string(),
1004 },
1005 ]);
1006 let (_, response) = orch.run(Context::default(), request()).await?;
1007 match response.llm_response.into_agg().await {
1008 Ok(_) => panic!("expected a mid-stream error, got an aggregate"),
1009 Err(err) => {
1010 assert!(err.to_string().contains("upstream exploded"));
1011 Ok(())
1012 }
1013 }
1014 }
1015
1016 #[tokio::test]
1017 async fn run_offloads_via_promise_then_returns_to_agent() -> Result<()> {
1018 let stream = orch(target_set(&[("offload/model", false)])).run_stream(
1021 Context::default(),
1022 request(),
1023 None,
1024 );
1025 tokio::pin!(stream);
1026
1027 let mut saw_call = false;
1028 let mut final_completion = None;
1029 while let Some(step) = stream.next().await {
1030 match step? {
1031 Step::CallLlm(call) => {
1032 saw_call = true;
1033 assert_eq!(call.get_decision().selected_model(), "offload/model");
1035 call.respond(Ok(Response {
1037 llm_response: LlmResponse::Agg(text_response(
1038 None,
1039 "fulfilled".to_string(),
1040 )),
1041 metadata: None,
1042 }))?;
1043 }
1044 Step::Decision(decision) => {
1045 assert_eq!(decision.selected_model(), "offload/model");
1046 }
1047 Step::ReturnToAgent(response) => {
1048 final_completion = Some(
1049 response
1050 .llm_response
1051 .as_agg()
1052 .map(completion_text)
1053 .unwrap_or_default(),
1054 );
1055 }
1056 }
1057 }
1058
1059 assert!(saw_call, "expected a CallLlm step before ReturnToAgent");
1060 assert_eq!(
1061 final_completion.ok_or_else(|| test_error("no ReturnToAgent step"))?,
1062 "fulfilled"
1063 );
1064 Ok(())
1065 }
1066
1067 #[tokio::test]
1068 async fn client_backed_target_offloads_with_a_default_client() -> Result<()> {
1069 let stream = orch(target_set(&[("direct/model", true)])).run_stream(
1072 Context::default(),
1073 request(),
1074 None,
1075 );
1076 tokio::pin!(stream);
1077
1078 let mut final_completion = None;
1079 while let Some(step) = stream.next().await {
1080 match step? {
1081 Step::CallLlm(call) => {
1082 let routed = call.get_routed().clone();
1083 let client = routed
1084 .default_client
1085 .clone()
1086 .ok_or_else(|| test_error("expected a default client"))?;
1087 let target = routed.decision.selected_model().to_string();
1088 let result = client
1089 .call(routed.ctx, routed.request, routed.decision)
1090 .await
1091 .map_err(|error| LibsyError::client_call(target, error));
1092 call.respond(result)?;
1093 }
1094 Step::Decision(_) => {}
1095 Step::ReturnToAgent(response) => {
1096 final_completion = Some(
1097 response
1098 .llm_response
1099 .as_agg()
1100 .map(completion_text)
1101 .unwrap_or_default(),
1102 );
1103 }
1104 }
1105 }
1106
1107 assert_eq!(
1109 final_completion.ok_or_else(|| test_error("no ReturnToAgent"))?,
1110 "direct/model"
1111 );
1112 Ok(())
1113 }
1114
1115 #[tokio::test]
1116 async fn run_returns_the_response_when_all_targets_have_clients() -> Result<()> {
1117 let (trace, response) = orch(target_set(&[("direct/model", true)]))
1120 .run(Context::default(), request())
1121 .await?;
1122 assert_eq!(
1124 response
1125 .llm_response
1126 .as_agg()
1127 .map(completion_text)
1128 .unwrap_or_default(),
1129 "direct/model"
1130 );
1131 assert_eq!(trace[0].selected_model(), "direct/model");
1132 Ok(())
1133 }
1134
1135 #[tokio::test]
1136 async fn run_errors_when_a_target_lacks_a_client() -> Result<()> {
1137 let error = orch(target_set(&[("offload/model", false)]))
1140 .run(Context::default(), request())
1141 .await
1142 .err()
1143 .ok_or_else(|| test_error("expected a missing-client error"))?;
1144 assert!(matches!(
1145 error,
1146 LibsyError::MissingClient { target } if target == "offload/model"
1147 ));
1148 Ok(())
1149 }
1150
1151 #[tokio::test(flavor = "multi_thread", worker_threads = 12)]
1152 async fn requests_are_processed_in_parallel() -> Result<()> {
1153 use std::time::Duration;
1154 use tokio::sync::Barrier;
1155
1156 const N: usize = 12;
1157
1158 struct BarrierClient {
1164 barrier: Arc<Barrier>,
1165 }
1166
1167 #[async_trait]
1168 impl RoutedLlmClient for BarrierClient {
1169 async fn call(
1170 &self,
1171 _ctx: Context,
1172 _request: Request,
1173 decision: Arc<dyn Decision>,
1174 ) -> std::result::Result<Response, LlmClientError> {
1175 self.barrier.wait().await;
1176 Ok(Response {
1177 llm_response: LlmResponse::Agg(text_response(
1178 None,
1179 decision.selected_model().to_string(),
1180 )),
1181 metadata: None,
1182 })
1183 }
1184 }
1185
1186 let barrier = Arc::new(Barrier::new(N));
1187 let targets = LlmTargetSet::new(vec![LlmTarget {
1188 semantic_name: "m".to_string(),
1189 llm_client: Some(Arc::new(BarrierClient {
1190 barrier: barrier.clone(),
1191 })),
1192 }]);
1193 let algo = orch(targets);
1195
1196 let mut handles = Vec::new();
1197 for _ in 0..N {
1198 let algo = algo.clone();
1199 handles.push(tokio::spawn(async move {
1200 algo.run(Context::default(), request())
1201 .await
1202 .map(|(_, response)| {
1203 response
1204 .llm_response
1205 .as_agg()
1206 .map(completion_text)
1207 .unwrap_or_default()
1208 })
1209 }));
1210 }
1211
1212 for handle in handles {
1213 let completion = tokio::time::timeout(Duration::from_secs(5), handle)
1215 .await
1216 .map_err(|error| LibsyError::external("waiting for test task", error))?
1217 .map_err(|source| LibsyError::AlgorithmTask { source })??;
1218 assert_eq!(completion, "m");
1219 }
1220 Ok(())
1221 }
1222
1223 #[tokio::test]
1224 async fn offload_error_propagates_back_to_the_algorithm() -> Result<()> {
1225 let stream = orch(target_set(&[("offload/model", false)])).run_stream(
1229 Context::default(),
1230 request(),
1231 None,
1232 );
1233 tokio::pin!(stream);
1234
1235 let mut saw_error = false;
1236 while let Some(step) = stream.next().await {
1237 match step {
1238 Ok(Step::CallLlm(call)) => {
1239 call.respond(Err(test_error("upstream model call failed")))?;
1240 }
1241 Ok(Step::Decision(_)) => {}
1242 Ok(Step::ReturnToAgent(..)) => {
1243 return Err(test_error(
1244 "expected the offload error to propagate, got a response",
1245 ));
1246 }
1247 Err(err) => {
1248 assert!(err.to_string().contains("upstream model call failed"));
1250 saw_error = true;
1251 }
1252 }
1253 }
1254
1255 assert!(saw_error, "expected an error step");
1256 Ok(())
1257 }
1258
1259 #[tokio::test]
1260 async fn dropping_the_stream_cancels_the_algorithm_task() -> Result<()> {
1261 use std::sync::atomic::{AtomicBool, Ordering};
1262 use std::time::Duration;
1263 use tokio::sync::mpsc;
1264
1265 struct DropGuard(Arc<AtomicBool>);
1268 impl Drop for DropGuard {
1269 fn drop(&mut self) {
1270 self.0.store(true, Ordering::SeqCst);
1271 }
1272 }
1273
1274 struct StuckAlgo {
1275 started: mpsc::UnboundedSender<()>,
1276 dropped: Arc<AtomicBool>,
1277 }
1278
1279 #[async_trait]
1280 impl Algorithm for StuckAlgo {
1281 fn name(&self) -> &str {
1282 "stuck"
1283 }
1284
1285 async fn create_run_task(
1286 self: Arc<Self>,
1287 _ctx: Context,
1288 _driver: Driver,
1289 _request: Request,
1290 ) -> Result<Response> {
1291 let _guard = DropGuard(self.dropped.clone());
1292 let _ = self.started.send(());
1293 std::future::pending::<()>().await;
1295 unreachable!()
1296 }
1297 }
1298
1299 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1300 let dropped = Arc::new(AtomicBool::new(false));
1301 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1302 started: started_tx,
1303 dropped: dropped.clone(),
1304 });
1305
1306 let stream = algo.run_stream(Context::default(), request(), None);
1307 started_rx
1308 .recv()
1309 .await
1310 .ok_or_else(|| test_error("task never started"))?;
1311 drop(stream);
1312 tokio::time::sleep(Duration::from_millis(100)).await;
1313
1314 assert!(
1315 dropped.load(Ordering::SeqCst),
1316 "algorithm task was NOT cancelled after dropping the stream"
1317 );
1318 Ok(())
1319 }
1320
1321 #[tokio::test]
1322 async fn create_run_task_panic_surfaces_as_a_stream_error() -> Result<()> {
1323 struct Panicky;
1326
1327 #[async_trait]
1328 impl Algorithm for Panicky {
1329 fn name(&self) -> &str {
1330 "panicky"
1331 }
1332
1333 async fn create_run_task(
1334 self: Arc<Self>,
1335 _ctx: Context,
1336 _driver: Driver,
1337 _request: Request,
1338 ) -> Result<Response> {
1339 panic!("boom");
1340 }
1341 }
1342
1343 let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
1344 let stream = algo.run_stream(Context::default(), request(), None);
1345 tokio::pin!(stream);
1346
1347 let mut saw_error = false;
1348 while let Some(step) = stream.next().await {
1349 match step {
1350 Err(err) => {
1351 assert!(matches!(err, LibsyError::AlgorithmTask { .. }));
1352 saw_error = true;
1353 }
1354 Ok(_) => return Err(test_error("expected the panic to surface as an error step")),
1355 }
1356 }
1357
1358 assert!(saw_error, "expected an error step from the panicked task");
1359 Ok(())
1360 }
1361
1362 #[tokio::test]
1363 async fn run_returns_an_error_when_the_algorithm_task_panics() -> Result<()> {
1364 struct Panicky;
1367
1368 #[async_trait]
1369 impl Algorithm for Panicky {
1370 fn name(&self) -> &str {
1371 "panicky"
1372 }
1373
1374 async fn create_run_task(
1375 self: Arc<Self>,
1376 _ctx: Context,
1377 _driver: Driver,
1378 _request: Request,
1379 ) -> Result<Response> {
1380 panic!("boom");
1381 }
1382 }
1383
1384 let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
1385 match algo.run(Context::default(), request()).await {
1386 Ok(_) => Err(test_error(
1387 "expected run to surface the algorithm panic as an error",
1388 )),
1389 Err(err) => {
1390 assert!(matches!(err, LibsyError::AlgorithmTask { .. }));
1391 Ok(())
1392 }
1393 }
1394 }
1395
1396 #[tokio::test]
1397 async fn cancelling_run_cancels_the_algorithm_task() -> Result<()> {
1398 use std::sync::atomic::{AtomicBool, Ordering};
1399 use std::time::Duration;
1400 use tokio::sync::mpsc;
1401
1402 struct DropGuard(Arc<AtomicBool>);
1405 impl Drop for DropGuard {
1406 fn drop(&mut self) {
1407 self.0.store(true, Ordering::SeqCst);
1408 }
1409 }
1410
1411 struct StuckAlgo {
1412 started: mpsc::UnboundedSender<()>,
1413 dropped: Arc<AtomicBool>,
1414 }
1415
1416 #[async_trait]
1417 impl Algorithm for StuckAlgo {
1418 fn name(&self) -> &str {
1419 "stuck"
1420 }
1421
1422 async fn create_run_task(
1423 self: Arc<Self>,
1424 _ctx: Context,
1425 _driver: Driver,
1426 _request: Request,
1427 ) -> Result<Response> {
1428 let _guard = DropGuard(self.dropped.clone());
1429 let _ = self.started.send(());
1430 std::future::pending::<()>().await;
1433 unreachable!()
1434 }
1435 }
1436
1437 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1438 let dropped = Arc::new(AtomicBool::new(false));
1439 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1440 started: started_tx,
1441 dropped: dropped.clone(),
1442 });
1443
1444 let run_task = tokio::spawn(async move { algo.run(Context::default(), request()).await });
1447 started_rx
1448 .recv()
1449 .await
1450 .ok_or_else(|| test_error("task never started"))?;
1451 run_task.abort();
1452 tokio::time::sleep(Duration::from_millis(100)).await;
1453
1454 assert!(
1455 dropped.load(Ordering::SeqCst),
1456 "algorithm task was NOT cancelled after cancelling run"
1457 );
1458 Ok(())
1459 }
1460
1461 struct LoserClient {
1467 started: Arc<tokio::sync::Notify>,
1468 delay: Option<std::time::Duration>,
1469 }
1470
1471 #[async_trait]
1472 impl RoutedLlmClient for LoserClient {
1473 async fn call(
1474 &self,
1475 _ctx: Context,
1476 _request: Request,
1477 decision: Arc<dyn Decision>,
1478 ) -> std::result::Result<Response, LlmClientError> {
1479 self.started.notify_one();
1480 match self.delay {
1481 Some(delay) => tokio::time::sleep(delay).await,
1482 None => std::future::pending::<()>().await,
1483 }
1484 Ok(Response {
1485 llm_response: LlmResponse::Agg(text_response(
1486 None,
1487 decision.selected_model().to_string(),
1488 )),
1489 metadata: None,
1490 })
1491 }
1492 }
1493
1494 struct GatedEchoClient {
1497 gate: Arc<tokio::sync::Notify>,
1498 }
1499
1500 #[async_trait]
1501 impl RoutedLlmClient for GatedEchoClient {
1502 async fn call(
1503 &self,
1504 _ctx: Context,
1505 _request: Request,
1506 decision: Arc<dyn Decision>,
1507 ) -> std::result::Result<Response, LlmClientError> {
1508 self.gate.notified().await;
1509 Ok(Response {
1510 llm_response: LlmResponse::Agg(text_response(
1511 None,
1512 decision.selected_model().to_string(),
1513 )),
1514 metadata: None,
1515 })
1516 }
1517 }
1518
1519 struct Hedge {
1522 winner: LlmTarget,
1523 loser: LlmTarget,
1524 }
1525
1526 #[async_trait]
1527 impl Algorithm for Hedge {
1528 fn name(&self) -> &str {
1529 "hedge"
1530 }
1531
1532 async fn create_run_task(
1533 self: Arc<Self>,
1534 ctx: Context,
1535 driver: Driver,
1536 request: Request,
1537 ) -> Result<Response> {
1538 let dec_w: Arc<dyn Decision> = Arc::new(TestDecision {
1539 model: self.winner.semantic_name.clone(),
1540 });
1541 let dec_l: Arc<dyn Decision> = Arc::new(TestDecision {
1542 model: self.loser.semantic_name.clone(),
1543 });
1544 let win = driver.call_llm_target(ctx.clone(), &self.winner, request.clone(), dec_w);
1545 let lose = driver.call_llm_target(ctx, &self.loser, request, dec_l);
1546 tokio::select! {
1548 res = win => res,
1549 res = lose => res,
1550 }
1551 }
1552 }
1553
1554 fn hedge(loser_delay: Option<std::time::Duration>) -> Arc<dyn Algorithm> {
1557 let started = Arc::new(tokio::sync::Notify::new());
1558 let winner = LlmTarget {
1559 semantic_name: "winner".to_string(),
1560 llm_client: Some(Arc::new(GatedEchoClient {
1561 gate: started.clone(),
1562 })),
1563 };
1564 let loser = LlmTarget {
1565 semantic_name: "loser".to_string(),
1566 llm_client: Some(Arc::new(LoserClient {
1567 started,
1568 delay: loser_delay,
1569 })),
1570 };
1571 Arc::new(Hedge { winner, loser })
1572 }
1573
1574 #[tokio::test]
1575 async fn run_returns_the_winner_without_a_late_loser_overwriting_it() -> Result<()> {
1576 let (_trace, response) = hedge(Some(std::time::Duration::from_millis(50)))
1579 .run(Context::default(), request())
1580 .await?;
1581 assert_eq!(
1582 response
1583 .llm_response
1584 .as_agg()
1585 .map(completion_text)
1586 .unwrap_or_default(),
1587 "winner"
1588 );
1589 Ok(())
1590 }
1591
1592 #[tokio::test]
1593 async fn run_returns_the_winner_without_hanging_on_a_pending_loser() -> Result<()> {
1594 let run = hedge(None).run(Context::default(), request());
1597 let (_trace, response) = tokio::time::timeout(std::time::Duration::from_secs(1), run)
1598 .await
1599 .map_err(|error| LibsyError::external("waiting for pending loser", error))??;
1600 assert_eq!(
1601 response
1602 .llm_response
1603 .as_agg()
1604 .map(completion_text)
1605 .unwrap_or_default(),
1606 "winner"
1607 );
1608 Ok(())
1609 }
1610
1611 #[tokio::test]
1612 async fn run_surfaces_a_terminal_error_with_many_calls_in_flight() -> Result<()> {
1613 use std::sync::atomic::{AtomicUsize, Ordering};
1614
1615 const N: usize = 10;
1618
1619 struct EnterThenPend {
1621 started: Arc<AtomicUsize>,
1622 all_started: Arc<tokio::sync::Notify>,
1623 n: usize,
1624 }
1625
1626 #[async_trait]
1627 impl RoutedLlmClient for EnterThenPend {
1628 async fn call(
1629 &self,
1630 _ctx: Context,
1631 _request: Request,
1632 _decision: Arc<dyn Decision>,
1633 ) -> std::result::Result<Response, LlmClientError> {
1634 if self.started.fetch_add(1, Ordering::SeqCst) + 1 == self.n {
1635 self.all_started.notify_one();
1636 }
1637 std::future::pending::<()>().await;
1638 unreachable!()
1639 }
1640 }
1641
1642 struct FanOutThenError {
1645 target: LlmTarget,
1646 all_started: Arc<tokio::sync::Notify>,
1647 n: usize,
1648 }
1649
1650 #[async_trait]
1651 impl Algorithm for FanOutThenError {
1652 fn name(&self) -> &str {
1653 "fan_out_then_error"
1654 }
1655
1656 async fn create_run_task(
1657 self: Arc<Self>,
1658 ctx: Context,
1659 driver: Driver,
1660 request: Request,
1661 ) -> Result<Response> {
1662 let offloads = futures::future::join_all((0..self.n).map(|i| {
1663 let decision: Arc<dyn Decision> = Arc::new(TestDecision {
1664 model: format!("m{i}"),
1665 });
1666 driver.call_llm_target(ctx.clone(), &self.target, request.clone(), decision)
1667 }));
1668 tokio::select! {
1669 _ = offloads => Err(test_error("offloads unexpectedly completed")),
1670 _ = self.all_started.notified() => {
1671 Err(test_error("terminal error while calls pending"))
1672 }
1673 }
1674 }
1675 }
1676
1677 let all_started = Arc::new(tokio::sync::Notify::new());
1678 let target = LlmTarget {
1679 semantic_name: "pending".to_string(),
1680 llm_client: Some(Arc::new(EnterThenPend {
1681 started: Arc::new(AtomicUsize::new(0)),
1682 all_started: all_started.clone(),
1683 n: N,
1684 })),
1685 };
1686 let algo: Arc<dyn Algorithm> = Arc::new(FanOutThenError {
1687 target,
1688 all_started,
1689 n: N,
1690 });
1691
1692 let run = algo.run(Context::default(), request());
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}