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 pub fn count_tokens_client(&self) -> Option<Arc<dyn RoutedLlmClient>> {
413 self.targets.iter().find_map(|target| {
414 target
415 .llm_client
416 .as_ref()
417 .filter(|client| client.supports_count_tokens())
418 .cloned()
419 })
420 }
421}
422
423const MAX_EVICTION_SESSIONS: usize = 1_024;
426
427#[derive(Default)]
433pub(crate) struct SessionEvictions {
434 sessions: Mutex<HashMap<String, HashSet<String>>>,
435}
436
437impl SessionEvictions {
438 fn evicted_in(&self, session: Option<&str>) -> Vec<String> {
440 let Some(session) = session else {
441 return Vec::new();
442 };
443 self.sessions
444 .lock()
445 .get(session)
446 .map(|targets| targets.iter().cloned().collect())
447 .unwrap_or_default()
448 }
449
450 fn record(&self, session: Option<&str>, target: &str) {
453 let Some(session) = session else { return };
454 let mut sessions = self.sessions.lock();
455 if sessions.len() >= MAX_EVICTION_SESSIONS
456 && !sessions.contains_key(session)
457 && let Some(oldest) = sessions.keys().next().cloned()
458 {
459 sessions.remove(&oldest);
460 }
461 sessions
462 .entry(session.to_string())
463 .or_default()
464 .insert(target.to_string());
465 }
466}
467
468fn eligible_targets(targets: &LlmTargetSet, ctx: &Context) -> usize {
470 targets
471 .targets()
472 .iter()
473 .filter(|t| !ctx.is_excluded(&t.semantic_name))
474 .count()
475}
476
477pub(crate) fn exclude_evicted(
480 ctx: &mut Context,
481 targets: &LlmTargetSet,
482 evictions: &SessionEvictions,
483 session: Option<&str>,
484) {
485 for target in evictions.evicted_in(session) {
486 if eligible_targets(targets, ctx) <= 1 {
489 break;
490 }
491 ctx.exclude_target(target);
492 }
493}
494
495#[allow(clippy::too_many_arguments)]
503pub(crate) async fn call_llm_with_overflow_fallback(
504 mut ctx: Context,
505 driver: &Driver,
506 targets: &LlmTargetSet,
507 mut target: LlmTarget,
508 mut decision: Arc<dyn Decision>,
509 request: Request,
510 session: Option<&str>,
511 evictions: &SessionEvictions,
512 fallback_decision: impl Fn(&LlmTarget, &LlmTarget) -> Arc<dyn Decision>,
513) -> Result<Response> {
514 loop {
515 let result = driver
516 .call_llm_target(ctx.clone(), &target, request.clone(), decision.clone())
517 .await;
518 let Err(error) = result else { return result };
519 let LibsyError::ClientCall {
520 target: failed,
521 source: LlmClientError::ContextWindowExceeded { .. },
522 } = &error
523 else {
524 return Err(error);
525 };
526 if !ctx.exclude_target(failed) {
529 return Err(error);
530 }
531 evictions.record(session, failed);
532 let Ok(next) = targets.resolve_target(&target.semantic_name, &ctx) else {
533 return Err(error);
534 };
535 decision = fallback_decision(&target, &next);
536 target = next;
537 driver.info(ctx.clone(), decision.clone()).await?;
538 }
539}
540
541#[async_trait]
562pub trait Algorithm: Send + Sync + 'static {
563 fn name(&self) -> &str;
567
568 async fn create_run_task(
574 self: Arc<Self>,
575 ctx: Context,
576 driver: Driver,
577 request: Request,
578 ) -> Result<Response>;
579
580 #[allow(unused_variables)]
584 async fn process_signals(self: Arc<Self>, signals: Signals) -> Result<()> {
585 Ok(())
586 }
587
588 fn count_tokens_client(&self) -> Option<Arc<dyn RoutedLlmClient>> {
598 None
599 }
600
601 async fn count_tokens(&self, request: Request) -> Result<serde_json::Value> {
609 let client = self
610 .count_tokens_client()
611 .ok_or_else(|| LibsyError::AlgorithmError {
612 message: "no target supports count_tokens (needs an Anthropic upstream)"
613 .to_string(),
614 })?;
615 client
616 .count_tokens(request)
617 .await
618 .map_err(|source| LibsyError::client_call("count_tokens", source))
619 }
620
621 fn run_stream(
631 self: Arc<Self>,
632 ctx: Context,
633 request: Request,
634 observer: Option<RunObserver>,
635 ) -> StepStream {
636 let mut ctx = ctx;
639 ctx.values.insert(
640 observability::ALGORITHM_KEY.to_string(),
641 self.name().to_string(),
642 );
643 let driver = Driver::with_observer(observer);
644 let task_driver = driver.clone();
645 let task_ctx = ctx.clone();
646 let stream = task_driver.stream();
647 let span = observability::run_span(self.name(), &request);
651 let observed_driver = task_driver.clone();
652 let handle = tokio::spawn(
653 async move {
654 observability::observe_run(
655 task_ctx.clone(),
656 observed_driver,
657 self.create_run_task(task_ctx, task_driver, request),
658 )
659 .await
660 }
661 .instrument(span),
662 );
663 let abort_guard = AbortOnDrop(handle.abort_handle());
665
666 let finish_driver = driver.clone();
667 let finish_ctx = ctx;
668 let tail: StepStream = Box::pin(
669 futures::stream::once(async move {
670 let result = match handle.await {
671 Ok(response) => response,
672 Err(source) => Err(LibsyError::AlgorithmTask { source }),
673 };
674 finish_driver.finish(finish_ctx, result).await
675 })
676 .filter_map(|finish_result| async move { finish_result.err().map(Err) }),
677 );
678
679 let stream: StepStream = Box::pin(stream);
680 Box::pin(futures::stream::select(stream, tail).map(move |step| {
681 let _keep_alive = &abort_guard;
683 step
684 }))
685 }
686
687 async fn run(
691 self: Arc<Self>,
692 ctx: Context,
693 request: Request,
694 ) -> Result<(Vec<Arc<dyn Decision>>, Response)> {
695 self.run_observed(ctx, request, None).await
696 }
697
698 async fn run_observed(
700 self: Arc<Self>,
701 ctx: Context,
702 request: Request,
703 observer: Option<RunObserver>,
704 ) -> Result<(Vec<Arc<dyn Decision>>, Response)> {
705 #[tracing::instrument(
711 target = "libsy",
712 name = "libsy.client_call",
713 skip_all,
714 fields(
715 algorithm = observability::algorithm_label(&call.get_routed().ctx),
716 switchyard.algorithm = observability::algorithm_label(&call.get_routed().ctx),
717 switchyard.routing.tier = tracing::field::Empty,
718 selected_model = call.get_decision().selected_model(),
719 otel.kind = "client",
720 otel.name = %format_args!("chat {}", call.get_decision().selected_model()),
721 openinference.span.kind = "LLM",
722 gen_ai.operation.name = "chat",
723 gen_ai.request.model = call.get_decision().selected_model(),
724 gen_ai.request.stream = tracing::field::Empty,
725 gen_ai.request.temperature = tracing::field::Empty,
726 gen_ai.request.top_p = tracing::field::Empty,
727 gen_ai.request.top_k = tracing::field::Empty,
728 gen_ai.request.max_tokens = tracing::field::Empty,
729 gen_ai.request.reasoning.level = tracing::field::Empty,
730 gen_ai.output.type = tracing::field::Empty,
731 gen_ai.conversation.id = tracing::field::Empty,
732 server.address = tracing::field::Empty,
733 server.port = tracing::field::Empty,
734 gen_ai.response.id = tracing::field::Empty,
735 gen_ai.response.model = tracing::field::Empty,
736 gen_ai.usage.input_tokens = tracing::field::Empty,
737 gen_ai.usage.output_tokens = tracing::field::Empty,
738 gen_ai.usage.cache_read.input_tokens = tracing::field::Empty,
739 gen_ai.usage.cache_creation.input_tokens = tracing::field::Empty,
740 gen_ai.usage.reasoning.output_tokens = tracing::field::Empty,
741 outcome = tracing::field::Empty,
742 otel.status_code = tracing::field::Empty,
743 error.type = tracing::field::Empty,
744 error = tracing::field::Empty,
745 )
746 )]
747 async fn serve(call: CallLlmRequest) -> Result<()> {
748 let span = tracing::Span::current();
749 observability::record_gen_ai_request(&span, &call.get_routed().request.llm_request);
750 if let Some(tier) = call.get_decision().routing_tier() {
751 span.record("switchyard.routing.tier", tier);
752 }
753 if let Some(session_id) = call
754 .get_routed()
755 .request
756 .metadata
757 .as_ref()
758 .and_then(|metadata| metadata.session_id.as_deref())
759 {
760 span.record("gen_ai.conversation.id", session_id);
761 }
762 let routed = call.get_routed().clone();
763 let target = routed.decision.selected_model().to_string();
764 let client =
765 routed
766 .default_client
767 .clone()
768 .ok_or_else(|| LibsyError::MissingClient {
769 target: target.clone(),
770 })?;
771 let result = client
772 .call(routed.ctx, routed.request, routed.decision)
773 .await
774 .map_err(|source| LibsyError::client_call(target, source));
775 let result = observability::observe_client_call(result);
776 call.respond(result)
777 }
778
779 let stream = self.run_stream(ctx, request, observer);
780 tokio::pin!(stream);
781
782 let mut trace: Vec<Arc<dyn Decision>> = Vec::new();
783 let mut in_flight = futures::stream::FuturesUnordered::new();
784 let mut final_response: Option<Response> = None;
785
786 loop {
787 tokio::select! {
788 Some(result) = in_flight.next() => match result {
789 Ok(()) => {}, Err(err) => return Err(err), },
792 step = stream.next() => {
793 match step {
794 None => break, Some(item) => match item? {
796 Step::CallLlm(call) => in_flight.push(serve(*call)),
797 Step::Decision(decision) => trace.push(decision),
798 Step::ReturnToAgent(response) => {
799 final_response = Some(*response);
800 break;
801 }
802 }
803 }
804 },
805 }
806 }
807 final_response
808 .map(|response| (trace, response))
809 .ok_or(LibsyError::MissingFinalResponse)
810 }
811}
812
813#[cfg(test)]
814mod tests {
815 use super::*;
816 use futures::StreamExt;
817 use switchyard_protocol::{
818 LlmResponse, LlmResponseChunk, completion_text, text_request, text_response,
819 };
820
821 #[derive(Debug, thiserror::Error)]
822 #[error("{0}")]
823 struct TestError(&'static str);
824
825 fn test_error(message: &'static str) -> LibsyError {
826 LibsyError::external("test", TestError(message))
827 }
828
829 struct EchoClient;
831
832 #[async_trait]
833 impl RoutedLlmClient for EchoClient {
834 async fn call(
835 &self,
836 _ctx: Context,
837 _request: Request,
838 decision: Arc<dyn Decision>,
839 ) -> std::result::Result<Response, LlmClientError> {
840 Ok(Response {
842 llm_response: LlmResponse::Agg(text_response(
843 None,
844 decision.selected_model().to_string(),
845 )),
846 metadata: None,
847 })
848 }
849 }
850
851 struct TestDecision {
854 model: String,
855 }
856
857 impl Decision for TestDecision {
858 fn selected_model(&self) -> &str {
859 &self.model
860 }
861 fn reasoning(&self) -> Option<&str> {
862 None
863 }
864 fn as_any(&self) -> &dyn std::any::Any {
865 self
866 }
867 }
868
869 struct TestAlgo {
870 target_set: LlmTargetSet,
871 }
872
873 #[async_trait]
874 impl Algorithm for TestAlgo {
875 fn name(&self) -> &str {
876 "test"
877 }
878
879 async fn create_run_task(
880 self: Arc<Self>,
881 ctx: Context,
882 driver: Driver,
883 request: Request,
884 ) -> Result<Response> {
885 let target = self
886 .target_set
887 .targets()
888 .first()
889 .ok_or(LibsyError::NoTargets)?
890 .clone();
891 let decision: Arc<dyn Decision> = Arc::new(TestDecision {
892 model: target.semantic_name.clone(),
893 });
894 driver.info(ctx.clone(), decision.clone()).await?;
895 driver
896 .call_llm_target(ctx, &target, request, decision)
897 .await
898 }
899 }
900
901 fn orch(target_set: LlmTargetSet) -> Arc<dyn Algorithm> {
903 Arc::new(TestAlgo { target_set })
904 }
905
906 fn request() -> Request {
907 Request {
908 llm_request: text_request(Some("auto".to_string()), "hi".to_string()),
909 raw_request: None,
910 metadata: None,
911 }
912 }
913
914 fn target_set(names: &[(&str, bool)]) -> LlmTargetSet {
916 let targets = names
917 .iter()
918 .map(|(name, has_client)| LlmTarget {
919 semantic_name: name.to_string(),
920 llm_client: has_client.then(|| Arc::new(EchoClient) as Arc<dyn RoutedLlmClient>),
921 })
922 .collect();
923 LlmTargetSet::new(targets)
924 }
925
926 #[tokio::test]
927 async fn observed_run_reports_one_successful_routed_call() -> Result<()> {
928 let observations = Arc::new(Mutex::new(Vec::new()));
929 let observed = observations.clone();
930 let observer: RunObserver = Arc::new(move |observation| observed.lock().push(observation));
931 let (_, response) = orch(target_set(&[("direct/model", true)]))
932 .run_observed(Context::default(), request(), Some(observer))
933 .await?;
934 assert_eq!(
935 response.llm_response.as_agg().map(completion_text),
936 Some("direct/model".to_string())
937 );
938 let observations = observations.lock();
939 assert_eq!(observations.len(), 2);
940 let RunObservation::LlmCall(observation) = &observations[0] else {
941 return Err(test_error("expected an LLM call observation"));
942 };
943 assert_eq!(observation.selected_model, "direct/model");
944 assert!(observation.is_routed);
945 assert!(observation.is_success);
946 assert!(observation.usage.is_some());
947 assert!(matches!(
948 observations[1],
949 RunObservation::RoutingOverhead(_)
950 ));
951 Ok(())
952 }
953
954 #[test]
955 fn target_lookup_returns_the_missing_target() {
956 let error = target_set(&[]).get_target("missing").err();
957 assert!(matches!(
958 error,
959 Some(LibsyError::TargetNotFound { target }) if target == "missing"
960 ));
961 }
962
963 struct StreamingClient {
966 chunks: Vec<LlmResponseChunk>,
967 }
968
969 #[async_trait]
970 impl RoutedLlmClient for StreamingClient {
971 async fn call(
972 &self,
973 _ctx: Context,
974 _request: Request,
975 _decision: Arc<dyn Decision>,
976 ) -> std::result::Result<Response, LlmClientError> {
977 let stream = futures::stream::iter(
978 self.chunks
979 .clone()
980 .into_iter()
981 .map(|chunk| Ok(chunk.into())),
982 )
983 .boxed();
984 Ok(Response {
985 llm_response: LlmResponse::Stream(stream),
986 metadata: None,
987 })
988 }
989 }
990
991 fn streaming_orch(chunks: Vec<LlmResponseChunk>) -> Arc<dyn Algorithm> {
993 let target = LlmTarget {
994 semantic_name: "stream/model".to_string(),
995 llm_client: Some(Arc::new(StreamingClient { chunks }) as Arc<dyn RoutedLlmClient>),
996 };
997 orch(LlmTargetSet::new(vec![target]))
998 }
999
1000 #[tokio::test]
1001 async fn run_returns_a_streamed_response_the_caller_aggregates() -> Result<()> {
1002 let orch = streaming_orch(vec![
1005 LlmResponseChunk::MessageStart {
1006 id: Some("m1".to_string()),
1007 model: Some("stream/model".to_string()),
1008 },
1009 LlmResponseChunk::TextDelta {
1010 index: 0,
1011 text: "hel".to_string(),
1012 },
1013 LlmResponseChunk::TextDelta {
1014 index: 0,
1015 text: "lo".to_string(),
1016 },
1017 LlmResponseChunk::MessageStop {
1018 reason: Some("stop".to_string()),
1019 },
1020 ]);
1021 let (trace, response) = orch.run(Context::default(), request()).await?;
1022 let agg = response
1024 .llm_response
1025 .into_agg()
1026 .await
1027 .map_err(|error| LibsyError::external("aggregating response stream", error))?;
1028 assert_eq!(completion_text(&agg), "hello");
1029 assert_eq!(agg.model.as_deref(), Some("stream/model"));
1030 assert_eq!(trace.len(), 1);
1031 Ok(())
1032 }
1033
1034 #[tokio::test]
1035 async fn aggregating_a_streamed_response_propagates_a_mid_stream_error() -> Result<()> {
1036 let orch = streaming_orch(vec![
1039 LlmResponseChunk::TextDelta {
1040 index: 0,
1041 text: "partial".to_string(),
1042 },
1043 LlmResponseChunk::StreamError {
1044 message: "upstream exploded".to_string(),
1045 },
1046 ]);
1047 let (_, response) = orch.run(Context::default(), request()).await?;
1048 match response.llm_response.into_agg().await {
1049 Ok(_) => panic!("expected a mid-stream error, got an aggregate"),
1050 Err(err) => {
1051 assert!(err.to_string().contains("upstream exploded"));
1052 Ok(())
1053 }
1054 }
1055 }
1056
1057 #[tokio::test]
1058 async fn run_offloads_via_promise_then_returns_to_agent() -> Result<()> {
1059 let stream = orch(target_set(&[("offload/model", false)])).run_stream(
1062 Context::default(),
1063 request(),
1064 None,
1065 );
1066 tokio::pin!(stream);
1067
1068 let mut saw_call = false;
1069 let mut final_completion = None;
1070 while let Some(step) = stream.next().await {
1071 match step? {
1072 Step::CallLlm(call) => {
1073 saw_call = true;
1074 assert_eq!(call.get_decision().selected_model(), "offload/model");
1076 call.respond(Ok(Response {
1078 llm_response: LlmResponse::Agg(text_response(
1079 None,
1080 "fulfilled".to_string(),
1081 )),
1082 metadata: None,
1083 }))?;
1084 }
1085 Step::Decision(decision) => {
1086 assert_eq!(decision.selected_model(), "offload/model");
1087 }
1088 Step::ReturnToAgent(response) => {
1089 final_completion = Some(
1090 response
1091 .llm_response
1092 .as_agg()
1093 .map(completion_text)
1094 .unwrap_or_default(),
1095 );
1096 }
1097 }
1098 }
1099
1100 assert!(saw_call, "expected a CallLlm step before ReturnToAgent");
1101 assert_eq!(
1102 final_completion.ok_or_else(|| test_error("no ReturnToAgent step"))?,
1103 "fulfilled"
1104 );
1105 Ok(())
1106 }
1107
1108 #[tokio::test]
1109 async fn client_backed_target_offloads_with_a_default_client() -> Result<()> {
1110 let stream = orch(target_set(&[("direct/model", true)])).run_stream(
1113 Context::default(),
1114 request(),
1115 None,
1116 );
1117 tokio::pin!(stream);
1118
1119 let mut final_completion = None;
1120 while let Some(step) = stream.next().await {
1121 match step? {
1122 Step::CallLlm(call) => {
1123 let routed = call.get_routed().clone();
1124 let client = routed
1125 .default_client
1126 .clone()
1127 .ok_or_else(|| test_error("expected a default client"))?;
1128 let target = routed.decision.selected_model().to_string();
1129 let result = client
1130 .call(routed.ctx, routed.request, routed.decision)
1131 .await
1132 .map_err(|error| LibsyError::client_call(target, error));
1133 call.respond(result)?;
1134 }
1135 Step::Decision(_) => {}
1136 Step::ReturnToAgent(response) => {
1137 final_completion = Some(
1138 response
1139 .llm_response
1140 .as_agg()
1141 .map(completion_text)
1142 .unwrap_or_default(),
1143 );
1144 }
1145 }
1146 }
1147
1148 assert_eq!(
1150 final_completion.ok_or_else(|| test_error("no ReturnToAgent"))?,
1151 "direct/model"
1152 );
1153 Ok(())
1154 }
1155
1156 #[tokio::test]
1157 async fn run_returns_the_response_when_all_targets_have_clients() -> Result<()> {
1158 let (trace, response) = orch(target_set(&[("direct/model", true)]))
1161 .run(Context::default(), request())
1162 .await?;
1163 assert_eq!(
1165 response
1166 .llm_response
1167 .as_agg()
1168 .map(completion_text)
1169 .unwrap_or_default(),
1170 "direct/model"
1171 );
1172 assert_eq!(trace[0].selected_model(), "direct/model");
1173 Ok(())
1174 }
1175
1176 #[tokio::test]
1177 async fn run_errors_when_a_target_lacks_a_client() -> Result<()> {
1178 let error = orch(target_set(&[("offload/model", false)]))
1181 .run(Context::default(), request())
1182 .await
1183 .err()
1184 .ok_or_else(|| test_error("expected a missing-client error"))?;
1185 assert!(matches!(
1186 error,
1187 LibsyError::MissingClient { target } if target == "offload/model"
1188 ));
1189 Ok(())
1190 }
1191
1192 #[tokio::test(flavor = "multi_thread", worker_threads = 12)]
1193 async fn requests_are_processed_in_parallel() -> Result<()> {
1194 use std::time::Duration;
1195 use tokio::sync::Barrier;
1196
1197 const N: usize = 12;
1198
1199 struct BarrierClient {
1205 barrier: Arc<Barrier>,
1206 }
1207
1208 #[async_trait]
1209 impl RoutedLlmClient for BarrierClient {
1210 async fn call(
1211 &self,
1212 _ctx: Context,
1213 _request: Request,
1214 decision: Arc<dyn Decision>,
1215 ) -> std::result::Result<Response, LlmClientError> {
1216 self.barrier.wait().await;
1217 Ok(Response {
1218 llm_response: LlmResponse::Agg(text_response(
1219 None,
1220 decision.selected_model().to_string(),
1221 )),
1222 metadata: None,
1223 })
1224 }
1225 }
1226
1227 let barrier = Arc::new(Barrier::new(N));
1228 let targets = LlmTargetSet::new(vec![LlmTarget {
1229 semantic_name: "m".to_string(),
1230 llm_client: Some(Arc::new(BarrierClient {
1231 barrier: barrier.clone(),
1232 })),
1233 }]);
1234 let algo = orch(targets);
1236
1237 let mut handles = Vec::new();
1238 for _ in 0..N {
1239 let algo = algo.clone();
1240 handles.push(tokio::spawn(async move {
1241 algo.run(Context::default(), request())
1242 .await
1243 .map(|(_, response)| {
1244 response
1245 .llm_response
1246 .as_agg()
1247 .map(completion_text)
1248 .unwrap_or_default()
1249 })
1250 }));
1251 }
1252
1253 for handle in handles {
1254 let completion = tokio::time::timeout(Duration::from_secs(5), handle)
1256 .await
1257 .map_err(|error| LibsyError::external("waiting for test task", error))?
1258 .map_err(|source| LibsyError::AlgorithmTask { source })??;
1259 assert_eq!(completion, "m");
1260 }
1261 Ok(())
1262 }
1263
1264 #[tokio::test]
1265 async fn offload_error_propagates_back_to_the_algorithm() -> Result<()> {
1266 let stream = orch(target_set(&[("offload/model", false)])).run_stream(
1270 Context::default(),
1271 request(),
1272 None,
1273 );
1274 tokio::pin!(stream);
1275
1276 let mut saw_error = false;
1277 while let Some(step) = stream.next().await {
1278 match step {
1279 Ok(Step::CallLlm(call)) => {
1280 call.respond(Err(test_error("upstream model call failed")))?;
1281 }
1282 Ok(Step::Decision(_)) => {}
1283 Ok(Step::ReturnToAgent(..)) => {
1284 return Err(test_error(
1285 "expected the offload error to propagate, got a response",
1286 ));
1287 }
1288 Err(err) => {
1289 assert!(err.to_string().contains("upstream model call failed"));
1291 saw_error = true;
1292 }
1293 }
1294 }
1295
1296 assert!(saw_error, "expected an error step");
1297 Ok(())
1298 }
1299
1300 #[tokio::test]
1301 async fn dropping_the_stream_cancels_the_algorithm_task() -> Result<()> {
1302 use std::sync::atomic::{AtomicBool, Ordering};
1303 use std::time::Duration;
1304 use tokio::sync::mpsc;
1305
1306 struct DropGuard(Arc<AtomicBool>);
1309 impl Drop for DropGuard {
1310 fn drop(&mut self) {
1311 self.0.store(true, Ordering::SeqCst);
1312 }
1313 }
1314
1315 struct StuckAlgo {
1316 started: mpsc::UnboundedSender<()>,
1317 dropped: Arc<AtomicBool>,
1318 }
1319
1320 #[async_trait]
1321 impl Algorithm for StuckAlgo {
1322 fn name(&self) -> &str {
1323 "stuck"
1324 }
1325
1326 async fn create_run_task(
1327 self: Arc<Self>,
1328 _ctx: Context,
1329 _driver: Driver,
1330 _request: Request,
1331 ) -> Result<Response> {
1332 let _guard = DropGuard(self.dropped.clone());
1333 let _ = self.started.send(());
1334 std::future::pending::<()>().await;
1336 unreachable!()
1337 }
1338 }
1339
1340 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1341 let dropped = Arc::new(AtomicBool::new(false));
1342 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1343 started: started_tx,
1344 dropped: dropped.clone(),
1345 });
1346
1347 let stream = algo.run_stream(Context::default(), request(), None);
1348 started_rx
1349 .recv()
1350 .await
1351 .ok_or_else(|| test_error("task never started"))?;
1352 drop(stream);
1353 tokio::time::sleep(Duration::from_millis(100)).await;
1354
1355 assert!(
1356 dropped.load(Ordering::SeqCst),
1357 "algorithm task was NOT cancelled after dropping the stream"
1358 );
1359 Ok(())
1360 }
1361
1362 #[tokio::test]
1363 async fn create_run_task_panic_surfaces_as_a_stream_error() -> 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 let stream = algo.run_stream(Context::default(), request(), None);
1386 tokio::pin!(stream);
1387
1388 let mut saw_error = false;
1389 while let Some(step) = stream.next().await {
1390 match step {
1391 Err(err) => {
1392 assert!(matches!(err, LibsyError::AlgorithmTask { .. }));
1393 saw_error = true;
1394 }
1395 Ok(_) => return Err(test_error("expected the panic to surface as an error step")),
1396 }
1397 }
1398
1399 assert!(saw_error, "expected an error step from the panicked task");
1400 Ok(())
1401 }
1402
1403 #[tokio::test]
1404 async fn run_returns_an_error_when_the_algorithm_task_panics() -> Result<()> {
1405 struct Panicky;
1408
1409 #[async_trait]
1410 impl Algorithm for Panicky {
1411 fn name(&self) -> &str {
1412 "panicky"
1413 }
1414
1415 async fn create_run_task(
1416 self: Arc<Self>,
1417 _ctx: Context,
1418 _driver: Driver,
1419 _request: Request,
1420 ) -> Result<Response> {
1421 panic!("boom");
1422 }
1423 }
1424
1425 let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
1426 match algo.run(Context::default(), request()).await {
1427 Ok(_) => Err(test_error(
1428 "expected run to surface the algorithm panic as an error",
1429 )),
1430 Err(err) => {
1431 assert!(matches!(err, LibsyError::AlgorithmTask { .. }));
1432 Ok(())
1433 }
1434 }
1435 }
1436
1437 #[tokio::test]
1438 async fn cancelling_run_cancels_the_algorithm_task() -> Result<()> {
1439 use std::sync::atomic::{AtomicBool, Ordering};
1440 use std::time::Duration;
1441 use tokio::sync::mpsc;
1442
1443 struct DropGuard(Arc<AtomicBool>);
1446 impl Drop for DropGuard {
1447 fn drop(&mut self) {
1448 self.0.store(true, Ordering::SeqCst);
1449 }
1450 }
1451
1452 struct StuckAlgo {
1453 started: mpsc::UnboundedSender<()>,
1454 dropped: Arc<AtomicBool>,
1455 }
1456
1457 #[async_trait]
1458 impl Algorithm for StuckAlgo {
1459 fn name(&self) -> &str {
1460 "stuck"
1461 }
1462
1463 async fn create_run_task(
1464 self: Arc<Self>,
1465 _ctx: Context,
1466 _driver: Driver,
1467 _request: Request,
1468 ) -> Result<Response> {
1469 let _guard = DropGuard(self.dropped.clone());
1470 let _ = self.started.send(());
1471 std::future::pending::<()>().await;
1474 unreachable!()
1475 }
1476 }
1477
1478 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1479 let dropped = Arc::new(AtomicBool::new(false));
1480 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1481 started: started_tx,
1482 dropped: dropped.clone(),
1483 });
1484
1485 let run_task = tokio::spawn(async move { algo.run(Context::default(), request()).await });
1488 started_rx
1489 .recv()
1490 .await
1491 .ok_or_else(|| test_error("task never started"))?;
1492 run_task.abort();
1493 tokio::time::sleep(Duration::from_millis(100)).await;
1494
1495 assert!(
1496 dropped.load(Ordering::SeqCst),
1497 "algorithm task was NOT cancelled after cancelling run"
1498 );
1499 Ok(())
1500 }
1501
1502 struct LoserClient {
1508 started: Arc<tokio::sync::Notify>,
1509 delay: Option<std::time::Duration>,
1510 }
1511
1512 #[async_trait]
1513 impl RoutedLlmClient for LoserClient {
1514 async fn call(
1515 &self,
1516 _ctx: Context,
1517 _request: Request,
1518 decision: Arc<dyn Decision>,
1519 ) -> std::result::Result<Response, LlmClientError> {
1520 self.started.notify_one();
1521 match self.delay {
1522 Some(delay) => tokio::time::sleep(delay).await,
1523 None => std::future::pending::<()>().await,
1524 }
1525 Ok(Response {
1526 llm_response: LlmResponse::Agg(text_response(
1527 None,
1528 decision.selected_model().to_string(),
1529 )),
1530 metadata: None,
1531 })
1532 }
1533 }
1534
1535 struct GatedEchoClient {
1538 gate: Arc<tokio::sync::Notify>,
1539 }
1540
1541 #[async_trait]
1542 impl RoutedLlmClient for GatedEchoClient {
1543 async fn call(
1544 &self,
1545 _ctx: Context,
1546 _request: Request,
1547 decision: Arc<dyn Decision>,
1548 ) -> std::result::Result<Response, LlmClientError> {
1549 self.gate.notified().await;
1550 Ok(Response {
1551 llm_response: LlmResponse::Agg(text_response(
1552 None,
1553 decision.selected_model().to_string(),
1554 )),
1555 metadata: None,
1556 })
1557 }
1558 }
1559
1560 struct Hedge {
1563 winner: LlmTarget,
1564 loser: LlmTarget,
1565 }
1566
1567 #[async_trait]
1568 impl Algorithm for Hedge {
1569 fn name(&self) -> &str {
1570 "hedge"
1571 }
1572
1573 async fn create_run_task(
1574 self: Arc<Self>,
1575 ctx: Context,
1576 driver: Driver,
1577 request: Request,
1578 ) -> Result<Response> {
1579 let dec_w: Arc<dyn Decision> = Arc::new(TestDecision {
1580 model: self.winner.semantic_name.clone(),
1581 });
1582 let dec_l: Arc<dyn Decision> = Arc::new(TestDecision {
1583 model: self.loser.semantic_name.clone(),
1584 });
1585 let win = driver.call_llm_target(ctx.clone(), &self.winner, request.clone(), dec_w);
1586 let lose = driver.call_llm_target(ctx, &self.loser, request, dec_l);
1587 tokio::select! {
1589 res = win => res,
1590 res = lose => res,
1591 }
1592 }
1593 }
1594
1595 fn hedge(loser_delay: Option<std::time::Duration>) -> Arc<dyn Algorithm> {
1598 let started = Arc::new(tokio::sync::Notify::new());
1599 let winner = LlmTarget {
1600 semantic_name: "winner".to_string(),
1601 llm_client: Some(Arc::new(GatedEchoClient {
1602 gate: started.clone(),
1603 })),
1604 };
1605 let loser = LlmTarget {
1606 semantic_name: "loser".to_string(),
1607 llm_client: Some(Arc::new(LoserClient {
1608 started,
1609 delay: loser_delay,
1610 })),
1611 };
1612 Arc::new(Hedge { winner, loser })
1613 }
1614
1615 #[tokio::test]
1616 async fn run_returns_the_winner_without_a_late_loser_overwriting_it() -> Result<()> {
1617 let (_trace, response) = hedge(Some(std::time::Duration::from_millis(50)))
1620 .run(Context::default(), request())
1621 .await?;
1622 assert_eq!(
1623 response
1624 .llm_response
1625 .as_agg()
1626 .map(completion_text)
1627 .unwrap_or_default(),
1628 "winner"
1629 );
1630 Ok(())
1631 }
1632
1633 #[tokio::test]
1634 async fn run_returns_the_winner_without_hanging_on_a_pending_loser() -> Result<()> {
1635 let run = hedge(None).run(Context::default(), request());
1638 let (_trace, response) = tokio::time::timeout(std::time::Duration::from_secs(1), run)
1639 .await
1640 .map_err(|error| LibsyError::external("waiting for pending loser", error))??;
1641 assert_eq!(
1642 response
1643 .llm_response
1644 .as_agg()
1645 .map(completion_text)
1646 .unwrap_or_default(),
1647 "winner"
1648 );
1649 Ok(())
1650 }
1651
1652 #[tokio::test]
1653 async fn run_surfaces_a_terminal_error_with_many_calls_in_flight() -> Result<()> {
1654 use std::sync::atomic::{AtomicUsize, Ordering};
1655
1656 const N: usize = 10;
1659
1660 struct EnterThenPend {
1662 started: Arc<AtomicUsize>,
1663 all_started: Arc<tokio::sync::Notify>,
1664 n: usize,
1665 }
1666
1667 #[async_trait]
1668 impl RoutedLlmClient for EnterThenPend {
1669 async fn call(
1670 &self,
1671 _ctx: Context,
1672 _request: Request,
1673 _decision: Arc<dyn Decision>,
1674 ) -> std::result::Result<Response, LlmClientError> {
1675 if self.started.fetch_add(1, Ordering::SeqCst) + 1 == self.n {
1676 self.all_started.notify_one();
1677 }
1678 std::future::pending::<()>().await;
1679 unreachable!()
1680 }
1681 }
1682
1683 struct FanOutThenError {
1686 target: LlmTarget,
1687 all_started: Arc<tokio::sync::Notify>,
1688 n: usize,
1689 }
1690
1691 #[async_trait]
1692 impl Algorithm for FanOutThenError {
1693 fn name(&self) -> &str {
1694 "fan_out_then_error"
1695 }
1696
1697 async fn create_run_task(
1698 self: Arc<Self>,
1699 ctx: Context,
1700 driver: Driver,
1701 request: Request,
1702 ) -> Result<Response> {
1703 let offloads = futures::future::join_all((0..self.n).map(|i| {
1704 let decision: Arc<dyn Decision> = Arc::new(TestDecision {
1705 model: format!("m{i}"),
1706 });
1707 driver.call_llm_target(ctx.clone(), &self.target, request.clone(), decision)
1708 }));
1709 tokio::select! {
1710 _ = offloads => Err(test_error("offloads unexpectedly completed")),
1711 _ = self.all_started.notified() => {
1712 Err(test_error("terminal error while calls pending"))
1713 }
1714 }
1715 }
1716 }
1717
1718 let all_started = Arc::new(tokio::sync::Notify::new());
1719 let target = LlmTarget {
1720 semantic_name: "pending".to_string(),
1721 llm_client: Some(Arc::new(EnterThenPend {
1722 started: Arc::new(AtomicUsize::new(0)),
1723 all_started: all_started.clone(),
1724 n: N,
1725 })),
1726 };
1727 let algo: Arc<dyn Algorithm> = Arc::new(FanOutThenError {
1728 target,
1729 all_started,
1730 n: N,
1731 });
1732
1733 let run = algo.run(Context::default(), request());
1736 let result = tokio::time::timeout(std::time::Duration::from_millis(500), run)
1737 .await
1738 .map_err(|error| {
1739 LibsyError::external("waiting for terminal error with full call cap", error)
1740 })?;
1741 match result {
1742 Ok(_) => Err(test_error("expected the terminal error, got a response")),
1743 Err(err) => {
1744 assert!(
1745 err.to_string()
1746 .contains("terminal error while calls pending")
1747 );
1748 Ok(())
1749 }
1750 }
1751 }
1752}