Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
60 changes: 48 additions & 12 deletions crates/switchyard-runner/src/algorithm.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,8 +15,8 @@ use libsy::{
CustomClassifierConfig, CustomClassifierPolicy, DecisionJudgeConfig, EscalationJudgeConfig,
GateTrigger, HandoffNoteConfig, LlmCapabilityConfig, LlmClassifierConfig, LlmFallback,
LlmTaskClassifier, Noop, Passthrough, PickerMode, PlanExecute, PlanExecuteConfig, Random,
StageRouter, StageRouterConfig, SubagentRouter, SubagentRouterConfig, TaskClassifierConfig,
ToolSemantics,
RuntimeModels, StageRouter, StageRouterConfig, SubagentRouter, SubagentRouterConfig,
TaskClassifierConfig, ToolSemantics,
};
use serde::Deserialize;
use serde_json::Value;
Expand Down Expand Up @@ -701,10 +701,7 @@ impl AlgorithmSpec {
}

/// Target names grouped as the runtime [`Driver`](libsy::Driver) expects them.
pub(crate) fn runtime_model_names(
&self,
route_name: &str,
) -> AlgorithmResult<RuntimeModelNames> {
fn runtime_model_names(&self, route_name: &str) -> AlgorithmResult<RuntimeModelNames> {
let parent = match self {
Self::Noop { .. } => HashMap::new(),
Self::Random { targets, .. } | Self::PrefillRouter { targets, .. } => {
Expand Down Expand Up @@ -820,22 +817,44 @@ impl AlgorithmSpec {
}
}

/// Builds this algorithm after resolving configured target names.
/// Builds only the routing algorithm from this specification.
///
/// Embedding hosts should normally use [`Self::build_with_runtime_models`], which also
/// resolves the model groups supplied to the algorithm at execution time.
pub fn build(
&self,
context: &str,
route_name: &str,
targets: &BTreeMap<String, ModelId>,
) -> AlgorithmResult<Arc<dyn Algorithm>> {
build_algorithm(context, self, targets)
build_algorithm(route_name, self, targets)
}

/// Builds the routing algorithm and its matching runtime model groups.
///
/// `route_name` identifies the route in validation errors. `targets` maps configured target
/// names to the model IDs served by the embedding host.
pub fn build_with_runtime_models(
&self,
route_name: &str,
targets: &BTreeMap<String, ModelId>,
) -> AlgorithmResult<(Arc<dyn Algorithm>, RuntimeModels)> {
let algorithm = self.build(route_name, targets)?;
let names = self.runtime_model_names(route_name)?;
let mut models =
RuntimeModels::new(resolve_runtime_models(route_name, names.parent, targets)?);
if let Some(subagent) = names.subagent {
models = models.with_subagent(resolve_runtime_models(route_name, subagent, targets)?);
}
Ok((algorithm, models))
}
}

/// One route's target names, grouped by category and by routing scope.
pub(crate) struct RuntimeModelNames {
struct RuntimeModelNames {
/// Groups the algorithm itself routes over.
pub(crate) parent: HashMap<Category, Vec<String>>,
parent: HashMap<Category, Vec<String>>,
/// Groups delegated sub-agent work routes over, when the route has a `subagents` table.
pub(crate) subagent: Option<HashMap<Category, Vec<String>>>,
subagent: Option<HashMap<Category, Vec<String>>>,
}

fn category_models(
Expand Down Expand Up @@ -1595,3 +1614,20 @@ fn resolve_target_model_id(
))
})
}

fn resolve_runtime_models(
route_name: &str,
names: HashMap<Category, Vec<String>>,
targets: &BTreeMap<String, ModelId>,
) -> AlgorithmResult<HashMap<Category, Vec<ModelId>>> {
names
.into_iter()
.map(|(category, names)| {
let models = names
.into_iter()
.map(|name| resolve_target_model_id(route_name, &name, targets))
.collect::<AlgorithmResult<Vec<_>>>()?;
Ok((category, models))
})
.collect()
}
39 changes: 4 additions & 35 deletions crates/switchyard-runner/src/config.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,15 +9,14 @@ use std::path::Path;
use std::sync::Arc;
use std::time::Duration;

use libsy::RuntimeModels;
use serde::de::DeserializeOwned;
use serde::{Deserialize, Deserializer};
use serde_json::Value;
use switchyard_llm_client::{
AuxiliaryOperation, Backend, ClientRouter, DEFAULT_MAX_RETRIES, HttpBackendConfig, ModelConfig,
SystemOneClient, TranslatingLlmClient,
};
use switchyard_protocol::{Category, ModelId, RoutedDecisionClient, RoutedLlmClient, WireFormat};
use switchyard_protocol::{ModelId, RoutedDecisionClient, RoutedLlmClient, WireFormat};

use crate::{
AlgorithmSpec, AuxiliaryTarget, CallerAuthKind, DecisionTarget, ModelCapabilities, Route,
Expand Down Expand Up @@ -265,9 +264,9 @@ impl DeploymentConfig {
"route {route_name} context_window must be greater than zero"
)));
}
let algorithm = config
let (algorithm, models) = config
.algorithm
.build(route_name, &targets)
.build_with_runtime_models(route_name, &targets)
.map_err(|error| RunnerError::configuration_source(error.to_string(), error))?;
let (route_clients, caller_auth) =
self.build_route_clients(route_name, config, &clients, &decision_clients)?;
Expand All @@ -280,14 +279,6 @@ impl DeploymentConfig {
.into_iter()
.filter_map(|name| self.decision_target(name))
.collect();
let names = config
.algorithm
.runtime_model_names(route_name)
.map_err(|error| RunnerError::configuration_source(error.to_string(), error))?;
let mut models = RuntimeModels::new(resolve_category_models(names.parent, &targets)?);
if let Some(subagent) = names.subagent {
models = models.with_subagent(resolve_category_models(subagent, &targets)?);
}
let route = Route::new(
algorithm,
route_clients,
Expand Down Expand Up @@ -751,29 +742,6 @@ impl ClientFormat {
}
}

/// Resolves one scope's configured target names to the models the driver serves.
fn resolve_category_models(
names: HashMap<Category, Vec<String>>,
targets: &BTreeMap<String, ModelId>,
) -> RunnerResult<HashMap<Category, Vec<ModelId>>> {
names
.into_iter()
.map(|(category, names)| {
let models = names
.into_iter()
.map(|name| {
targets.get(&name).cloned().ok_or_else(|| {
RunnerError::configuration(format!(
"route references unknown target {name}"
))
})
})
.collect::<RunnerResult<Vec<_>>>()?;
Ok((category, models))
})
.collect()
}

fn build_backend(
client_name: &str,
config: &LlmClientConfig,
Expand Down Expand Up @@ -917,6 +885,7 @@ bogus = true
mod deployment_tests {
use super::*;
use serde_json::json;
use switchyard_protocol::Category;

const VALID_CONFIG: &str = r#"
schema_version = 1
Expand Down
29 changes: 23 additions & 6 deletions crates/switchyard-runner/tests/route.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,11 +7,10 @@ use std::sync::{Arc, Mutex};

use async_trait::async_trait;
use futures_util::StreamExt;
use libsy::RuntimeModels;
use switchyard_llm_client::{ClientRouter, RunObservation};
use switchyard_protocol::{
Category, LlmClientError, LlmResponse, ModelId, Request, Response, RoutedLlmClient,
text_request, text_response,
LlmClientError, LlmResponse, ModelId, Request, Response, RoutedLlmClient, text_request,
text_response,
};
use switchyard_runner::{AlgorithmSpec, ModelCapabilities, Route};

Expand Down Expand Up @@ -40,8 +39,8 @@ fn plugin_route(client: Arc<dyn RoutedLlmClient>) -> Route {
"semantic-target".to_string(),
ModelId::from("semantic-target"),
)]);
let algorithm = spec
.build("switchyard", &targets)
let (algorithm, models) = spec
.build_with_runtime_models("switchyard", &targets)
.expect("identity target map should build");
let clients = ClientRouter::new(
BTreeMap::from([(ModelId::from("semantic-target"), client)])
Expand All @@ -56,10 +55,28 @@ fn plugin_route(client: Arc<dyn RoutedLlmClient>) -> Route {
None,
None,
Vec::new(),
RuntimeModels::new([(Category::Any, vec![ModelId::from("semantic-target")])].into()),
models,
)
}

#[test]
fn runtime_model_builder_rejects_an_unknown_target() {
// Embedding hosts do not get DeploymentConfig's target prevalidation.
let spec = AlgorithmSpec::Passthrough {
target: "missing".to_string(),
subagents: None,
};
let error = match spec.build_with_runtime_models("embedded", &BTreeMap::new()) {
Ok(_) => panic!("unknown target should fail"),
Err(error) => error,
};

assert_eq!(
error.to_string(),
"route embedded references unknown target missing"
);
}

#[tokio::test]
async fn plugin_shaped_route_executes_without_runner_model_or_toml() {
let route = plugin_route(Arc::new(StubClient));
Expand Down
Loading