use std::{
collections::{HashMap, HashSet},
pin::Pin,
sync::Arc,
time::{Duration, Instant},
};
use async_trait::async_trait;
use futures::{Stream, StreamExt};
use parking_lot::Mutex;
use tracing::Instrument;
use switchyard_protocol::{
Context, Decision, LlmClientError, Request, Response, RoutedLlmClient, RoutingFallbackReason,
Signals, Usage,
};
use super::driver::{DriverRequest, DriverStep, TypeErasedDriver};
use crate::{DriverError, LibsyError, Result, observability};
pub type StepStream = Pin<Box<dyn Stream<Item = Result<Step>> + Send>>;
#[derive(Clone, Debug)]
pub struct LlmCallObservation {
pub selected_model: String,
pub tier: Option<String>,
pub is_routed: bool,
pub is_success: bool,
pub duration: Duration,
pub usage: Option<Usage>,
}
#[derive(Clone, Debug)]
pub enum RunObservation {
LlmCall(LlmCallObservation),
RoutingOverhead(Duration),
}
pub type RunObserver = Arc<dyn Fn(RunObservation) + Send + Sync>;
#[derive(Clone)]
pub struct RoutedRequest {
pub request: Request,
pub decision: Arc<dyn Decision>,
pub default_client: Option<Arc<dyn RoutedLlmClient>>,
pub ctx: Context,
}
pub struct CallLlmRequest {
inner: DriverRequest,
routed: RoutedRequest,
}
impl CallLlmRequest {
fn new(inner: DriverRequest) -> Self {
let routed = match inner.request::<RoutedRequest>() {
Ok(routed) => routed.clone(),
Err(_) => unreachable!("CallLlmRequest payload is always a RoutedRequest"),
};
Self { inner, routed }
}
pub fn get_routed(&self) -> &RoutedRequest {
&self.routed
}
pub fn get_request(&self) -> &Request {
&self.get_routed().request
}
pub fn get_decision(&self) -> &dyn Decision {
self.get_routed().decision.as_ref()
}
pub fn respond(self, result: Result<Response>) -> Result<()> {
self.inner.respond::<Response>(result)
}
}
#[derive(Clone)]
pub struct Driver {
driver: TypeErasedDriver,
routed_call: Arc<Mutex<Option<Duration>>>,
observer: Option<RunObserver>,
}
impl Driver {
pub(crate) fn new() -> Self {
Self::with_observer(None)
}
fn with_observer(observer: Option<RunObserver>) -> Self {
Self {
driver: TypeErasedDriver::new(),
routed_call: Arc::new(Mutex::new(None)),
observer,
}
}
pub(crate) fn routed_call_duration(&self) -> Option<Duration> {
*self.routed_call.lock()
}
pub(crate) fn observe_routing_overhead(&self, duration: Duration) {
if let Some(observer) = &self.observer {
observer(RunObservation::RoutingOverhead(duration));
}
}
#[tracing::instrument(
target = "libsy",
name = "libsy.llm_call",
skip_all,
fields(
algorithm = observability::algorithm_label(&routed.ctx),
selected_model = routed.decision.selected_model(),
openinference.span.kind = "CHAIN",
outcome = tracing::field::Empty,
error = tracing::field::Empty,
input_tokens = tracing::field::Empty,
output_tokens = tracing::field::Empty,
total_tokens = tracing::field::Empty,
reasoning_tokens = tracing::field::Empty,
)
)]
pub async fn call_llm(&self, routed: RoutedRequest) -> Result<Response> {
let algorithm = observability::algorithm_label(&routed.ctx).to_string();
let selected_model = routed.decision.selected_model().to_string();
let tier = routed.decision.routing_tier().map(str::to_string);
let is_routed = routed.decision.is_routed_call();
let started = Instant::now();
let result = self
.driver
.fulfill_request::<RoutedRequest, Response>(routed.ctx.clone(), routed)
.await;
let elapsed = started.elapsed();
observability::record_llm_call(
&algorithm,
&selected_model,
tier.as_deref(),
is_routed,
elapsed,
&result,
&tracing::Span::current(),
);
if let Some(observer) = &self.observer {
observer(RunObservation::LlmCall(LlmCallObservation {
selected_model,
tier,
is_routed,
is_success: result.is_ok(),
duration: elapsed,
usage: result
.as_ref()
.ok()
.and_then(|response| response.llm_response.as_agg())
.map(|response| response.usage.clone()),
}));
}
if is_routed && result.is_ok() {
*self.routed_call.lock() = Some(elapsed);
}
result
}
pub async fn call_llm_target(
&self,
ctx: Context,
target: &LlmTarget,
request: Request,
decision: Arc<dyn Decision>,
) -> Result<Response> {
self.call_llm(RoutedRequest {
request,
decision,
default_client: target.llm_client.clone(),
ctx,
})
.await
}
pub async fn info(&self, ctx: Context, decision: Arc<dyn Decision>) -> Result<()> {
self.driver.info(ctx.clone(), decision.clone()).await?;
observability::record_decision(&ctx, decision.as_ref());
Ok(())
}
pub(crate) async fn finish(&self, ctx: Context, result: Result<Response>) -> Result<()> {
match result {
Ok(response) => self.driver.done(ctx, response).await,
Err(err) => self.driver.fail(ctx, err).await,
}
}
pub(crate) fn stream(&self) -> impl Stream<Item = Result<Step>> + use<> {
self.driver.stream().map(|item| match item? {
DriverStep::Request(req) => Ok(Step::CallLlm(Box::new(CallLlmRequest::new(req)))),
DriverStep::Info(payload) => payload
.downcast::<Arc<dyn Decision>>()
.map(|decision| Step::Decision(*decision))
.map_err(|_| {
DriverError::TypeMismatch {
expected: "Arc<dyn Decision>",
}
.into()
}),
DriverStep::Done(payload) => payload
.downcast::<Response>()
.map(Step::ReturnToAgent)
.map_err(|_| {
DriverError::TypeMismatch {
expected: "Response",
}
.into()
}),
})
}
}
impl Default for Driver {
fn default() -> Self {
Self::new()
}
}
pub enum Step {
CallLlm(Box<CallLlmRequest>),
Decision(Arc<dyn Decision>),
ReturnToAgent(Box<Response>),
}
struct AbortOnDrop(tokio::task::AbortHandle);
impl Drop for AbortOnDrop {
fn drop(&mut self) {
self.0.abort();
}
}
#[derive(Clone)]
pub struct LlmTarget {
pub semantic_name: String,
pub llm_client: Option<Arc<dyn RoutedLlmClient>>,
}
#[derive(Clone)]
pub struct LlmTargetSet {
targets: Vec<LlmTarget>,
}
impl LlmTargetSet {
pub fn new(targets: Vec<LlmTarget>) -> Self {
Self { targets }
}
pub fn targets(&self) -> &[LlmTarget] {
&self.targets
}
pub fn get_target(&self, name: &str) -> Result<LlmTarget> {
self.targets
.iter()
.find(|t| t.semantic_name == name)
.cloned()
.ok_or_else(|| LibsyError::TargetNotFound {
target: name.to_string(),
})
}
pub fn resolve_target(&self, name: &str, ctx: &Context) -> Result<LlmTarget> {
let target = self.get_target(name)?;
if !ctx.is_excluded(&target.semantic_name) {
return Ok(target);
}
self.targets
.iter()
.find(|t| !ctx.is_excluded(&t.semantic_name))
.cloned()
.ok_or(LibsyError::AllTargetsExcluded)
}
}
#[derive(Clone, Hash, PartialEq, Eq)]
pub(crate) enum RoutingIdentity {
Session(String),
Subagent { session: String, agent: String },
}
impl RoutingIdentity {
pub(crate) fn from_request(request: &Request) -> Option<Self> {
let metadata = request.metadata.as_ref()?;
let session = metadata.session_id.as_deref().filter(|id| !id.is_empty())?;
if metadata.is_subagent {
let agent = metadata.agent_id.as_deref().filter(|id| !id.is_empty())?;
Some(Self::Subagent {
session: session.to_string(),
agent: agent.to_string(),
})
} else {
Some(Self::Session(session.to_string()))
}
}
fn session(&self) -> &str {
match self {
Self::Session(session) | Self::Subagent { session, .. } => session,
}
}
}
const MAX_EVICTION_IDENTITIES: usize = 1_024;
#[derive(Default)]
pub(crate) struct SessionEvictions {
by_identity: Mutex<HashMap<RoutingIdentity, HashSet<String>>>,
}
impl SessionEvictions {
pub(crate) fn remove_session(&self, session: &str) {
self.by_identity
.lock()
.retain(|identity, _| identity.session() != session);
}
fn evicted_for(&self, identity: Option<&RoutingIdentity>) -> Vec<String> {
let Some(identity) = identity else {
return Vec::new();
};
self.by_identity
.lock()
.get(identity)
.map(|targets| targets.iter().cloned().collect())
.unwrap_or_default()
}
fn record(&self, identity: Option<&RoutingIdentity>, target: &str) {
let Some(identity) = identity else { return };
let mut histories = self.by_identity.lock();
if histories.len() >= MAX_EVICTION_IDENTITIES
&& !histories.contains_key(identity)
&& let Some(oldest) = histories.keys().next().cloned()
{
histories.remove(&oldest);
}
histories
.entry(identity.clone())
.or_default()
.insert(target.to_string());
}
}
fn eligible_targets(targets: &LlmTargetSet, ctx: &Context) -> usize {
targets
.targets()
.iter()
.filter(|t| !ctx.is_excluded(&t.semantic_name))
.count()
}
pub(crate) fn exclude_evicted(
ctx: &mut Context,
targets: &LlmTargetSet,
evictions: &SessionEvictions,
identity: Option<&RoutingIdentity>,
) {
for target in evictions.evicted_for(identity) {
if eligible_targets(targets, ctx) <= 1 {
break;
}
ctx.exclude_target(target);
}
}
fn classify_fallback(error: &LibsyError) -> Option<(&str, RoutingFallbackReason)> {
let LibsyError::ClientCall { target, source } = error else {
return None;
};
let reason = match source {
LlmClientError::ContextWindowExceeded { .. } => RoutingFallbackReason::ContextWindow,
LlmClientError::Transport { .. } | LlmClientError::Timeout { .. } => {
RoutingFallbackReason::Unavailable
}
LlmClientError::UpstreamHttp { status, .. }
if matches!(*status, 403 | 408 | 429) || (500..=599).contains(status) =>
{
RoutingFallbackReason::Unavailable
}
_ => return None,
};
Some((target, reason))
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn call_llm_with_fallback(
mut ctx: Context,
driver: &Driver,
targets: &LlmTargetSet,
mut target: LlmTarget,
mut decision: Arc<dyn Decision>,
request: Request,
identity: Option<&RoutingIdentity>,
evictions: &SessionEvictions,
target_unavailable: impl Fn(&Request, &str),
fallback_decision: impl Fn(&LlmTarget, &LlmTarget, RoutingFallbackReason) -> Arc<dyn Decision>,
) -> Result<Response> {
loop {
let result = driver
.call_llm_target(ctx.clone(), &target, request.clone(), decision.clone())
.await;
let Err(error) = result else { return result };
let Some((failed, reason)) = classify_fallback(&error) else {
return Err(error);
};
if !ctx.exclude_target(failed) {
return Err(error);
}
match reason {
RoutingFallbackReason::ContextWindow => evictions.record(identity, failed),
RoutingFallbackReason::Unavailable => target_unavailable(&request, failed),
}
let Ok(next) = targets.resolve_target(&target.semantic_name, &ctx) else {
return Err(error);
};
decision = fallback_decision(&target, &next, reason);
target = next;
driver.info(ctx.clone(), decision.clone()).await?;
}
}
#[async_trait]
pub trait Algorithm: Send + Sync + 'static {
fn name(&self) -> &str;
async fn create_run_task(
self: Arc<Self>,
ctx: Context,
driver: Driver,
request: Request,
) -> Result<Response>;
#[allow(unused_variables)]
async fn process_signals(self: Arc<Self>, signals: Signals) -> Result<()> {
Ok(())
}
fn run_stream(
self: Arc<Self>,
ctx: Context,
request: Request,
observer: Option<RunObserver>,
) -> StepStream {
let mut ctx = ctx;
ctx.values.insert(
observability::ALGORITHM_KEY.to_string(),
self.name().to_string(),
);
let driver = Driver::with_observer(observer);
let task_driver = driver.clone();
let task_ctx = ctx.clone();
let stream = task_driver.stream();
let span = observability::run_span(self.name(), &request);
let observed_driver = task_driver.clone();
let handle = tokio::spawn(
async move {
observability::observe_run(
task_ctx.clone(),
observed_driver,
self.create_run_task(task_ctx, task_driver, request),
)
.await
}
.instrument(span),
);
let abort_guard = AbortOnDrop(handle.abort_handle());
let finish_driver = driver.clone();
let finish_ctx = ctx;
let tail: StepStream = Box::pin(
futures::stream::once(async move {
let result = match handle.await {
Ok(response) => response,
Err(source) => Err(LibsyError::AlgorithmTask { source }),
};
finish_driver.finish(finish_ctx, result).await
})
.filter_map(|finish_result| async move { finish_result.err().map(Err) }),
);
let stream: StepStream = Box::pin(stream);
Box::pin(futures::stream::select(stream, tail).map(move |step| {
let _keep_alive = &abort_guard;
step
}))
}
async fn run(
self: Arc<Self>,
ctx: Context,
request: Request,
) -> Result<(Vec<Arc<dyn Decision>>, Response)> {
self.run_observed(ctx, request, None).await
}
async fn run_observed(
self: Arc<Self>,
ctx: Context,
request: Request,
observer: Option<RunObserver>,
) -> Result<(Vec<Arc<dyn Decision>>, Response)> {
#[tracing::instrument(
target = "libsy",
name = "libsy.client_call",
skip_all,
fields(
algorithm = observability::algorithm_label(&call.get_routed().ctx),
switchyard.algorithm = observability::algorithm_label(&call.get_routed().ctx),
switchyard.routing.tier = tracing::field::Empty,
selected_model = call.get_decision().selected_model(),
otel.kind = "client",
otel.name = %format_args!("chat {}", call.get_decision().selected_model()),
openinference.span.kind = "LLM",
gen_ai.operation.name = "chat",
gen_ai.request.model = call.get_decision().selected_model(),
gen_ai.request.stream = tracing::field::Empty,
gen_ai.request.temperature = tracing::field::Empty,
gen_ai.request.top_p = tracing::field::Empty,
gen_ai.request.top_k = tracing::field::Empty,
gen_ai.request.max_tokens = tracing::field::Empty,
gen_ai.request.reasoning.level = tracing::field::Empty,
gen_ai.output.type = tracing::field::Empty,
gen_ai.conversation.id = tracing::field::Empty,
server.address = tracing::field::Empty,
server.port = tracing::field::Empty,
gen_ai.response.id = tracing::field::Empty,
gen_ai.response.model = tracing::field::Empty,
gen_ai.usage.input_tokens = tracing::field::Empty,
gen_ai.usage.output_tokens = tracing::field::Empty,
gen_ai.usage.cache_read.input_tokens = tracing::field::Empty,
gen_ai.usage.cache_creation.input_tokens = tracing::field::Empty,
gen_ai.usage.reasoning.output_tokens = tracing::field::Empty,
outcome = tracing::field::Empty,
otel.status_code = tracing::field::Empty,
error.type = tracing::field::Empty,
error = tracing::field::Empty,
)
)]
async fn serve(call: CallLlmRequest) -> Result<()> {
let span = tracing::Span::current();
observability::record_gen_ai_request(&span, &call.get_routed().request.llm_request);
if let Some(tier) = call.get_decision().routing_tier() {
span.record("switchyard.routing.tier", tier);
}
if let Some(session_id) = call
.get_routed()
.request
.metadata
.as_ref()
.and_then(|metadata| metadata.session_id.as_deref())
{
span.record("gen_ai.conversation.id", session_id);
}
let routed = call.get_routed().clone();
let target = routed.decision.selected_model().to_string();
let client =
routed
.default_client
.clone()
.ok_or_else(|| LibsyError::MissingClient {
target: target.clone(),
})?;
let result = client
.call(routed.ctx, routed.request, routed.decision)
.await
.map_err(|source| LibsyError::client_call(target, source));
let result = observability::observe_client_call(result);
call.respond(result)
}
let stream = self.run_stream(ctx, request, observer);
tokio::pin!(stream);
let mut trace: Vec<Arc<dyn Decision>> = Vec::new();
let mut in_flight = futures::stream::FuturesUnordered::new();
let mut final_response: Option<Response> = None;
loop {
tokio::select! {
Some(result) = in_flight.next() => match result {
Ok(()) => {}, Err(err) => return Err(err), },
step = stream.next() => {
match step {
None => break, Some(item) => match item? {
Step::CallLlm(call) => in_flight.push(serve(*call)),
Step::Decision(decision) => trace.push(decision),
Step::ReturnToAgent(response) => {
final_response = Some(*response);
break;
}
}
}
},
}
}
final_response
.map(|response| (trace, response))
.ok_or(LibsyError::MissingFinalResponse)
}
}
#[cfg(test)]
mod tests {
use super::*;
use futures::StreamExt;
use switchyard_protocol::{
LlmResponse, LlmResponseChunk, completion_text, text_request, text_response,
};
#[derive(Debug, thiserror::Error)]
#[error("{0}")]
struct TestError(&'static str);
fn test_error(message: &'static str) -> LibsyError {
LibsyError::external("test", TestError(message))
}
fn classified_client_error(source: LlmClientError) -> Option<RoutingFallbackReason> {
classify_fallback(&LibsyError::client_call("target", source)).map(|(_, reason)| reason)
}
#[test]
fn route_fallback_only_accepts_context_and_unavailable_failures() {
assert_eq!(
classified_client_error(LlmClientError::ContextWindowExceeded {
model: "target".to_string(),
message: "too long".to_string(),
}),
Some(RoutingFallbackReason::ContextWindow)
);
for source in [
LlmClientError::Transport {
source: Box::new(std::io::Error::other("connection failed")),
},
LlmClientError::Timeout {
source: Box::new(std::io::Error::other("request timed out")),
},
] {
assert_eq!(
classified_client_error(source),
Some(RoutingFallbackReason::Unavailable)
);
}
for (status, expected) in [
(400, None),
(401, None),
(403, Some(RoutingFallbackReason::Unavailable)),
(404, None),
(408, Some(RoutingFallbackReason::Unavailable)),
(409, None),
(429, Some(RoutingFallbackReason::Unavailable)),
(499, None),
(500, Some(RoutingFallbackReason::Unavailable)),
(599, Some(RoutingFallbackReason::Unavailable)),
(600, None),
] {
assert_eq!(
classified_client_error(LlmClientError::UpstreamHttp {
status,
body: "failed".to_string(),
}),
expected
);
}
assert_eq!(
classified_client_error(LlmClientError::InvalidResponse {
source: Box::new(std::io::Error::other("invalid response")),
}),
None
);
}
struct EchoClient;
#[async_trait]
impl RoutedLlmClient for EchoClient {
async fn call(
&self,
_ctx: Context,
_request: Request,
decision: Arc<dyn Decision>,
) -> std::result::Result<Response, LlmClientError> {
Ok(Response {
llm_response: LlmResponse::Agg(text_response(
None,
decision.selected_model().to_string(),
)),
metadata: None,
})
}
}
struct TestDecision {
model: String,
}
impl Decision for TestDecision {
fn selected_model(&self) -> &str {
&self.model
}
fn reasoning(&self) -> Option<&str> {
None
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
}
struct TestAlgo {
target_set: LlmTargetSet,
}
#[async_trait]
impl Algorithm for TestAlgo {
fn name(&self) -> &str {
"test"
}
async fn create_run_task(
self: Arc<Self>,
ctx: Context,
driver: Driver,
request: Request,
) -> Result<Response> {
let target = self
.target_set
.targets()
.first()
.ok_or(LibsyError::NoTargets)?
.clone();
let decision: Arc<dyn Decision> = Arc::new(TestDecision {
model: target.semantic_name.clone(),
});
driver.info(ctx.clone(), decision.clone()).await?;
driver
.call_llm_target(ctx, &target, request, decision)
.await
}
}
fn orch(target_set: LlmTargetSet) -> Arc<dyn Algorithm> {
Arc::new(TestAlgo { target_set })
}
fn request() -> Request {
Request {
llm_request: text_request(Some("auto".to_string()), "hi".to_string()),
raw_request: None,
metadata: None,
}
}
fn target_set(names: &[(&str, bool)]) -> LlmTargetSet {
let targets = names
.iter()
.map(|(name, has_client)| LlmTarget {
semantic_name: name.to_string(),
llm_client: has_client.then(|| Arc::new(EchoClient) as Arc<dyn RoutedLlmClient>),
})
.collect();
LlmTargetSet::new(targets)
}
#[tokio::test]
async fn observed_run_reports_one_successful_routed_call() -> Result<()> {
let observations = Arc::new(Mutex::new(Vec::new()));
let observed = observations.clone();
let observer: RunObserver = Arc::new(move |observation| observed.lock().push(observation));
let (_, response) = orch(target_set(&[("direct/model", true)]))
.run_observed(Context::default(), request(), Some(observer))
.await?;
assert_eq!(
response.llm_response.as_agg().map(completion_text),
Some("direct/model".to_string())
);
let observations = observations.lock();
assert_eq!(observations.len(), 2);
let RunObservation::LlmCall(observation) = &observations[0] else {
return Err(test_error("expected an LLM call observation"));
};
assert_eq!(observation.selected_model, "direct/model");
assert!(observation.is_routed);
assert!(observation.is_success);
assert!(observation.usage.is_some());
assert!(matches!(
observations[1],
RunObservation::RoutingOverhead(_)
));
Ok(())
}
#[test]
fn target_lookup_returns_the_missing_target() {
let error = target_set(&[]).get_target("missing").err();
assert!(matches!(
error,
Some(LibsyError::TargetNotFound { target }) if target == "missing"
));
}
struct StreamingClient {
chunks: Vec<LlmResponseChunk>,
}
#[async_trait]
impl RoutedLlmClient for StreamingClient {
async fn call(
&self,
_ctx: Context,
_request: Request,
_decision: Arc<dyn Decision>,
) -> std::result::Result<Response, LlmClientError> {
let stream = futures::stream::iter(
self.chunks
.clone()
.into_iter()
.map(|chunk| Ok(chunk.into())),
)
.boxed();
Ok(Response {
llm_response: LlmResponse::Stream(stream),
metadata: None,
})
}
}
fn streaming_orch(chunks: Vec<LlmResponseChunk>) -> Arc<dyn Algorithm> {
let target = LlmTarget {
semantic_name: "stream/model".to_string(),
llm_client: Some(Arc::new(StreamingClient { chunks }) as Arc<dyn RoutedLlmClient>),
};
orch(LlmTargetSet::new(vec![target]))
}
#[tokio::test]
async fn run_returns_a_streamed_response_the_caller_aggregates() -> Result<()> {
let orch = streaming_orch(vec![
LlmResponseChunk::MessageStart {
id: Some("m1".to_string()),
model: Some("stream/model".to_string()),
},
LlmResponseChunk::TextDelta {
index: 0,
text: "hel".to_string(),
},
LlmResponseChunk::TextDelta {
index: 0,
text: "lo".to_string(),
},
LlmResponseChunk::MessageStop {
reason: Some("stop".to_string()),
},
]);
let (trace, response) = orch.run(Context::default(), request()).await?;
let agg = response
.llm_response
.into_agg()
.await
.map_err(|error| LibsyError::external("aggregating response stream", error))?;
assert_eq!(completion_text(&agg), "hello");
assert_eq!(agg.model.as_deref(), Some("stream/model"));
assert_eq!(trace.len(), 1);
Ok(())
}
#[tokio::test]
async fn aggregating_a_streamed_response_propagates_a_mid_stream_error() -> Result<()> {
let orch = streaming_orch(vec![
LlmResponseChunk::TextDelta {
index: 0,
text: "partial".to_string(),
},
LlmResponseChunk::StreamError {
message: "upstream exploded".to_string(),
},
]);
let (_, response) = orch.run(Context::default(), request()).await?;
match response.llm_response.into_agg().await {
Ok(_) => panic!("expected a mid-stream error, got an aggregate"),
Err(err) => {
assert!(err.to_string().contains("upstream exploded"));
Ok(())
}
}
}
#[tokio::test]
async fn run_offloads_via_promise_then_returns_to_agent() -> Result<()> {
let stream = orch(target_set(&[("offload/model", false)])).run_stream(
Context::default(),
request(),
None,
);
tokio::pin!(stream);
let mut saw_call = false;
let mut final_completion = None;
while let Some(step) = stream.next().await {
match step? {
Step::CallLlm(call) => {
saw_call = true;
assert_eq!(call.get_decision().selected_model(), "offload/model");
call.respond(Ok(Response {
llm_response: LlmResponse::Agg(text_response(
None,
"fulfilled".to_string(),
)),
metadata: None,
}))?;
}
Step::Decision(decision) => {
assert_eq!(decision.selected_model(), "offload/model");
}
Step::ReturnToAgent(response) => {
final_completion = Some(
response
.llm_response
.as_agg()
.map(completion_text)
.unwrap_or_default(),
);
}
}
}
assert!(saw_call, "expected a CallLlm step before ReturnToAgent");
assert_eq!(
final_completion.ok_or_else(|| test_error("no ReturnToAgent step"))?,
"fulfilled"
);
Ok(())
}
#[tokio::test]
async fn client_backed_target_offloads_with_a_default_client() -> Result<()> {
let stream = orch(target_set(&[("direct/model", true)])).run_stream(
Context::default(),
request(),
None,
);
tokio::pin!(stream);
let mut final_completion = None;
while let Some(step) = stream.next().await {
match step? {
Step::CallLlm(call) => {
let routed = call.get_routed().clone();
let client = routed
.default_client
.clone()
.ok_or_else(|| test_error("expected a default client"))?;
let target = routed.decision.selected_model().to_string();
let result = client
.call(routed.ctx, routed.request, routed.decision)
.await
.map_err(|error| LibsyError::client_call(target, error));
call.respond(result)?;
}
Step::Decision(_) => {}
Step::ReturnToAgent(response) => {
final_completion = Some(
response
.llm_response
.as_agg()
.map(completion_text)
.unwrap_or_default(),
);
}
}
}
assert_eq!(
final_completion.ok_or_else(|| test_error("no ReturnToAgent"))?,
"direct/model"
);
Ok(())
}
#[tokio::test]
async fn run_returns_the_response_when_all_targets_have_clients() -> Result<()> {
let (trace, response) = orch(target_set(&[("direct/model", true)]))
.run(Context::default(), request())
.await?;
assert_eq!(
response
.llm_response
.as_agg()
.map(completion_text)
.unwrap_or_default(),
"direct/model"
);
assert_eq!(trace[0].selected_model(), "direct/model");
Ok(())
}
#[tokio::test]
async fn run_errors_when_a_target_lacks_a_client() -> Result<()> {
let error = orch(target_set(&[("offload/model", false)]))
.run(Context::default(), request())
.await
.err()
.ok_or_else(|| test_error("expected a missing-client error"))?;
assert!(matches!(
error,
LibsyError::MissingClient { target } if target == "offload/model"
));
Ok(())
}
#[tokio::test(flavor = "multi_thread", worker_threads = 12)]
async fn requests_are_processed_in_parallel() -> Result<()> {
use std::time::Duration;
use tokio::sync::Barrier;
const N: usize = 12;
struct BarrierClient {
barrier: Arc<Barrier>,
}
#[async_trait]
impl RoutedLlmClient for BarrierClient {
async fn call(
&self,
_ctx: Context,
_request: Request,
decision: Arc<dyn Decision>,
) -> std::result::Result<Response, LlmClientError> {
self.barrier.wait().await;
Ok(Response {
llm_response: LlmResponse::Agg(text_response(
None,
decision.selected_model().to_string(),
)),
metadata: None,
})
}
}
let barrier = Arc::new(Barrier::new(N));
let targets = LlmTargetSet::new(vec![LlmTarget {
semantic_name: "m".to_string(),
llm_client: Some(Arc::new(BarrierClient {
barrier: barrier.clone(),
})),
}]);
let algo = orch(targets);
let mut handles = Vec::new();
for _ in 0..N {
let algo = algo.clone();
handles.push(tokio::spawn(async move {
algo.run(Context::default(), request())
.await
.map(|(_, response)| {
response
.llm_response
.as_agg()
.map(completion_text)
.unwrap_or_default()
})
}));
}
for handle in handles {
let completion = tokio::time::timeout(Duration::from_secs(5), handle)
.await
.map_err(|error| LibsyError::external("waiting for test task", error))?
.map_err(|source| LibsyError::AlgorithmTask { source })??;
assert_eq!(completion, "m");
}
Ok(())
}
#[tokio::test]
async fn offload_error_propagates_back_to_the_algorithm() -> Result<()> {
let stream = orch(target_set(&[("offload/model", false)])).run_stream(
Context::default(),
request(),
None,
);
tokio::pin!(stream);
let mut saw_error = false;
while let Some(step) = stream.next().await {
match step {
Ok(Step::CallLlm(call)) => {
call.respond(Err(test_error("upstream model call failed")))?;
}
Ok(Step::Decision(_)) => {}
Ok(Step::ReturnToAgent(..)) => {
return Err(test_error(
"expected the offload error to propagate, got a response",
));
}
Err(err) => {
assert!(err.to_string().contains("upstream model call failed"));
saw_error = true;
}
}
}
assert!(saw_error, "expected an error step");
Ok(())
}
#[tokio::test]
async fn dropping_the_stream_cancels_the_algorithm_task() -> Result<()> {
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use tokio::sync::mpsc;
struct DropGuard(Arc<AtomicBool>);
impl Drop for DropGuard {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
struct StuckAlgo {
started: mpsc::UnboundedSender<()>,
dropped: Arc<AtomicBool>,
}
#[async_trait]
impl Algorithm for StuckAlgo {
fn name(&self) -> &str {
"stuck"
}
async fn create_run_task(
self: Arc<Self>,
_ctx: Context,
_driver: Driver,
_request: Request,
) -> Result<Response> {
let _guard = DropGuard(self.dropped.clone());
let _ = self.started.send(());
std::future::pending::<()>().await;
unreachable!()
}
}
let (started_tx, mut started_rx) = mpsc::unbounded_channel();
let dropped = Arc::new(AtomicBool::new(false));
let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
started: started_tx,
dropped: dropped.clone(),
});
let stream = algo.run_stream(Context::default(), request(), None);
started_rx
.recv()
.await
.ok_or_else(|| test_error("task never started"))?;
drop(stream);
tokio::time::sleep(Duration::from_millis(100)).await;
assert!(
dropped.load(Ordering::SeqCst),
"algorithm task was NOT cancelled after dropping the stream"
);
Ok(())
}
#[tokio::test]
async fn create_run_task_panic_surfaces_as_a_stream_error() -> Result<()> {
struct Panicky;
#[async_trait]
impl Algorithm for Panicky {
fn name(&self) -> &str {
"panicky"
}
async fn create_run_task(
self: Arc<Self>,
_ctx: Context,
_driver: Driver,
_request: Request,
) -> Result<Response> {
panic!("boom");
}
}
let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
let stream = algo.run_stream(Context::default(), request(), None);
tokio::pin!(stream);
let mut saw_error = false;
while let Some(step) = stream.next().await {
match step {
Err(err) => {
assert!(matches!(err, LibsyError::AlgorithmTask { .. }));
saw_error = true;
}
Ok(_) => return Err(test_error("expected the panic to surface as an error step")),
}
}
assert!(saw_error, "expected an error step from the panicked task");
Ok(())
}
#[tokio::test]
async fn run_returns_an_error_when_the_algorithm_task_panics() -> Result<()> {
struct Panicky;
#[async_trait]
impl Algorithm for Panicky {
fn name(&self) -> &str {
"panicky"
}
async fn create_run_task(
self: Arc<Self>,
_ctx: Context,
_driver: Driver,
_request: Request,
) -> Result<Response> {
panic!("boom");
}
}
let algo: Arc<dyn Algorithm> = Arc::new(Panicky);
match algo.run(Context::default(), request()).await {
Ok(_) => Err(test_error(
"expected run to surface the algorithm panic as an error",
)),
Err(err) => {
assert!(matches!(err, LibsyError::AlgorithmTask { .. }));
Ok(())
}
}
}
#[tokio::test]
async fn cancelling_run_cancels_the_algorithm_task() -> Result<()> {
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use tokio::sync::mpsc;
struct DropGuard(Arc<AtomicBool>);
impl Drop for DropGuard {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
struct StuckAlgo {
started: mpsc::UnboundedSender<()>,
dropped: Arc<AtomicBool>,
}
#[async_trait]
impl Algorithm for StuckAlgo {
fn name(&self) -> &str {
"stuck"
}
async fn create_run_task(
self: Arc<Self>,
_ctx: Context,
_driver: Driver,
_request: Request,
) -> Result<Response> {
let _guard = DropGuard(self.dropped.clone());
let _ = self.started.send(());
std::future::pending::<()>().await;
unreachable!()
}
}
let (started_tx, mut started_rx) = mpsc::unbounded_channel();
let dropped = Arc::new(AtomicBool::new(false));
let algo: Arc<dyn Algorithm> = Arc::new(StuckAlgo {
started: started_tx,
dropped: dropped.clone(),
});
let run_task = tokio::spawn(async move { algo.run(Context::default(), request()).await });
started_rx
.recv()
.await
.ok_or_else(|| test_error("task never started"))?;
run_task.abort();
tokio::time::sleep(Duration::from_millis(100)).await;
assert!(
dropped.load(Ordering::SeqCst),
"algorithm task was NOT cancelled after cancelling run"
);
Ok(())
}
struct LoserClient {
started: Arc<tokio::sync::Notify>,
delay: Option<std::time::Duration>,
}
#[async_trait]
impl RoutedLlmClient for LoserClient {
async fn call(
&self,
_ctx: Context,
_request: Request,
decision: Arc<dyn Decision>,
) -> std::result::Result<Response, LlmClientError> {
self.started.notify_one();
match self.delay {
Some(delay) => tokio::time::sleep(delay).await,
None => std::future::pending::<()>().await,
}
Ok(Response {
llm_response: LlmResponse::Agg(text_response(
None,
decision.selected_model().to_string(),
)),
metadata: None,
})
}
}
struct GatedEchoClient {
gate: Arc<tokio::sync::Notify>,
}
#[async_trait]
impl RoutedLlmClient for GatedEchoClient {
async fn call(
&self,
_ctx: Context,
_request: Request,
decision: Arc<dyn Decision>,
) -> std::result::Result<Response, LlmClientError> {
self.gate.notified().await;
Ok(Response {
llm_response: LlmResponse::Agg(text_response(
None,
decision.selected_model().to_string(),
)),
metadata: None,
})
}
}
struct Hedge {
winner: LlmTarget,
loser: LlmTarget,
}
#[async_trait]
impl Algorithm for Hedge {
fn name(&self) -> &str {
"hedge"
}
async fn create_run_task(
self: Arc<Self>,
ctx: Context,
driver: Driver,
request: Request,
) -> Result<Response> {
let dec_w: Arc<dyn Decision> = Arc::new(TestDecision {
model: self.winner.semantic_name.clone(),
});
let dec_l: Arc<dyn Decision> = Arc::new(TestDecision {
model: self.loser.semantic_name.clone(),
});
let win = driver.call_llm_target(ctx.clone(), &self.winner, request.clone(), dec_w);
let lose = driver.call_llm_target(ctx, &self.loser, request, dec_l);
tokio::select! {
res = win => res,
res = lose => res,
}
}
}
fn hedge(loser_delay: Option<std::time::Duration>) -> Arc<dyn Algorithm> {
let started = Arc::new(tokio::sync::Notify::new());
let winner = LlmTarget {
semantic_name: "winner".to_string(),
llm_client: Some(Arc::new(GatedEchoClient {
gate: started.clone(),
})),
};
let loser = LlmTarget {
semantic_name: "loser".to_string(),
llm_client: Some(Arc::new(LoserClient {
started,
delay: loser_delay,
})),
};
Arc::new(Hedge { winner, loser })
}
#[tokio::test]
async fn run_returns_the_winner_without_a_late_loser_overwriting_it() -> Result<()> {
let (_trace, response) = hedge(Some(std::time::Duration::from_millis(50)))
.run(Context::default(), request())
.await?;
assert_eq!(
response
.llm_response
.as_agg()
.map(completion_text)
.unwrap_or_default(),
"winner"
);
Ok(())
}
#[tokio::test]
async fn run_returns_the_winner_without_hanging_on_a_pending_loser() -> Result<()> {
let run = hedge(None).run(Context::default(), request());
let (_trace, response) = tokio::time::timeout(std::time::Duration::from_secs(1), run)
.await
.map_err(|error| LibsyError::external("waiting for pending loser", error))??;
assert_eq!(
response
.llm_response
.as_agg()
.map(completion_text)
.unwrap_or_default(),
"winner"
);
Ok(())
}
#[tokio::test]
async fn run_surfaces_a_terminal_error_with_many_calls_in_flight() -> Result<()> {
use std::sync::atomic::{AtomicUsize, Ordering};
const N: usize = 10;
struct EnterThenPend {
started: Arc<AtomicUsize>,
all_started: Arc<tokio::sync::Notify>,
n: usize,
}
#[async_trait]
impl RoutedLlmClient for EnterThenPend {
async fn call(
&self,
_ctx: Context,
_request: Request,
_decision: Arc<dyn Decision>,
) -> std::result::Result<Response, LlmClientError> {
if self.started.fetch_add(1, Ordering::SeqCst) + 1 == self.n {
self.all_started.notify_one();
}
std::future::pending::<()>().await;
unreachable!()
}
}
struct FanOutThenError {
target: LlmTarget,
all_started: Arc<tokio::sync::Notify>,
n: usize,
}
#[async_trait]
impl Algorithm for FanOutThenError {
fn name(&self) -> &str {
"fan_out_then_error"
}
async fn create_run_task(
self: Arc<Self>,
ctx: Context,
driver: Driver,
request: Request,
) -> Result<Response> {
let offloads = futures::future::join_all((0..self.n).map(|i| {
let decision: Arc<dyn Decision> = Arc::new(TestDecision {
model: format!("m{i}"),
});
driver.call_llm_target(ctx.clone(), &self.target, request.clone(), decision)
}));
tokio::select! {
_ = offloads => Err(test_error("offloads unexpectedly completed")),
_ = self.all_started.notified() => {
Err(test_error("terminal error while calls pending"))
}
}
}
}
let all_started = Arc::new(tokio::sync::Notify::new());
let target = LlmTarget {
semantic_name: "pending".to_string(),
llm_client: Some(Arc::new(EnterThenPend {
started: Arc::new(AtomicUsize::new(0)),
all_started: all_started.clone(),
n: N,
})),
};
let algo: Arc<dyn Algorithm> = Arc::new(FanOutThenError {
target,
all_started,
n: N,
});
let run = algo.run(Context::default(), request());
let result = tokio::time::timeout(std::time::Duration::from_millis(500), run)
.await
.map_err(|error| {
LibsyError::external("waiting for terminal error with full call cap", error)
})?;
match result {
Ok(_) => Err(test_error("expected the terminal error, got a response")),
Err(err) => {
assert!(
err.to_string()
.contains("terminal error while calls pending")
);
Ok(())
}
}
}
}