diff --git a/Cargo.lock b/Cargo.lock index aeeeaf37b..48b15ca63 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2358,6 +2358,7 @@ dependencies = [ "opentelemetry_sdk", "parking_lot", "reqwest", + "serde", "serde_json", "switchyard-libsy", "switchyard-protocol", diff --git a/crates/libsy-llm-client/Cargo.toml b/crates/libsy-llm-client/Cargo.toml index cb8ea4f93..ded0ecab8 100644 --- a/crates/libsy-llm-client/Cargo.toml +++ b/crates/libsy-llm-client/Cargo.toml @@ -29,6 +29,7 @@ parking_lot.workspace = true http.workspace = true httpdate.workspace = true serde_json.workspace = true +serde.workspace = true tokio.workspace = true tracing.workspace = true tracing-opentelemetry.workspace = true diff --git a/crates/libsy-llm-client/src/client.rs b/crates/libsy-llm-client/src/client.rs index 6ba089d9c..b66a38b42 100644 --- a/crates/libsy-llm-client/src/client.rs +++ b/crates/libsy-llm-client/src/client.rs @@ -1,7 +1,7 @@ // SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -//! [`TranslatingLlmClient`] — the crate's single public entry point: encode a neutral +//! [`TranslatingLlmClient`]: encode a neutral //! request, call the configured backend over HTTP, decode the neutral response. use std::collections::{BTreeMap, BTreeSet, HashMap}; @@ -896,7 +896,7 @@ fn record_gen_ai_request(url: &str, model: &str, streaming: bool) { } } -fn convert_reqwest_error(error: reqwest::Error) -> LlmClientError { +pub(crate) fn convert_reqwest_error(error: reqwest::Error) -> LlmClientError { // Reqwest labels truncated or otherwise unreadable response bodies as decode // errors, so distinguish them from serde JSON failures at the call site. let error = error.without_url(); diff --git a/crates/libsy-llm-client/src/lib.rs b/crates/libsy-llm-client/src/lib.rs index d5f579c7d..ee1d9d616 100644 --- a/crates/libsy-llm-client/src/lib.rs +++ b/crates/libsy-llm-client/src/lib.rs @@ -11,6 +11,10 @@ //! back to a [`switchyard_protocol::Response`] — supporting both buffered and //! streamed responses. //! +//! [`SystemOneClient`] serves typed decision requests. Register decision targets +//! with [`ClientRouter::new_with_decision_clients`] or +//! [`ClientRouter::single_with_decision_clients`] to serve them alongside LLM calls. +//! //! [`run()`] pairs the client with a libsy algorithm: it drives //! [`switchyard_libsy::Algorithm::run_stream`], serves routing-time calls, and makes the terminal //! answer call from the routing outcome when needed. A host that just wants the answer does not @@ -24,14 +28,16 @@ mod observability; mod observation; pub mod raw; pub mod run; +mod system_one; pub use backend::{Backend, DEFAULT_MAX_RETRIES, HttpBackendConfig}; pub use client::{AuxiliaryOperation, ModelConfig, TranslatingLlmClient}; pub use error::{LlmClientError, Result}; -pub use observation::{LlmCallObservation, RunObservation, RunObserver}; +pub use observation::{ModelCallObservation, RunObservation, RunObserver}; pub use raw::RawResponse; pub use run::{ClientRouter, decide, run}; pub use switchyard_translation::RawEventStream; +pub use system_one::SystemOneClient; /// Registers process-wide compatibility gauges with the global meter provider. pub fn initialize_metrics() { diff --git a/crates/libsy-llm-client/src/observation.rs b/crates/libsy-llm-client/src/observation.rs index 308ad4b28..53e660f14 100644 --- a/crates/libsy-llm-client/src/observation.rs +++ b/crates/libsy-llm-client/src/observation.rs @@ -11,7 +11,7 @@ use switchyard_protocol::{ModelId, Usage}; /// One completed model call observed while serving an algorithm run. #[derive(Clone, Debug)] -pub struct LlmCallObservation { +pub struct ModelCallObservation { /// Model selected for the completed call. pub selected_model: ModelId, /// Whether the call completed successfully. @@ -22,15 +22,17 @@ pub struct LlmCallObservation { pub usage: Option, } -/// Events emitted inline while [`crate::run`] serves a routing request. +/// Events emitted inline while [`crate::run()`] serves a routing request. #[derive(Clone, Debug)] pub enum RunObservation { /// Metadata attached to the completed routing outcome. Outcome(OutcomeMetadata), /// A completed model call requested by the algorithm for routing work. - LlmCall(LlmCallObservation), + LlmCall(ModelCallObservation), + /// A completed decision call requested by the algorithm for routing work. + DecisionCall(ModelCallObservation), /// A completed terminal model call made from the routing outcome. - AnswerCall(LlmCallObservation), + AnswerCall(ModelCallObservation), /// Routing time recorded by the `switchyard.routing_overhead_ms` metric. RoutingOverhead(Duration), } diff --git a/crates/libsy-llm-client/src/run.rs b/crates/libsy-llm-client/src/run.rs index a3dfc339a..dcb4a6f7d 100644 --- a/crates/libsy-llm-client/src/run.rs +++ b/crates/libsy-llm-client/src/run.rs @@ -32,12 +32,12 @@ use switchyard_libsy::{ }; use switchyard_protocol::{ AggLlmResponse, LlmClientError, LlmResponse, LlmResponseChunk, LlmResponseStream, Message, - ModelId, Request, Response, ResponseAccumulator, RoutedLlmClient, RoutingFallbackReason, - WireFormat, + ModelId, Request, Response, ResponseAccumulator, RoutedDecisionClient, RoutedLlmClient, + RoutingFallbackReason, WireFormat, }; use switchyard_translation::prepare_request_for_target; -use crate::observation::{LlmCallObservation, RunObservation, RunObserver}; +use crate::observation::{ModelCallObservation, RunObservation, RunObserver}; use crate::{metrics, observability}; /// Run one request to completion, serving every offloaded model call with `client`. @@ -152,18 +152,38 @@ pub async fn decide( Ok(outcome) } -async fn unsupported_decision(call: CallDecision) -> Result<()> { - let model = call.model.clone(); - call.respond(Err(LibsyError::client_call( - model, - LlmClientError::General("decision calls are not supported by this client".to_string()), - ))) +#[tracing::instrument( + name = "libsy.decision_client_call", + skip_all, + fields(algorithm = %call.algorithm, selected_model = %call.model, outcome = tracing::field::Empty), +)] +async fn serve_decision( + clients: &ClientRouter, + call: CallDecision, + observations: &Option>>>, +) -> Result<()> { + let client = clients.route_decision(&call.model); + let started = Instant::now(); + let result = async { client?.call(call.request.clone()).await }.await; + tracing::Span::current().record("outcome", if result.is_ok() { "ok" } else { "error" }); + if let Some(observations) = observations { + observations + .lock() + .push(RunObservation::DecisionCall(ModelCallObservation { + selected_model: call.model.clone(), + is_success: result.is_ok(), + duration: started.elapsed(), + usage: result.as_ref().ok().map(|response| response.usage.clone()), + })); + } + let result = result.map_err(|error| LibsyError::client_call(call.model.clone(), error)); + call.respond(result) } /// Emits completed routing calls after the outcome reveals whether one response became the answer. fn emit_routing_observations( observer: &Option, - observations: &Option>>>, + observations: &Option>>>, answered_model: Option<&ModelId>, ) { let (Some(observer), Some(observations)) = (observer, observations) else { @@ -171,31 +191,36 @@ fn emit_routing_observations( }; let mut answer_observed = false; for observation in observations.lock().drain(..) { - if !answer_observed && answered_model == Some(&observation.selected_model) { - answer_observed = true; - observer(RunObservation::AnswerCall(observation)); - } else { - observer(RunObservation::LlmCall(observation)); - } + observer(match observation { + RunObservation::LlmCall(call) + if !answer_observed && answered_model == Some(&call.selected_model) => + { + answer_observed = true; + RunObservation::AnswerCall(call) + } + observation => observation, + }); } } /// Serve one offloaded call and fulfill its promise. /// -/// LLM failures stop the run unless the call enables recovery. Unsupported decisions -/// return an error to the algorithm. +/// LLM failures stop the run unless the call enables recovery. Decision results, +/// including client failures, go back to the algorithm for its routing policy. async fn serve( clients: ClientRouter, call: Call, - observations: Option>>>, + observations: Option>>>, ) -> Result<()> { let call = match call { Call::Model(call) => *call, - Call::Decision(call) => return unsupported_decision(*call).await, + Call::Decision(call) => return serve_decision(&clients, *call, &observations).await, }; let observe = |observation| { if let Some(observations) = &observations { - observations.lock().push(observation); + observations + .lock() + .push(RunObservation::LlmCall(observation)); } }; let target = call.models.first().ok_or(LibsyError::NoTargets)?; @@ -224,7 +249,7 @@ async fn call_first_available( algorithm: &str, request: &Request, models: &[ModelId], - observe: &(dyn Fn(LlmCallObservation) + Send + Sync), + observe: &(dyn Fn(ModelCallObservation) + Send + Sync), ) -> Result { for (index, target) in models.iter().enumerate() { let request = clients.prepare_completion_request(request.clone(), target); @@ -301,7 +326,7 @@ async fn call_one( model_id: &ModelId, request: Request, algorithm: &str, - observe: &(dyn Fn(LlmCallObservation) + Send + Sync), + observe: &(dyn Fn(ModelCallObservation) + Send + Sync), // index is for span log index: usize, // count is for span log @@ -342,7 +367,7 @@ async fn call_one( } else { result }; - observe(LlmCallObservation { + observe(ModelCallObservation { selected_model: model_id.clone(), is_success: result.is_ok(), duration, @@ -557,8 +582,7 @@ fn conversation_id(fields: &serde_json::Map) -> Option<&str> { /// /// An algorithm routes among named targets; which provider each target lives on is the /// host's concern, and two targets in one run may sit on different providers. A router owns -/// that mapping. It is *not* itself a client: it hands back a [`RoutedLlmClient`] and the -/// caller makes the call. +/// that mapping and resolves the client for each kind of call. /// /// Cloning is cheap — the mapping is shared, so one router can serve every request. #[derive(Clone)] @@ -568,6 +592,7 @@ pub struct ClientRouter { struct ClientRouting { routing: Routing, + decision_clients: HashMap>, target_prompts: HashMap, routing_answer_target: Option, /// Provider-owned routing pins and materialized cross-format history. @@ -613,7 +638,7 @@ impl ClientRouter { /// Build a router with an explicit list of models that can answer requests. /// - /// `by_model` contains every callable model, including classifiers. Only the model IDs in + /// `by_model` contains callable LLMs, including LLM classifiers. Only the model IDs in /// `completion_targets` determine whether native Responses state needs a routing pin. /// Cross-format state is recorded lazily when a Responses request uses a Chat or Anthropic /// target. `target_prompts` and `routing_answer_target` have the same meaning as in @@ -623,6 +648,24 @@ impl ClientRouter { target_prompts: HashMap, routing_answer_target: Option, completion_targets: &[ModelId], + ) -> Self { + Self::new_with_decision_clients( + by_model, + HashMap::new(), + target_prompts, + routing_answer_target, + completion_targets, + ) + } + + /// Build a router with separate LLM and decision mappings, shared across clones. + /// LLM prompt and completion settings follow [`Self::new_with_completion_targets`]. + pub fn new_with_decision_clients( + by_model: HashMap>, + decision_clients: HashMap>, + target_prompts: HashMap, + routing_answer_target: Option, + completion_targets: &[ModelId], ) -> Self { let mut clients = completion_targets .iter() @@ -633,6 +676,7 @@ impl ClientRouter { Self { inner: Arc::new(ClientRouting { routing: Routing::ByModel(by_model), + decision_clients, target_prompts, routing_answer_target, state_owners: Mutex::default(), @@ -647,9 +691,18 @@ impl ClientRouter { /// backends internally and rejects ones it does not know, so enumerating them here would /// only duplicate that. pub fn single(client: Arc) -> Self { + Self::single_with_decision_clients(client, HashMap::new()) + } + + /// Use one client for all LLM targets and explicit clients for decision targets. + pub fn single_with_decision_clients( + client: Arc, + decision_clients: HashMap>, + ) -> Self { Self { inner: Arc::new(ClientRouting { routing: Routing::Single(client), + decision_clients, target_prompts: HashMap::new(), routing_answer_target: None, state_owners: Mutex::default(), @@ -678,6 +731,19 @@ impl ClientRouter { } } + /// Resolve a decision target without falling back to an LLM client. + fn route_decision( + &self, + model: &ModelId, + ) -> std::result::Result<&Arc, LlmClientError> { + self.inner + .decision_clients + .get(model) + .ok_or_else(|| LlmClientError::Configuration { + message: format!("no decision client is configured for model {model:?}"), + }) + } + /// Return the recorded model for the requested response or conversation ID. fn stored_state_owner(&self, request: &Request) -> Option { let fields = &request.llm_request.extensions.fields; @@ -1072,6 +1138,239 @@ mod tests { } } + #[tokio::test] + async fn system_one_serves_decisions_and_returns_failures_to_the_classifier() + -> std::result::Result<(), Box> { + use std::time::Duration; + use switchyard_libsy::{ + CapabilityJudgeConfig, DecisionJudgeConfig, LlmClassifierConfig, LlmTaskClassifier, + TaskClassifierConfig, + }; + use switchyard_protocol::{ + BooleanEstimate, DecisionRequest, DecisionValue, Probability, ScoreValue, + }; + use wiremock::matchers::{body_partial_json, header, path}; + + let server = MockServer::start().await; + let client = Arc::new(crate::SystemOneClient::new( + format!("{}/v1/systemone", server.uri()).parse()?, + "test-key".into(), + Duration::from_secs(2), + )?); + assert!(matches!( + crate::SystemOneClient::new( + format!("{}/v1/systemone", server.uri()).parse()?, + "invalid\nkey".into(), + Duration::from_secs(2), + ), + Err(LlmClientError::Configuration { .. }) + )); + let mixed: DecisionRequest = serde_json::from_value(json!({ + "model": "judge", "context": {"task": "A simple task"}, + "questions": { + "boolean": {"instructions": "Is it simple?", "kind": {"type": "boolean", "data": { + "true_description": ["Simple"], "false_description": null + }}}, + "route": {"instructions": {"task": "Choose"}, "kind": {"type": "choice", "data": { + "options": [{"id": "advantage", "description": "A wins"}, {"id": "no_advantage"}] + }}}, + "score": {"instructions": "Rate difficulty", "kind": {"type": "score", "data": { + "levels": ["Easy", {"description": "Hard"}] + }}} + } + }))?; + let mut body = json!({ + "model": "jev-1.13.0", "usage": {"input_tokens": 42, "output_tokens": 3}, + "answers": { + "boolean": {"type": "noul", "noul": 0.9}, + "route": {"type": "choice", "choice": "no_advantage", "confidence": 0.7, + "probabilities": {"advantage": 0.2, "no_advantage": 0.8}}, + "score": {"type": "score", "score": 0.6, "confidence": 0.3, + "legend": {"0": "Easy", "1": {"description": "Hard"}}, + "probabilities": {"1": 0.6, "0": 0.4}} + } + }); + Mock::given(method("POST")) + .and(path("/v1/systemone")) + .and(header("authorization", "Bearer test-key")) + .and(body_partial_json(json!({ + "model": "judge", "state": mixed.context, + "questions": { + "boolean": {"type": "noul", "criteria": {"true": ["Simple"]}}, + "route": {"type": "choice", "instructions": {"task": "Choose"}, + "criteria": {"advantage": "A wins", "no_advantage": null}}, + "score": {"type": "score", "criteria": ["Easy", {"description": "Hard"}]} + } + }))) + .respond_with(ResponseTemplate::new(200).set_body_json(&body)) + .expect(1) + .mount(&server) + .await; + let response = client.call(mixed.clone()).await?; + assert_eq!(response.model.as_deref(), Some("jev-1.13.0")); + assert_eq!(response.id, None); + assert_eq!(response.usage.input_tokens, Some(42)); + assert_eq!(response.usage.total_tokens, None); + assert_eq!( + response.answers["boolean"].value, + DecisionValue::Boolean(BooleanEstimate::ProbabilityTrue(Probability(0.9))) + ); + assert_eq!( + response.answers["score"].value, + DecisionValue::Score { + value: ScoreValue(0.6), + probabilities: Some(vec![Probability(0.4), Probability(0.6)]), + } + ); + assert_eq!( + response.answers["route"].provider_confidence.map(|c| c.0), + Some(0.7) + ); + server.verify().await; + + let mut partial = body.clone(); + partial["answers"] + .as_object_mut() + .unwrap() + .remove("boolean"); + let mut invalid_responses = vec![json!({"answers": {}}), partial]; + for (id, answer) in [ + ("unknown-question", json!({"type": "noul", "noul": 0.9})), + ("boolean", json!({"type": "score", "score": 0.0})), + ( + "route", + json!({"type": "choice", "choice": "unknown-option", + "probabilities": {"advantage": 0.2, "no_advantage": 0.8}}), + ), + ("score", json!({"type": "score", "score": -0.1})), + ( + "score", + json!({"type": "score", "score": 1.1, + "probabilities": {"0": 0.4, "1": 0.6}}), + ), + ] { + let mut invalid = body.clone(); + invalid["answers"][id] = answer; + if id == "unknown-question" { + invalid["answers"] + .as_object_mut() + .unwrap() + .remove("boolean"); + } + invalid_responses.push(invalid); + } + for template in invalid_responses + .into_iter() + .map(|body| ResponseTemplate::new(200).set_body_json(body)) + .chain([ResponseTemplate::new(200).set_body_string("invalid JSON")]) + { + server.reset().await; + Mock::given(method("POST")) + .respond_with(template) + .expect(1) + .mount(&server) + .await; + assert!(matches!( + client.call(mixed.clone()).await, + Err(LlmClientError::ResponseTranslation(_) | LlmClientError::InvalidResponse { .. }) + )); + server.verify().await; + } + + body["answers"].as_object_mut().unwrap().remove("boolean"); + body["answers"].as_object_mut().unwrap().remove("score"); + let models = Arc::new(RuntimeModels::new(HashMap::from([ + (Category::Judge, vec!["judge".into()]), + (Category::Capable, vec!["capable".into()]), + (Category::Efficient, vec!["efficient".into()]), + (Category::Any, vec!["efficient".into(), "capable".into()]), + ]))); + let algorithm = || -> Result> { + Ok(Arc::new(LlmTaskClassifier::new( + LlmClassifierConfig::Capability { + config: TaskClassifierConfig { + judge: CapabilityJudgeConfig::Decision(DecisionJudgeConfig { + cutoff: 0.4, + instructions: None, + candidates: BTreeMap::from([ + ("a".into(), "capable".into()), + ("b".into(), "efficient".into()), + ]), + evidence: json!({}), + }), + ..TaskClassifierConfig::default() + }, + }, + )?)) + }; + for (template, expected, is_success) in [ + ( + ResponseTemplate::new(200).set_body_json(&body), + "efficient", + true, + ), + ( + ResponseTemplate::new(503).set_body_string("unavailable"), + "capable", + false, + ), + ] { + server.reset().await; + Mock::given(method("POST")) + .and(path("/v1/systemone")) + .respond_with(template) + .expect(2) + .mount(&server) + .await; + let llm = Arc::new(CandidateClient { + calls: Mutex::default(), + requests: Mutex::default(), + first: FirstOutcome::Unauthorized, + }); + let clients = ClientRouter::single_with_decision_clients( + llm.clone(), + HashMap::from([( + ModelId::from("judge"), + client.clone() as Arc, + )]), + ); + let outcome = decide(algorithm()?, clients.clone(), request(), models.clone()).await?; + assert_eq!(outcome.selected_model_id()?, expected); + assert!(llm.calls.lock().is_empty()); + let events = Arc::new(Mutex::new(Vec::new())); + let captured = events.clone(); + let (selected, _) = run( + algorithm()?, + clients, + request(), + models.clone(), + Some(Arc::new(move |event| captured.lock().push(event))), + ) + .await?; + assert_eq!(selected, expected); + assert_eq!(&*llm.calls.lock(), &[ModelId::from(expected)]); + { + let events = events.lock(); + let RunObservation::DecisionCall(call) = &events[0] else { + panic!("missing decision observation") + }; + assert_eq!(call.selected_model, "judge"); + assert_eq!(call.is_success, is_success); + assert_eq!( + call.usage.as_ref().and_then(|u| u.input_tokens), + is_success.then_some(42) + ); + assert!( + events + .iter() + .any(|event| matches!(event, RunObservation::AnswerCall(_))) + ); + } + server.verify().await; + } + Ok(()) + } + fn instruction_text(request: &Request) -> Vec<&str> { request .llm_request @@ -1159,12 +1458,13 @@ mod tests { fn answer_observation_keeps_call_order() { let pending = Some(Arc::new(Mutex::new( ["answer", "judge"] - .map(|model| LlmCallObservation { + .map(|model| ModelCallObservation { selected_model: model.into(), is_success: true, duration: std::time::Duration::ZERO, usage: None, }) + .map(RunObservation::LlmCall) .into(), ))); let emitted = Arc::new(Mutex::new(Vec::new())); diff --git a/crates/libsy-llm-client/src/system_one.rs b/crates/libsy-llm-client/src/system_one.rs new file mode 100644 index 000000000..611832ca8 --- /dev/null +++ b/crates/libsy-llm-client/src/system_one.rs @@ -0,0 +1,243 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Buffered System One calls using the provider-neutral decision IR. + +use std::collections::BTreeMap; +use std::time::Duration; + +use async_trait::async_trait; +use reqwest::{Url, header::HeaderValue}; +use serde::Deserialize; +use serde_json::{Value, json}; +use switchyard_protocol::{ + BooleanEstimate, DecisionAnswer, DecisionKind, DecisionRequest, DecisionResponse, + DecisionValue, ModelId, Probability, ProviderConfidence, RoutedDecisionClient, ScoreValue, + Usage, +}; + +use crate::client::convert_reqwest_error; +use crate::{LlmClientError, Result, metrics}; + +/// Serves decision requests through a System One endpoint, such as TypeSafe's Jev API. +pub struct SystemOneClient { + client: reqwest::Client, + endpoint: Url, + api_key: String, +} + +impl SystemOneClient { + /// `endpoint` is the full URL, including `/v1/systemone`. Each call makes one + /// attempt; `timeout` covers sending the request and reading the response body. + pub fn new(endpoint: Url, api_key: String, timeout: Duration) -> Result { + if HeaderValue::try_from(format!("Bearer {api_key}")).is_err() { + return Err(LlmClientError::Configuration { + message: "System One API key cannot be encoded as an HTTP header".into(), + }); + } + let client = reqwest::Client::builder() + .timeout(timeout) + .redirect(reqwest::redirect::Policy::none()) + .build() + .map_err(convert_reqwest_error)?; + Ok(Self { + client, + endpoint, + api_key, + }) + } +} + +#[async_trait] +impl RoutedDecisionClient for SystemOneClient { + async fn call(&self, request: DecisionRequest) -> Result { + let response = self + .client + .post(self.endpoint.clone()) + .bearer_auth(&self.api_key) + .json(&encode(&request)?) + .send() + .await + .inspect_err(|_| metrics::record_upstream_attempt(None)) + .map_err(convert_reqwest_error)?; + let status = response.status(); + let body = response + .bytes() + .await + .inspect_err(|_| metrics::record_upstream_attempt(None)) + .map_err(convert_reqwest_error)?; + metrics::record_upstream_attempt(Some(status.as_u16())); + if !status.is_success() { + return Err(LlmClientError::UpstreamHttp { + status, + body: String::from_utf8_lossy(&body).replace(&self.api_key, "[REDACTED]"), + }); + } + let response: WireResponse = + serde_json::from_slice(&body).map_err(|source| LlmClientError::InvalidResponse { + source: Box::new(source), + })?; + if !response.answers.keys().eq(request.questions.keys()) { + return Err(LlmClientError::ResponseTranslation( + "System One answer keys do not match the requested question keys".into(), + )); + } + let answers = response + .answers + .into_iter() + .map(|(id, answer)| { + let value = match (request.questions.get(&id).map(|q| &q.kind), answer.value) { + (Some(DecisionKind::Boolean { .. }), WireValue::Noul { noul }) => { + DecisionValue::Boolean(BooleanEstimate::ProbabilityTrue(noul)) + } + ( + Some(DecisionKind::Choice { options }), + WireValue::Choice { + choice, + probabilities, + }, + ) if options.iter().any(|option| option.id == choice) => { + DecisionValue::Choice { + selected: choice, + probabilities, + } + } + ( + Some(DecisionKind::Score { levels }), + WireValue::Score { + score, + probabilities, + }, + ) if !levels.is_empty() + && (0.0..=(levels.len() - 1) as f64).contains(&score.0) => + { + let probabilities = probabilities + .map(|mut probabilities| { + let invalid_rubric = || { + LlmClientError::ResponseTranslation(format!( + "score probabilities for {id:?} do not match its rubric" + )) + }; + if probabilities.len() != levels.len() { + return Err(invalid_rubric()); + } + // JSON keys are strings; order probabilities by the request's rubric. + (0..levels.len()) + .map(|index| { + probabilities + .remove(&index.to_string()) + .ok_or_else(invalid_rubric) + }) + .collect() + }) + .transpose()?; + DecisionValue::Score { + value: score, + probabilities, + } + } + _ => { + return Err(LlmClientError::ResponseTranslation(format!( + "answer {id:?} has an unexpected question ID, kind, or value" + ))); + } + }; + Ok(( + id, + DecisionAnswer { + value, + provider_confidence: answer.confidence, + }, + )) + }) + .collect::>()?; + Ok(DecisionResponse { + id: response.id, + model: response.model, + answers, + usage: response.usage, + }) + } +} + +fn encode(request: &DecisionRequest) -> Result { + let model = request + .model + .as_ref() + .ok_or_else(|| LlmClientError::InvalidRequest { + message: "System One requires a selected model".into(), + })?; + let mut questions = serde_json::Map::new(); + for (id, question) in &request.questions { + let (kind, criteria) = match &question.kind { + DecisionKind::Boolean { + true_description, + false_description, + } => { + let mut criteria = serde_json::Map::new(); + for (key, description) in [("true", true_description), ("false", false_description)] + { + if let Some(description) = description { + criteria.insert(key.into(), description.clone()); + } + } + ("noul", Value::Object(criteria)) + } + DecisionKind::Choice { options } => { + let mut criteria = serde_json::Map::new(); + for option in options { + let description = option.description.as_ref().unwrap_or(&Value::Null); + if criteria + .insert(option.id.clone(), description.clone()) + .is_some() + { + return Err(LlmClientError::InvalidRequest { + message: format!("duplicate option {:?} in questions.{id}", option.id), + }); + } + } + ("choice", Value::Object(criteria)) + } + DecisionKind::Score { levels } => ("score", json!(levels)), + }; + questions.insert( + id.clone(), + json!({ + "type": kind, "instructions": question.instructions, "criteria": criteria, + }), + ); + } + Ok(json!({"model": model, "state": request.context, "questions": questions})) +} + +#[derive(Deserialize)] +struct WireResponse { + id: Option, + model: Option, + answers: BTreeMap, + #[serde(default)] + usage: Usage, +} + +#[derive(Deserialize)] +struct WireAnswer { + #[serde(flatten)] + value: WireValue, + confidence: Option, +} + +#[derive(Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +enum WireValue { + Noul { + noul: Probability, + }, + Choice { + choice: String, + probabilities: Option>, + }, + Score { + score: ScoreValue, + probabilities: Option>, + }, +} diff --git a/crates/protocol/src/client.rs b/crates/protocol/src/client.rs index 8d312e72d..fa37f01c1 100644 --- a/crates/protocol/src/client.rs +++ b/crates/protocol/src/client.rs @@ -1,17 +1,15 @@ // SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -//! The routed-call server trait and its shared error types. +//! Routed-call client contracts and shared error types. //! -//! [`RoutedLlmClient`] is the one piece of I/O the protocol does not own: a host -//! implements it to actually perform a model call. It lives here — rather than in -//! libsy's orchestration crate — so a client crate that depends only on the protocol -//! can serve routed calls without pulling in the orchestrator. +//! Hosts implement these traits to perform model calls. Keeping the contracts here +//! lets clients depend on the protocol without pulling in libsy's orchestration. use async_trait::async_trait; use thiserror::Error; -use crate::{ModelId, Request, Response}; +use crate::{DecisionRequest, DecisionResponse, ModelId, Request, Response}; /// A boxed client-specific error preserved as the source of a routed call failure. pub type BoxError = Box; @@ -157,3 +155,11 @@ pub trait RoutedLlmClient: Send + Sync { /// Make a request async fn call(&self, request: Request) -> Result; } + +/// Serves typed decision calls over a host-owned transport. +/// Implementations may be shared across targets and called concurrently. +#[async_trait] +pub trait RoutedDecisionClient: Send + Sync { + /// Evaluate the request using its selected model. + async fn call(&self, request: DecisionRequest) -> Result; +} diff --git a/crates/switchyard-nemo-relay-plugin/src/runtime.rs b/crates/switchyard-nemo-relay-plugin/src/runtime.rs index 95794f0d8..4c1b7ab05 100644 --- a/crates/switchyard-nemo-relay-plugin/src/runtime.rs +++ b/crates/switchyard-nemo-relay-plugin/src/runtime.rs @@ -11,7 +11,7 @@ use nemo_relay_plugin::{ MetricValueType, PluginRuntime, }; use serde_json::{Map, json}; -use switchyard_llm_client::{LlmCallObservation, RunObservation, RunObserver}; +use switchyard_llm_client::{ModelCallObservation, RunObservation, RunObserver}; use switchyard_protocol::{ LlmClientError, LlmResponse, LlmResponseChunk, LlmStreamError, Metadata, ProviderExtensions, Request, Response, Usage, WireFormat, @@ -336,7 +336,23 @@ impl SwitchyardRuntime { } RunObservation::LlmCall(call) => { call_index += 1; - self.routing_call_events(events, call, call_index, metadata); + self.routing_call_events( + events, + call, + call_index, + metadata, + "switchyard.routing.llm_call", + ); + } + RunObservation::DecisionCall(call) => { + call_index += 1; + self.routing_call_events( + events, + call, + call_index, + metadata, + "switchyard.routing.decision_call", + ); } RunObservation::RoutingOverhead(duration) => { let latency_ms = duration.as_secs_f64() * 1_000.0; @@ -374,15 +390,16 @@ impl SwitchyardRuntime { fn routing_call_events( &self, events: &mut Vec, - call: LlmCallObservation, + call: ModelCallObservation, call_index: usize, metadata: &Json, + mark_name: &str, ) { let outcome = if call.is_success { "ok" } else { "error" }; let latency_ms = call.duration.as_secs_f64() * 1_000.0; let token_metrics = token_usage_metrics("routing", &call, metadata); events.push(RoutingEvent::Mark(RoutingMark { - name: "switchyard.routing.llm_call".into(), + name: mark_name.into(), data: json!({ "call_index": call_index, "selected_model": call.selected_model.as_str(), @@ -717,7 +734,7 @@ fn routing_overhead_metric(latency_ms: f64, metadata: Json) -> RoutingEvent { fn token_usage_metrics( call_role: &str, - call: &LlmCallObservation, + call: &ModelCallObservation, metadata: &Json, ) -> Vec { let Some(usage) = call.usage.as_ref() else { @@ -1299,7 +1316,7 @@ mod tests { runtime.emit_observations( &mut events, vec![ - RunObservation::LlmCall(LlmCallObservation { + RunObservation::LlmCall(ModelCallObservation { selected_model: ModelId::from("routing-model"), is_success: false, duration: std::time::Duration::from_millis(12), @@ -1390,7 +1407,7 @@ mod tests { #[test] fn token_usage_metrics_distinguish_routing_and_answer_targets() { - let call = LlmCallObservation { + let call = ModelCallObservation { selected_model: ModelId::from("judge-model"), is_success: true, duration: std::time::Duration::from_millis(1), @@ -1452,13 +1469,13 @@ mod tests { runtime.emit_observations( &mut events, vec![ - RunObservation::AnswerCall(LlmCallObservation { + RunObservation::AnswerCall(ModelCallObservation { selected_model: ModelId::from("weak-target"), is_success: false, duration: std::time::Duration::from_millis(2), usage: None, }), - RunObservation::AnswerCall(LlmCallObservation { + RunObservation::AnswerCall(ModelCallObservation { selected_model: ModelId::from("strong-target"), is_success: true, duration: std::time::Duration::from_millis(3), diff --git a/crates/switchyard-server/src/lib.rs b/crates/switchyard-server/src/lib.rs index 0f1738936..842f77d92 100644 --- a/crates/switchyard-server/src/lib.rs +++ b/crates/switchyard-server/src/lib.rs @@ -445,7 +445,7 @@ fn stats_observer( stats.record_error(&call.selected_model); } } - RunObservation::LlmCall(call) => { + RunObservation::LlmCall(call) | RunObservation::DecisionCall(call) => { let latency_ms = call.duration.as_secs_f64() * 1_000.0; if call.is_success { if let (Some((log, context)), Some(usage)) = @@ -1685,7 +1685,7 @@ fn endpoint_listing(has_routing_log: bool) -> String { #[cfg(test)] mod tests { - use switchyard_llm_client::LlmCallObservation; + use switchyard_llm_client::ModelCallObservation; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::sync::{Notify, oneshot}; @@ -1706,7 +1706,7 @@ mod tests { let observer = stats_observer(StatsAccumulator::default(), Some((log.clone(), context))); let call = |model: &str, answer: bool| { - let observation = LlmCallObservation { + let observation = ModelCallObservation { selected_model: ModelId::from(model), is_success: true, duration: Duration::from_millis(3),