Skip to content
Merged
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
120 changes: 106 additions & 14 deletions crates/switchyard-py/src/libsy_bindings.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,28 +3,26 @@

//! Minimal Python API for running Rust-owned libsy algorithms.

use std::collections::HashMap;
use std::collections::{BTreeMap, HashMap};
use std::sync::Arc;

use futures::StreamExt;
use http::HeaderMap;
use http::header::{HeaderName, HeaderValue};
use pyo3::exceptions::{
PyBaseException, PyNotImplementedError, PyStopAsyncIteration, PyTypeError, PyValueError,
};
use pyo3::exceptions::{PyBaseException, PyStopAsyncIteration, PyTypeError, PyValueError};
use pyo3::prelude::*;
use serde_json::Value;
use switchyard_libsy::{
Algorithm, CallModel, CapabilityJudgeConfig, ClassifierContractConfig,
Algorithm, CallDecision, CallModel, CapabilityJudgeConfig, ClassifierContractConfig,
ClassifierResponseFormat, ClassifyTrigger, CustomClassifierConfig, CustomClassifierPolicy,
DeescalationConfig, EscalationJudgeConfig, HandoffNoteConfig, LibsyError as RustLibsyError,
LlmCapabilityConfig, LlmClassifierConfig, LlmFallback, LlmTaskClassifier, Noop, PickerMode,
Random, RoutingOutcome, RuntimeModels, StageRouter, StageRouterConfig, Step as RustStep,
StepStream, TaskClassifierConfig, ToolSemantics,
DecisionJudgeConfig, DeescalationConfig, EscalationJudgeConfig, HandoffNoteConfig,
LibsyError as RustLibsyError, LlmCapabilityConfig, LlmClassifierConfig, LlmFallback,
LlmTaskClassifier, Noop, PickerMode, Random, RoutingOutcome, RuntimeModels, StageRouter,
StageRouterConfig, Step as RustStep, StepStream, TaskClassifierConfig, ToolSemantics,
};
use switchyard_protocol::{
Category, LlmClientError, LlmResponse, LlmResponseStream, LlmResponseStreamEvent, Metadata,
ModelId, Request, Response,
Category, DecisionResponse, LlmClientError, LlmResponse, LlmResponseStream,
LlmResponseStreamEvent, Metadata, ModelId, Request, Response,
};
use tokio::sync::Mutex;

Expand Down Expand Up @@ -306,6 +304,31 @@ impl PyLlmClassifierConfig {

#[pymethods]
impl PyTaskClassifierConfig {
/// Use a decision model's relative-advantage probability to select a target.
#[staticmethod]
#[pyo3(signature = (*, cutoff, candidates, evidence, instructions=None))]
fn decision(
cutoff: f64,
candidates: BTreeMap<String, String>,
evidence: &Bound<'_, PyAny>,
instructions: Option<&Bound<'_, PyAny>>,
) -> PyResult<Self> {
Ok(Self {
inner: TaskClassifierConfig {
judge: CapabilityJudgeConfig::Decision(DecisionJudgeConfig {
cutoff,
candidates: candidates
.into_iter()
.map(|(label, model)| (label, ModelId::from(model)))
.collect(),
evidence: from_python(evidence)?,
instructions: instructions.map(from_python).transpose()?,
}),
..TaskClassifierConfig::default()
},
})
}

#[new]
#[pyo3(signature = (
base_threshold,
Expand Down Expand Up @@ -564,6 +587,70 @@ impl PyModelCall {
}
}

/// A decision request owned by the Python host until it replies or drops the call.
#[pyclass(name = "DecisionCall", module = "switchyard.libsy")]
struct PyDecisionCall {
inner: Option<CallDecision>,
algorithm: String,
request: Py<PyAny>,
model: ModelId,
}

impl PyDecisionCall {
fn new(py: Python<'_>, call: CallDecision) -> PyResult<Self> {
let request = to_python(py, &call.request)?;
Ok(Self {
algorithm: call.algorithm.clone(),
model: call.model.clone(),
inner: Some(call),
request,
})
}

fn take(&mut self) -> PyResult<CallDecision> {
self.inner
.take()
.ok_or_else(|| py_libsy_error("decision call has already been completed"))
}
}

#[pymethods]
impl PyDecisionCall {
#[getter]
fn algorithm(&self) -> &str {
&self.algorithm
}

/// The provider-neutral decision request as a Python dictionary.
#[getter]
fn request(&self, py: Python<'_>) -> Py<PyAny> {
self.request.clone_ref(py)
}

#[getter]
fn model(&self) -> &str {
&self.model
}

/// Return a provider-neutral decision response to the waiting algorithm.
fn respond(&mut self, response: &Bound<'_, PyAny>) -> PyResult<()> {
let response: DecisionResponse = from_python(response)?;
self.take()?.respond(Ok(response)).map_err(py_libsy_error)
}

/// Return a Python client failure so the algorithm can apply its fallback policy.
fn fail(&mut self, error: &Bound<'_, PyAny>) -> PyResult<()> {
if !error.is_instance_of::<PyBaseException>() {
return Err(PyTypeError::new_err("error must derive from BaseException"));
}
let source = python_client_error(error.py(), PyErr::from_value(error.clone()), &self.model);
let model = self.model.clone();
self.take()?
.respond(Err(RustLibsyError::client_call(model, source)))
.map_err(py_libsy_error)
}
}

/// Identity and optional JSON evidence from the Rust routing outcome.
#[pyclass(name = "OutcomeMetadata", module = "switchyard.libsy", frozen)]
struct PyOutcomeMetadata {
Expand Down Expand Up @@ -677,6 +764,8 @@ fn response_to_python(py: Python<'_>, response: LlmResponse) -> PyResult<Py<PyAn
enum PyStep {
/// The host must serve the model call before the algorithm can continue.
CallModel { call: Py<PyModelCall> },
/// The host must serve the decision call before the algorithm can continue.
CallDecision { call: Py<PyDecisionCall> },
/// The terminal routing outcome.
Done { outcome: Py<PyRoutingOutcome> },
}
Expand Down Expand Up @@ -769,9 +858,11 @@ impl PyAlgorithm {

fn step_to_python(step: RustStep) -> PyResult<PyStep> {
match step {
RustStep::CallDecision(_) => Err(PyNotImplementedError::new_err(
"decision calls are not supported by the Python bindings",
)),
RustStep::CallDecision(call) => Python::attach(|py| {
Ok(PyStep::CallDecision {
call: Py::new(py, PyDecisionCall::new(py, *call)?)?,
})
}),
RustStep::CallModel(call) => Python::attach(|py| {
Ok(PyStep::CallModel {
call: Py::new(py, PyModelCall::new(py, *call)?)?,
Expand Down Expand Up @@ -939,6 +1030,7 @@ pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
let libsy_module = PyModule::new(module.py(), "libsy")?;
libsy_module.add_class::<PyAlgorithm>()?;
libsy_module.add_class::<PyCustomClassifierConfig>()?;
libsy_module.add_class::<PyDecisionCall>()?;
libsy_module.add_class::<PyDeescalationConfig>()?;
libsy_module.add_class::<PyEscalationClassifierConfig>()?;
libsy_module.add_class::<PyLlmClassifierConfig>()?;
Expand Down
4 changes: 3 additions & 1 deletion switchyard/libsy/__init__.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,13 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Run Rust-owned libsy algorithms with Python-hosted LLM clients."""
"""Run Rust-owned libsy algorithms with Python-hosted model clients."""

from switchyard_rust.libsy import (
Algorithm,
ContextWindowExceededError,
CustomClassifierConfig,
DecisionCall,
DeescalationConfig,
EscalationClassifierConfig,
LibsyError,
Expand All @@ -26,6 +27,7 @@
"Algorithm",
"ContextWindowExceededError",
"CustomClassifierConfig",
"DecisionCall",
"DeescalationConfig",
"EscalationClassifierConfig",
"LibsyError",
Expand Down
34 changes: 33 additions & 1 deletion switchyard_rust/libsy.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
"Algorithm",
"ContextWindowExceededError",
"CustomClassifierConfig",
"DecisionCall",
"DeescalationConfig",
"EscalationClassifierConfig",
"LibsyError",
Expand Down Expand Up @@ -133,6 +134,21 @@ def respond(self, response: LlmResponse.Agg | LlmResponse.Stream) -> None: ...

def fail(self, error: BaseException) -> None: ...

@final
class DecisionCall:
@property
def algorithm(self) -> str: ...

@property
def request(self) -> dict[str, object]: ...

@property
def model(self) -> str: ...

def respond(self, response: Mapping[str, object]) -> None: ...

def fail(self, error: BaseException) -> None: ...

@final
class OutcomeMetadata:
"""Read-only outcome identity and optional algorithm evidence."""
Expand Down Expand Up @@ -166,6 +182,11 @@ class CallModel:
__match_args__: ClassVar[tuple[Literal["call"]]] = ("call",)
call: ModelCall

@final
class CallDecision:
__match_args__: ClassVar[tuple[Literal["call"]]] = ("call",)
call: DecisionCall

@final
class Done:
__match_args__: ClassVar[tuple[Literal["outcome"]]] = ("outcome",)
Expand All @@ -179,6 +200,17 @@ class TaskClassifierConfig:
positive, and ``message_hash_fallback`` requires ``session_affinity``.
"""

@staticmethod
def decision(
*,
cutoff: float,
candidates: Mapping[str, str],
evidence: object,
instructions: object | None = None,
) -> TaskClassifierConfig:
"""Use relative advantage with the default routing and fallback settings."""
...

def __init__(
self,
base_threshold: float,
Expand Down Expand Up @@ -246,7 +278,7 @@ def run_stream(
models: Mapping[str, Sequence[str]],
subagent_models: Mapping[str, Sequence[str]] | None = None,
headers: Mapping[str, str] | None = None,
) -> AsyncIterator[Step.CallModel | Step.Done]: ...
) -> AsyncIterator[Step.CallModel | Step.CallDecision | Step.Done]: ...

def noop() -> Algorithm: ...

Expand Down
Loading