Skip to content
Open
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
10 changes: 8 additions & 2 deletions crates/libsy/src/algorithms/llm_class.rs
Original file line number Diff line number Diff line change
Expand Up @@ -454,7 +454,8 @@ impl CustomClassifierPolicy {
pub struct CustomClassifierConfig {
/// System prompt sent to the classifier judge.
pub prompt: String,
/// Inner JSON Schema placed inside the provider's structured-output wrapper.
/// Inner JSON Schema for the verdict. Sent in the provider's structured-output wrapper in
/// JSON Schema mode, or appended to the prompt in JSON Object mode.
pub response_schema: Value,
/// Deterministic policy applied after the verdict passes schema validation.
pub policy: CustomClassifierPolicy,
Expand All @@ -466,6 +467,8 @@ pub struct CustomClassifierConfig {
pub recent_turn_window: Option<usize>,
/// Maximum completion tokens available to the classifier verdict.
pub max_output_tokens: u64,
/// Structured-output mode requested from the classifier judge.
pub response_format_type: ClassifierResponseFormat,
}

impl CustomClassifierConfig {
Expand All @@ -483,6 +486,7 @@ impl CustomClassifierConfig {
message_hash_fallback: false,
recent_turn_window: None,
max_output_tokens: DEFAULT_JUDGE_MAX_OUTPUT_TOKENS,
response_format_type: ClassifierResponseFormat::default(),
}
}

Expand Down Expand Up @@ -666,8 +670,10 @@ impl LlmTaskClassifier {
message_hash_fallback,
recent_turn_window,
max_output_tokens,
response_format_type,
} = config;
let contract = ClassifierContract::from_inner_schema(&prompt, response_schema)?;
let contract =
ClassifierContract::from_inner_schema(&prompt, response_schema, response_format_type)?;
let policy = match policy {
CustomClassifierPolicy::TargetSelector { selector } => {
CustomPolicyRuntime::TargetSelector(TargetSelectorPolicy::new(selector)?)
Expand Down
212 changes: 175 additions & 37 deletions crates/libsy/src/algorithms/util/classifier_contract.rs
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,7 @@ impl ClassifierContract {
response_format_json: &str,
) -> Result<Self> {
let prompt_template = config.prompt().unwrap_or(default_prompt);
validate_prompt(prompt_template)?;
let response_format: Value =
serde_json::from_str(response_format_json).map_err(|error| {
LibsyError::AlgorithmError {
Expand All @@ -90,29 +91,33 @@ impl ClassifierContract {
message: "response schema has no json_schema.schema".to_string(),
})?;
match config.response_format_type() {
ClassifierResponseFormat::JsonSchema => {
Self::from_response_format(prompt_template, response_format, None)
}
ClassifierResponseFormat::JsonSchema => Ok(Self::from_response_format(
prompt_template.to_string(),
response_format,
None,
)),
ClassifierResponseFormat::JsonObject => {
validate_prompt(prompt_template)?;
let validator = compile_schema(schema)?;
let rendered_schema = serde_json::to_string_pretty(schema).map_err(|error| {
algorithm_error(format!("response schema could not be rendered: {error}"))
})?;
let system_prompt = format!(
"{prompt_template}\n\nReturn exactly one JSON object matching this JSON Schema:\n{rendered_schema}"
);
Self::from_response_format(
&system_prompt,
Ok(Self::from_response_format(
schema_in_prompt(prompt_template, schema)?,
json!({"type": "json_object"}),
Some(validator),
)
))
}
}
}

/// Builds a provider response format around a user-supplied inner JSON Schema.
pub(crate) fn from_inner_schema(prompt_template: &str, schema: Value) -> Result<Self> {
/// Builds a contract around a user-supplied inner JSON Schema.
///
/// The verdict is validated against the schema locally in both modes. JSON Schema mode sends it through the
/// provider's strict wrapper; JSON Object mode appends it to the prompt, because the
/// provider then guarantees only that the reply is JSON.
pub(crate) fn from_inner_schema(
prompt_template: &str,
schema: Value,
response_format_type: ClassifierResponseFormat,
) -> Result<Self> {
validate_prompt(prompt_template)?;
if schema.get("json_schema").is_some() {
return Err(LibsyError::AlgorithmError {
message:
Expand All @@ -121,32 +126,53 @@ impl ClassifierContract {
});
}
let validator = compile_schema(&schema)?;
Self::from_response_format(
prompt_template,
json!({
"type": "json_schema",
"json_schema": {
"name": "switchyard_classifier_response",
"strict": true,
"schema": schema,
match response_format_type {
ClassifierResponseFormat::JsonSchema => Ok(Self::from_response_format(
prompt_template.to_string(),
json!({
"type": "json_schema",
"json_schema": {
"name": "switchyard_classifier_response",
"strict": true,
"schema": schema,
}
}),
Some(validator),
)),
ClassifierResponseFormat::JsonObject => {
// Catches the common case of a root `type` that excludes objects: the provider
// only returns objects, so every verdict would fail and route to the default.
// A schema with no `type` (e.g. `{"enum": [1, 2]}`) is not checked and can
// still never match; that is left to the author.
let is_object_allowed = match schema.get("type") {
Some(Value::String(root)) => root == "object",
Some(Value::Array(roots)) => roots.iter().any(|root| root == "object"),
_ => true,
};
if !is_object_allowed {
return Err(algorithm_error(
"response_schema must describe a JSON object in json_object mode",
));
}
}),
Some(validator),
)
Ok(Self::from_response_format(
schema_in_prompt(prompt_template, &schema)?,
json!({"type": "json_object"}),
Some(validator),
))
}
}
}

fn from_response_format(
prompt_template: &str,
system_prompt: String,
response_format: Value,
validator: Option<Validator>,
) -> Result<Self> {
validate_prompt(prompt_template)?;

Ok(Self {
system_prompt: prompt_template.to_string(),
) -> Self {
Self {
system_prompt,
response_format,
validator,
})
}
}

pub(crate) fn system_prompt(&self) -> &str {
Expand Down Expand Up @@ -187,6 +213,19 @@ pub(super) fn validate_prompt(prompt_template: &str) -> Result<()> {
Ok(())
}

/// Appends the verdict schema to a prompt for JSON Object mode.
///
/// Callers validate the template, not this result: the appended schema makes the prompt
/// non-empty and may itself contain `{{RESPONSE_SCHEMA}}` text.
fn schema_in_prompt(prompt_template: &str, schema: &Value) -> Result<String> {
let rendered_schema = serde_json::to_string_pretty(schema).map_err(|error| {
algorithm_error(format!("response schema could not be rendered: {error}"))
})?;
Ok(format!(
"{prompt_template}\n\nReturn exactly one JSON object matching this JSON Schema:\n{rendered_schema}"
))
}

fn compile_schema(schema: &Value) -> Result<Validator> {
if !schema.is_object() {
return Err(algorithm_error("response_schema must be a JSON object"));
Expand Down Expand Up @@ -307,6 +346,7 @@ mod tests {
"required": ["decision"],
"additionalProperties": false
}),
ClassifierResponseFormat::JsonSchema,
)?;

assert_eq!(
Expand Down Expand Up @@ -335,12 +375,110 @@ mod tests {

#[test]
fn a_provider_wrapper_is_rejected_as_an_inner_schema() {
for response_format_type in [
ClassifierResponseFormat::JsonSchema,
ClassifierResponseFormat::JsonObject,
] {
let error = ClassifierContract::from_inner_schema(
"classify",
json!({"json_schema": {"schema": {"type": "object"}}}),
response_format_type,
)
.expect_err("a wrapped schema should error");

assert!(
error.to_string().contains("inner JSON Schema"),
"{response_format_type:?}: {error}"
);
}
}

#[test]
fn a_custom_contract_can_request_a_json_object() -> Result<()> {
let contract = ClassifierContract::from_inner_schema(
"Choose a target.",
json!({
"type": "object",
"properties": {"target": {"type": "string", "enum": ["sonnet", "opus"]}},
"required": ["target"],
"additionalProperties": false
}),
ClassifierResponseFormat::JsonObject,
)?;

assert_eq!(contract.response_format(), &json!({"type": "json_object"}));
assert!(contract.system_prompt().starts_with("Choose a target."));
assert!(contract.system_prompt().contains("JSON Schema"));
assert!(contract.system_prompt().contains("\"target\""));
// The provider no longer enforces the schema, so the local validator must.
contract.validate_verdict(&json!({"target": "sonnet"}))?;
assert!(
contract
.validate_verdict(&json!({"target": "unknown"}))
.is_err()
);
Ok(())
}

#[test]
fn json_object_mode_accepts_only_schemas_whose_root_allows_an_object() {
for (schema, accepted) in [
(json!({"type": "object"}), true),
(json!({"type": ["object"]}), true),
(json!({"type": ["object", "null"]}), true),
(json!({"type": "string"}), false),
(json!({"type": "array"}), false),
(json!({"type": ["string", "null"]}), false),
] {
let result = ClassifierContract::from_inner_schema(
"classify",
schema.clone(),
ClassifierResponseFormat::JsonObject,
);

match result {
Ok(_) => assert!(accepted, "{schema} should be rejected"),
Err(error) => {
assert!(!accepted, "{schema}: {error}");
assert!(
error.to_string().contains("must describe a JSON object"),
"{schema}: {error}"
);
}
}
}
}

#[test]
fn a_schema_that_mentions_the_placeholder_is_accepted_in_both_modes() {
let schema = json!({
"type": "object",
"properties": {"target": {"type": "string"}},
"required": ["target"],
"description": "{{RESPONSE_SCHEMA}}"
});
for response_format_type in [
ClassifierResponseFormat::JsonSchema,
ClassifierResponseFormat::JsonObject,
] {
ClassifierContract::from_inner_schema("pick", schema.clone(), response_format_type)
.unwrap_or_else(|error| panic!("{response_format_type:?}: {error}"));
}
}

#[test]
fn an_empty_prompt_is_rejected_in_json_object_mode() {
let error = ClassifierContract::from_inner_schema(
"classify",
json!({"json_schema": {"schema": {"type": "object"}}}),
" ",
json!({"type": "object"}),
ClassifierResponseFormat::JsonObject,
)
.expect_err("provider wrapper should be rejected");
.expect_err("an empty prompt should be rejected before the schema is appended");

assert!(error.to_string().contains("inner JSON Schema"));
assert!(
error
.to_string()
.contains("classifier prompt must not be empty")
);
}
}
26 changes: 15 additions & 11 deletions crates/switchyard-py/src/libsy_bindings.rs
Original file line number Diff line number Diff line change
Expand Up @@ -218,7 +218,8 @@ impl PyCustomClassifierConfig {
session_affinity=false,
message_hash_fallback=false,
recent_turn_window=None,
max_output_tokens=4096
max_output_tokens=4096,
response_format_type="json_schema"
))]
#[allow(clippy::too_many_arguments)]
fn new(
Expand All @@ -229,6 +230,7 @@ impl PyCustomClassifierConfig {
message_hash_fallback: bool,
recent_turn_window: Option<usize>,
max_output_tokens: u64,
response_format_type: &str,
) -> PyResult<Self> {
// Convert the Python schema into serde JSON and pair it with the target-selector policy;
// conversion failures propagate to Python through `PyResult`.
Expand All @@ -241,6 +243,7 @@ impl PyCustomClassifierConfig {
inner.message_hash_fallback = message_hash_fallback;
inner.recent_turn_window = recent_turn_window;
inner.max_output_tokens = max_output_tokens;
inner.response_format_type = parse_response_format_type(response_format_type)?;
Ok(Self { inner })
}
}
Expand Down Expand Up @@ -348,16 +351,17 @@ fn classifier_contract(
if let Some(prompt) = prompt {
contract = contract.with_prompt(prompt);
}
let response_format_type = match response_format_type {
"json_schema" => ClassifierResponseFormat::JsonSchema,
"json_object" => ClassifierResponseFormat::JsonObject,
other => {
return Err(PyValueError::new_err(format!(
"response_format_type must be 'json_schema' or 'json_object', got {other:?}"
)));
}
};
Ok(contract.with_response_format_type(response_format_type))
Ok(contract.with_response_format_type(parse_response_format_type(response_format_type)?))
}

fn parse_response_format_type(response_format_type: &str) -> PyResult<ClassifierResponseFormat> {
match response_format_type {
"json_schema" => Ok(ClassifierResponseFormat::JsonSchema),
"json_object" => Ok(ClassifierResponseFormat::JsonObject),
other => Err(PyValueError::new_err(format!(
"response_format_type must be 'json_schema' or 'json_object', got {other:?}"
))),
}
}

/// Judge target and policy used when stage-router signals are inconclusive.
Expand Down
Loading
Loading