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 crate::text::{completion_text, text_request, text_response};
817 use futures::StreamExt;
818 use switchyard_protocol::{LlmResponse, LlmResponseChunk};
819
820 #[derive(Debug, thiserror::Error)]
821 #[error("{0}")]
822 struct TestError(&'static str);
823
824 fn test_error(message: &'static str) -> LibsyError {
825 LibsyError::external("test", TestError(message))
826 }
827
828 struct EchoClient;
830
831 #[async_trait]
832 impl RoutedLlmClient for EchoClient {
833 async fn call(
834 &self,
835 _ctx: Context,
836 _request: Request,
837 decision: Arc<dyn Decision>,
838 ) -> std::result::Result<Response, LlmClientError> {
839 Ok(Response {
841 llm_response: LlmResponse::Agg(text_response(
842 None,
843 decision.selected_model().to_string(),
844 )),
845 metadata: None,
846 })
847 }
848 }
849
850 struct TestDecision {
853 model: String,
854 }
855
856 impl Decision for TestDecision {
857 fn selected_model(&self) -> &str {
858 &self.model
859 }
860 fn reasoning(&self) -> Option<&str> {
861 None
862 }
863 fn as_any(&self) -> &dyn std::any::Any {
864 self
865 }
866 }
867
868 struct TestAlgo {
869 target_set: LlmTargetSet,
870 }
871
872 #[async_trait]
873 impl Algorithm for TestAlgo {
874 fn name(&self) -> &str {
875 "test"
876 }
877
878 async fn create_run_task(
879 self: Arc<Self>,
880 ctx: Context,
881 driver: Driver,
882 request: Request,
883 ) -> Result<Response> {
884 let target = self
885 .target_set
886 .targets()
887 .first()
888 .ok_or(LibsyError::NoTargets)?
889 .clone();
890 let decision: Arc<dyn Decision> = Arc::new(TestDecision {
891 model: target.semantic_name.clone(),
892 });
893 driver.info(ctx.clone(), decision.clone()).await?;
894 driver
895 .call_llm_target(ctx, &target, request, decision)
896 .await
897 }
898 }
899
900 fn orch(target_set: LlmTargetSet) -> Arc<dyn Algorithm> {
902 Arc::new(TestAlgo { target_set })
903 }
904
905 fn request() -> Request {
906 Request {
907 llm_request: text_request(Some("auto".to_string()), "hi".to_string()),
908 raw_request: None,
909 metadata: None,
910 }
911 }
912
913 fn target_set(names: &[(&str, bool)]) -> LlmTargetSet {
915 let targets = names
916 .iter()
917 .map(|(name, has_client)| LlmTarget {
918 semantic_name: name.to_string(),
919 llm_client: has_client.then(|| Arc::new(EchoClient) as Arc<dyn RoutedLlmClient>),
920 })
921 .collect();
922 LlmTargetSet::new(targets)
923 }
924
925 #[tokio::test]
926 async fn observed_run_reports_one_successful_routed_call() -> Result<()> {
927 let observations = Arc::new(Mutex::new(Vec::new()));
928 let observed = observations.clone();
929 let observer: RunObserver = Arc::new(move |observation| observed.lock().push(observation));
930 let (_, response) = orch(target_set(&[("direct/model", true)]))
931 .run_observed(Context::default(), request(), Some(observer))
932 .await?;
933 assert_eq!(
934 response.llm_response.as_agg().map(completion_text),
935 Some("direct/model".to_string())
936 );
937 let observations = observations.lock();
938 assert_eq!(observations.len(), 2);
939 let RunObservation::LlmCall(observation) = &observations[0] else {
940 return Err(test_error("expected an LLM call observation"));
941 };
942 assert_eq!(observation.selected_model, "direct/model");
943 assert!(observation.is_routed);
944 assert!(observation.is_success);
945 assert!(observation.usage.is_some());
946 assert!(matches!(
947 observations[1],
948 RunObservation::RoutingOverhead(_)
949 ));
950 Ok(())
951 }
952
953 #[test]
954 fn target_lookup_returns_the_missing_target() {
955 let error = target_set(&[]).get_target("missing").err();
956 assert!(matches!(
957 error,
958 Some(LibsyError::TargetNotFound { target }) if target == "missing"
959 ));
960 }
961
962 struct StreamingClient {
965 chunks: Vec<LlmResponseChunk>,
966 }
967
968 #[async_trait]
969 impl RoutedLlmClient for StreamingClient {
970 async fn call(
971 &self,
972 _ctx: Context,
973 _request: Request,
974 _decision: Arc<dyn Decision>,
975 ) -> std::result::Result<Response, LlmClientError> {
976 let stream = futures::stream::iter(
977 self.chunks
978 .clone()
979 .into_iter()
980 .map(|chunk| Ok(chunk.into())),
981 )
982 .boxed();
983 Ok(Response {
984 llm_response: LlmResponse::Stream(stream),
985 metadata: None,
986 })
987 }
988 }
989
990 fn streaming_orch(chunks: Vec<LlmResponseChunk>) -> Arc<dyn Algorithm> {
992 let target = LlmTarget {
993 semantic_name: "stream/model".to_string(),
994 llm_client: Some(Arc::new(StreamingClient { chunks }) as Arc<dyn RoutedLlmClient>),
995 };
996 orch(LlmTargetSet::new(vec![target]))
997 }
998
999 #[tokio::test]
1000 async fn run_returns_a_streamed_response_the_caller_aggregates() -> Result<()> {
1001 let orch = streaming_orch(vec![
1004 LlmResponseChunk::MessageStart {
1005 id: Some("m1".to_string()),
1006 model: Some("stream/model".to_string()),
1007 },
1008 LlmResponseChunk::TextDelta {
1009 index: 0,
1010 text: "hel".to_string(),
1011 },
1012 LlmResponseChunk::TextDelta {
1013 index: 0,
1014 text: "lo".to_string(),
1015 },
1016 LlmResponseChunk::MessageStop {
1017 reason: Some("stop".to_string()),
1018 },
1019 ]);
1020 let (trace, response) = orch.run(Context::default(), request()).await?;
1021 let agg = response
1023 .llm_response
1024 .into_agg()
1025 .await
1026 .map_err(|error| LibsyError::external("aggregating response stream", error))?;
1027 assert_eq!(completion_text(&agg), "hello");
1028 assert_eq!(agg.model.as_deref(), Some("stream/model"));
1029 assert_eq!(trace.len(), 1);
1030 Ok(())
1031 }
1032
1033 #[tokio::test]
1034 async fn aggregating_a_streamed_response_propagates_a_mid_stream_error() -> Result<()> {
1035 let orch = streaming_orch(vec![
1038 LlmResponseChunk::TextDelta {
1039 index: 0,
1040 text: "partial".to_string(),
1041 },
1042 LlmResponseChunk::StreamError {
1043 message: "upstream exploded".to_string(),
1044 },
1045 ]);
1046 let (_, response) = orch.run(Context::default(), request()).await?;
1047 match response.llm_response.into_agg().await {
1048 Ok(_) => panic!("expected a mid-stream error, got an aggregate"),
1049 Err(err) => {
1050 assert!(err.to_string().contains("upstream exploded"));
1051 Ok(())
1052 }
1053 }
1054 }
1055
1056 #[tokio::test]
1057 async fn run_offloads_via_promise_then_returns_to_agent() -> Result<()> {
1058 let stream = orch(target_set(&[("offload/model", false)])).run_stream(
1061 Context::default(),
1062 request(),
1063 None,
1064 );
1065 tokio::pin!(stream);
1066
1067 let mut saw_call = false;
1068 let mut final_completion = None;
1069 while let Some(step) = stream.next().await {
1070 match step? {
1071 Step::CallLlm(call) => {
1072 saw_call = true;
1073 assert_eq!(call.get_decision().selected_model(), "offload/model");
1075 call.respond(Ok(Response {
1077 llm_response: LlmResponse::Agg(text_response(
1078 None,
1079 "fulfilled".to_string(),
1080 )),
1081 metadata: None,
1082 }))?;
1083 }
1084 Step::Decision(decision) => {
1085 assert_eq!(decision.selected_model(), "offload/model");
1086 }
1087 Step::ReturnToAgent(response) => {
1088 final_completion = Some(
1089 response
1090 .llm_response
1091 .as_agg()
1092 .map(completion_text)
1093 .unwrap_or_default(),
1094 );
1095 }
1096 }
1097 }
1098
1099 assert!(saw_call, "expected a CallLlm step before ReturnToAgent");
1100 assert_eq!(
1101 final_completion.ok_or_else(|| test_error("no ReturnToAgent step"))?,
1102 "fulfilled"
1103 );
1104 Ok(())
1105 }
1106
1107 #[tokio::test]
1108 async fn client_backed_target_offloads_with_a_default_client() -> Result<()> {
1109 let stream = orch(target_set(&[("direct/model", true)])).run_stream(
1112 Context::default(),
1113 request(),
1114 None,
1115 );
1116 tokio::pin!(stream);
1117
1118 let mut final_completion = None;
1119 while let Some(step) = stream.next().await {
1120 match step? {
1121 Step::CallLlm(call) => {
1122 let routed = call.get_routed().clone();
1123 let client = routed
1124 .default_client
1125 .clone()
1126 .ok_or_else(|| test_error("expected a default client"))?;
1127 let target = routed.decision.selected_model().to_string();
1128 let result = client
1129 .call(routed.ctx, routed.request, routed.decision)
1130 .await
1131 .map_err(|error| LibsyError::client_call(target, error));
1132 call.respond(result)?;
1133 }
1134 Step::Decision(_) => {}
1135 Step::ReturnToAgent(response) => {
1136 final_completion = Some(
1137 response
1138 .llm_response
1139 .as_agg()
1140 .map(completion_text)
1141 .unwrap_or_default(),
1142 );
1143 }
1144 }
1145 }
1146
1147 assert_eq!(
1149 final_completion.ok_or_else(|| test_error("no ReturnToAgent"))?,
1150 "direct/model"
1151 );
1152 Ok(())
1153 }
1154
1155 #[tokio::test]
1156 async fn run_returns_the_response_when_all_targets_have_clients() -> Result<()> {
1157 let (trace, response) = orch(target_set(&[("direct/model", true)]))
1160 .run(Context::default(), request())
1161 .await?;
1162 assert_eq!(
1164 response
1165 .llm_response
1166 .as_agg()
1167 .map(completion_text)
1168 .unwrap_or_default(),
1169 "direct/model"
1170 );
1171 assert_eq!(trace[0].selected_model(), "direct/model");
1172 Ok(())
1173 }
1174
1175 #[tokio::test]
1176 async fn run_errors_when_a_target_lacks_a_client() -> Result<()> {
1177 let error = orch(target_set(&[("offload/model", false)]))
1180 .run(Context::default(), request())
1181 .await
1182 .err()
1183 .ok_or_else(|| test_error("expected a missing-client error"))?;
1184 assert!(matches!(
1185 error,
1186 LibsyError::MissingClient { target } if target == "offload/model"
1187 ));
1188 Ok(())
1189 }
1190
1191 #[tokio::test(flavor = "multi_thread", worker_threads = 12)]
1192 async fn requests_are_processed_in_parallel() -> Result<()> {
1193 use std::time::Duration;
1194 use tokio::sync::Barrier;
1195
1196 const N: usize = 12;
1197
1198 struct BarrierClient {
1204 barrier: Arc<Barrier>,
1205 }
1206
1207 #[async_trait]
1208 impl RoutedLlmClient for BarrierClient {
1209 async fn call(
1210 &self,
1211 _ctx: Context,
1212 _request: Request,
1213 decision: Arc<dyn Decision>,
1214 ) -> std::result::Result<Response, LlmClientError> {
1215 self.barrier.wait().await;
1216 Ok(Response {
1217 llm_response: LlmResponse::Agg(text_response(
1218 None,
1219 decision.selected_model().to_string(),
1220 )),
1221 metadata: None,
1222 })
1223 }
1224 }
1225
1226 let barrier = Arc::new(Barrier::new(N));
1227 let targets = LlmTargetSet::new(vec![LlmTarget {
1228 semantic_name: "m".to_string(),
1229 llm_client: Some(Arc::new(BarrierClient {
1230 barrier: barrier.clone(),
1231 })),
1232 }]);
1233 let algo = orch(targets);
1235
1236 let mut handles = Vec::new();
1237 for _ in 0..N {
1238 let algo = algo.clone();
1239 handles.push(tokio::spawn(async move {
1240 algo.run(Context::default(), request())
1241 .await
1242 .map(|(_, response)| {
1243 response
1244 .llm_response
1245 .as_agg()
1246 .map(completion_text)
1247 .unwrap_or_default()
1248 })
1249 }));
1250 }
1251
1252 for handle in handles {
1253 let completion = tokio::time::timeout(Duration::from_secs(5), handle)
1255 .await
1256 .map_err(|error| LibsyError::external("waiting for test task", error))?
1257 .map_err(|source| LibsyError::AlgorithmTask { source })??;
1258 assert_eq!(completion, "m");
1259 }
1260 Ok(())
1261 }
1262
1263 #[tokio::test]
1264 async fn offload_error_propagates_back_to_the_algorithm() -> Result<()> {
1265 let stream = orch(target_set(&[("offload/model", false)])).run_stream(
1269 Context::default(),
1270 request(),
1271 None,
1272 );
1273 tokio::pin!(stream);
1274
1275 let mut saw_error = false;
1276 while let Some(step) = stream.next().await {
1277 match step {
1278 Ok(Step::CallLlm(call)) => {
1279 call.respond(Err(test_error("upstream model call failed")))?;
1280 }
1281 Ok(Step::Decision(_)) => {}
1282 Ok(Step::ReturnToAgent(..)) => {
1283 return Err(test_error(
1284 "expected the offload error to propagate, got a response",
1285 ));
1286 }
1287 Err(err) => {
1288 assert!(err.to_string().contains("upstream model call failed"));
1290 saw_error = true;
1291 }
1292 }
1293 }
1294
1295 assert!(saw_error, "expected an error step");
1296 Ok(())
1297 }
1298
1299 #[tokio::test]
1300 async fn dropping_the_stream_cancels_the_algorithm_task() -> Result<()> {
1301 use std::sync::atomic::{AtomicBool, Ordering};
1302 use std::time::Duration;
1303 use tokio::sync::mpsc;
1304
1305 struct DropGuard(Arc<AtomicBool>);
1308 impl Drop for DropGuard {
1309 fn drop(&mut self) {
1310 self.0.store(true, Ordering::SeqCst);
1311 }
1312 }
1313
1314 struct StuckAlgo {
1315 started: mpsc::UnboundedSender<()>,
1316 dropped: Arc<AtomicBool>,
1317 }
1318
1319 #[async_trait]
1320 impl Algorithm for StuckAlgo {
1321 fn name(&self) -> &str {
1322 "stuck"
1323 }
1324
1325 async fn create_run_task(
1326 self: Arc<Self>,
1327 _ctx: Context,
1328 _driver: Driver,
1329 _request: Request,
1330 ) -> Result<Response> {
1331 let _guard = DropGuard(self.dropped.clone());
1332 let _ = self.started.send(());
1333 std::future::pending::<()>().await;
1335 unreachable!()
1336 }
1337 }
1338
1339 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1340 let dropped = Arc::new(AtomicBool::new(false));
1341 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1342 started: started_tx,
1343 dropped: dropped.clone(),
1344 });
1345
1346 let stream = algo.run_stream(Context::default(), request(), None);
1347 started_rx
1348 .recv()
1349 .await
1350 .ok_or_else(|| test_error("task never started"))?;
1351 drop(stream);
1352 tokio::time::sleep(Duration::from_millis(100)).await;
1353
1354 assert!(
1355 dropped.load(Ordering::SeqCst),
1356 "algorithm task was NOT cancelled after dropping the stream"
1357 );
1358 Ok(())
1359 }
1360
1361 #[tokio::test]
1362 async fn create_run_task_panic_surfaces_as_a_stream_error() -> Result<()> {
1363 struct Panicky;
1366
1367 #[async_trait]
1368 impl Algorithm for Panicky {
1369 fn name(&self) -> &str {
1370 "panicky"
1371 }
1372
1373 async fn create_run_task(
1374 self: Arc<Self>,
1375 _ctx: Context,
1376 _driver: Driver,
1377 _request: Request,
1378 ) -> Result<Response> {
1379 panic!("boom");
1380 }
1381 }
1382
1383 let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
1384 let stream = algo.run_stream(Context::default(), request(), None);
1385 tokio::pin!(stream);
1386
1387 let mut saw_error = false;
1388 while let Some(step) = stream.next().await {
1389 match step {
1390 Err(err) => {
1391 assert!(matches!(err, LibsyError::AlgorithmTask { .. }));
1392 saw_error = true;
1393 }
1394 Ok(_) => return Err(test_error("expected the panic to surface as an error step")),
1395 }
1396 }
1397
1398 assert!(saw_error, "expected an error step from the panicked task");
1399 Ok(())
1400 }
1401
1402 #[tokio::test]
1403 async fn run_returns_an_error_when_the_algorithm_task_panics() -> Result<()> {
1404 struct Panicky;
1407
1408 #[async_trait]
1409 impl Algorithm for Panicky {
1410 fn name(&self) -> &str {
1411 "panicky"
1412 }
1413
1414 async fn create_run_task(
1415 self: Arc<Self>,
1416 _ctx: Context,
1417 _driver: Driver,
1418 _request: Request,
1419 ) -> Result<Response> {
1420 panic!("boom");
1421 }
1422 }
1423
1424 let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
1425 match algo.run(Context::default(), request()).await {
1426 Ok(_) => Err(test_error(
1427 "expected run to surface the algorithm panic as an error",
1428 )),
1429 Err(err) => {
1430 assert!(matches!(err, LibsyError::AlgorithmTask { .. }));
1431 Ok(())
1432 }
1433 }
1434 }
1435
1436 #[tokio::test]
1437 async fn cancelling_run_cancels_the_algorithm_task() -> Result<()> {
1438 use std::sync::atomic::{AtomicBool, Ordering};
1439 use std::time::Duration;
1440 use tokio::sync::mpsc;
1441
1442 struct DropGuard(Arc<AtomicBool>);
1445 impl Drop for DropGuard {
1446 fn drop(&mut self) {
1447 self.0.store(true, Ordering::SeqCst);
1448 }
1449 }
1450
1451 struct StuckAlgo {
1452 started: mpsc::UnboundedSender<()>,
1453 dropped: Arc<AtomicBool>,
1454 }
1455
1456 #[async_trait]
1457 impl Algorithm for StuckAlgo {
1458 fn name(&self) -> &str {
1459 "stuck"
1460 }
1461
1462 async fn create_run_task(
1463 self: Arc<Self>,
1464 _ctx: Context,
1465 _driver: Driver,
1466 _request: Request,
1467 ) -> Result<Response> {
1468 let _guard = DropGuard(self.dropped.clone());
1469 let _ = self.started.send(());
1470 std::future::pending::<()>().await;
1473 unreachable!()
1474 }
1475 }
1476
1477 let (started_tx, mut started_rx) = mpsc::unbounded_channel();
1478 let dropped = Arc::new(AtomicBool::new(false));
1479 let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
1480 started: started_tx,
1481 dropped: dropped.clone(),
1482 });
1483
1484 let run_task = tokio::spawn(async move { algo.run(Context::default(), request()).await });
1487 started_rx
1488 .recv()
1489 .await
1490 .ok_or_else(|| test_error("task never started"))?;
1491 run_task.abort();
1492 tokio::time::sleep(Duration::from_millis(100)).await;
1493
1494 assert!(
1495 dropped.load(Ordering::SeqCst),
1496 "algorithm task was NOT cancelled after cancelling run"
1497 );
1498 Ok(())
1499 }
1500
1501 struct LoserClient {
1507 started: Arc<tokio::sync::Notify>,
1508 delay: Option<std::time::Duration>,
1509 }
1510
1511 #[async_trait]
1512 impl RoutedLlmClient for LoserClient {
1513 async fn call(
1514 &self,
1515 _ctx: Context,
1516 _request: Request,
1517 decision: Arc<dyn Decision>,
1518 ) -> std::result::Result<Response, LlmClientError> {
1519 self.started.notify_one();
1520 match self.delay {
1521 Some(delay) => tokio::time::sleep(delay).await,
1522 None => std::future::pending::<()>().await,
1523 }
1524 Ok(Response {
1525 llm_response: LlmResponse::Agg(text_response(
1526 None,
1527 decision.selected_model().to_string(),
1528 )),
1529 metadata: None,
1530 })
1531 }
1532 }
1533
1534 struct GatedEchoClient {
1537 gate: Arc<tokio::sync::Notify>,
1538 }
1539
1540 #[async_trait]
1541 impl RoutedLlmClient for GatedEchoClient {
1542 async fn call(
1543 &self,
1544 _ctx: Context,
1545 _request: Request,
1546 decision: Arc<dyn Decision>,
1547 ) -> std::result::Result<Response, LlmClientError> {
1548 self.gate.notified().await;
1549 Ok(Response {
1550 llm_response: LlmResponse::Agg(text_response(
1551 None,
1552 decision.selected_model().to_string(),
1553 )),
1554 metadata: None,
1555 })
1556 }
1557 }
1558
1559 struct Hedge {
1562 winner: LlmTarget,
1563 loser: LlmTarget,
1564 }
1565
1566 #[async_trait]
1567 impl Algorithm for Hedge {
1568 fn name(&self) -> &str {
1569 "hedge"
1570 }
1571
1572 async fn create_run_task(
1573 self: Arc<Self>,
1574 ctx: Context,
1575 driver: Driver,
1576 request: Request,
1577 ) -> Result<Response> {
1578 let dec_w: Arc<dyn Decision> = Arc::new(TestDecision {
1579 model: self.winner.semantic_name.clone(),
1580 });
1581 let dec_l: Arc<dyn Decision> = Arc::new(TestDecision {
1582 model: self.loser.semantic_name.clone(),
1583 });
1584 let win = driver.call_llm_target(ctx.clone(), &self.winner, request.clone(), dec_w);
1585 let lose = driver.call_llm_target(ctx, &self.loser, request, dec_l);
1586 tokio::select! {
1588 res = win => res,
1589 res = lose => res,
1590 }
1591 }
1592 }
1593
1594 fn hedge(loser_delay: Option<std::time::Duration>) -> Arc<dyn Algorithm> {
1597 let started = Arc::new(tokio::sync::Notify::new());
1598 let winner = LlmTarget {
1599 semantic_name: "winner".to_string(),
1600 llm_client: Some(Arc::new(GatedEchoClient {
1601 gate: started.clone(),
1602 })),
1603 };
1604 let loser = LlmTarget {
1605 semantic_name: "loser".to_string(),
1606 llm_client: Some(Arc::new(LoserClient {
1607 started,
1608 delay: loser_delay,
1609 })),
1610 };
1611 Arc::new(Hedge { winner, loser })
1612 }
1613
1614 #[tokio::test]
1615 async fn run_returns_the_winner_without_a_late_loser_overwriting_it() -> Result<()> {
1616 let (_trace, response) = hedge(Some(std::time::Duration::from_millis(50)))
1619 .run(Context::default(), request())
1620 .await?;
1621 assert_eq!(
1622 response
1623 .llm_response
1624 .as_agg()
1625 .map(completion_text)
1626 .unwrap_or_default(),
1627 "winner"
1628 );
1629 Ok(())
1630 }
1631
1632 #[tokio::test]
1633 async fn run_returns_the_winner_without_hanging_on_a_pending_loser() -> Result<()> {
1634 let run = hedge(None).run(Context::default(), request());
1637 let (_trace, response) = tokio::time::timeout(std::time::Duration::from_secs(1), run)
1638 .await
1639 .map_err(|error| LibsyError::external("waiting for pending loser", error))??;
1640 assert_eq!(
1641 response
1642 .llm_response
1643 .as_agg()
1644 .map(completion_text)
1645 .unwrap_or_default(),
1646 "winner"
1647 );
1648 Ok(())
1649 }
1650
1651 #[tokio::test]
1652 async fn run_surfaces_a_terminal_error_with_many_calls_in_flight() -> Result<()> {
1653 use std::sync::atomic::{AtomicUsize, Ordering};
1654
1655 const N: usize = 10;
1658
1659 struct EnterThenPend {
1661 started: Arc<AtomicUsize>,
1662 all_started: Arc<tokio::sync::Notify>,
1663 n: usize,
1664 }
1665
1666 #[async_trait]
1667 impl RoutedLlmClient for EnterThenPend {
1668 async fn call(
1669 &self,
1670 _ctx: Context,
1671 _request: Request,
1672 _decision: Arc<dyn Decision>,
1673 ) -> std::result::Result<Response, LlmClientError> {
1674 if self.started.fetch_add(1, Ordering::SeqCst) + 1 == self.n {
1675 self.all_started.notify_one();
1676 }
1677 std::future::pending::<()>().await;
1678 unreachable!()
1679 }
1680 }
1681
1682 struct FanOutThenError {
1685 target: LlmTarget,
1686 all_started: Arc<tokio::sync::Notify>,
1687 n: usize,
1688 }
1689
1690 #[async_trait]
1691 impl Algorithm for FanOutThenError {
1692 fn name(&self) -> &str {
1693 "fan_out_then_error"
1694 }
1695
1696 async fn create_run_task(
1697 self: Arc<Self>,
1698 ctx: Context,
1699 driver: Driver,
1700 request: Request,
1701 ) -> Result<Response> {
1702 let offloads = futures::future::join_all((0..self.n).map(|i| {
1703 let decision: Arc<dyn Decision> = Arc::new(TestDecision {
1704 model: format!("m{i}"),
1705 });
1706 driver.call_llm_target(ctx.clone(), &self.target, request.clone(), decision)
1707 }));
1708 tokio::select! {
1709 _ = offloads => Err(test_error("offloads unexpectedly completed")),
1710 _ = self.all_started.notified() => {
1711 Err(test_error("terminal error while calls pending"))
1712 }
1713 }
1714 }
1715 }
1716
1717 let all_started = Arc::new(tokio::sync::Notify::new());
1718 let target = LlmTarget {
1719 semantic_name: "pending".to_string(),
1720 llm_client: Some(Arc::new(EnterThenPend {
1721 started: Arc::new(AtomicUsize::new(0)),
1722 all_started: all_started.clone(),
1723 n: N,
1724 })),
1725 };
1726 let algo: Arc<dyn Algorithm> = Arc::new(FanOutThenError {
1727 target,
1728 all_started,
1729 n: N,
1730 });
1731
1732 let run = algo.run(Context::default(), request());
1735 let result = tokio::time::timeout(std::time::Duration::from_millis(500), run)
1736 .await
1737 .map_err(|error| {
1738 LibsyError::external("waiting for terminal error with full call cap", error)
1739 })?;
1740 match result {
1741 Ok(_) => Err(test_error("expected the terminal error, got a response")),
1742 Err(err) => {
1743 assert!(
1744 err.to_string()
1745 .contains("terminal error while calls pending")
1746 );
1747 Ok(())
1748 }
1749 }
1750 }
1751}