use crate::message::Message;
pub fn validate_structured(schema: &serde_json::Value, answer: &str) -> Result<(), String> {
let instance: serde_json::Value =
serde_json::from_str(answer).map_err(|e| format!("answer is not valid JSON: {e}"))?;
let validator =
jsonschema::validator_for(schema).map_err(|e| format!("invalid JSON schema: {e}"))?;
if let Err(e) = validator.validate(&instance) {
return Err(format!("answer does not match the JSON schema: {e}"));
}
Ok(())
}
pub fn structured_retry_message(error: &str) -> Message {
Message::user(format!(
"your previous answer failed JSON schema validation: {error}; \
please reply with a single JSON value conforming to the schema"
))
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum StructuredOutcome {
Passed,
Retry {
message: Message,
},
Exhausted {
max_retries: usize,
},
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StructuredValidator {
schema: serde_json::Value,
max_retries: usize,
retries_used: usize,
}
impl StructuredValidator {
pub fn new(schema: serde_json::Value, max_retries: usize) -> Self {
Self {
schema,
max_retries,
retries_used: 0,
}
}
pub fn validate(&mut self, answer: &str) -> StructuredOutcome {
match validate_structured(&self.schema, answer) {
Ok(()) => StructuredOutcome::Passed,
Err(error) => {
self.retries_used += 1;
if self.retries_used > self.max_retries {
StructuredOutcome::Exhausted {
max_retries: self.max_retries,
}
} else {
StructuredOutcome::Retry {
message: structured_retry_message(&error),
}
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::message::ContentBlock;
fn schema() -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": { "city": { "type": "string" } },
"required": ["city"],
})
}
#[test]
fn validator_three_outcomes() {
let mut validator = StructuredValidator::new(schema(), 1);
match validator.validate("bad") {
StructuredOutcome::Retry { message } => {
let Message::User(blocks) = message else {
panic!("retry message must be a user message")
};
assert!(blocks.iter().any(|b| matches!(
b,
ContentBlock::Text(t) if t.contains("JSON schema validation")
)));
}
other => panic!("expected Retry, got {other:?}"),
}
assert!(matches!(
validator.validate("bad2"),
StructuredOutcome::Exhausted { max_retries: 1 }
));
}
#[test]
fn validator_passes_after_retry() {
let mut validator = StructuredValidator::new(schema(), 3);
assert!(matches!(
validator.validate("bad"),
StructuredOutcome::Retry { .. }
));
assert!(matches!(
validator.validate(r#"{"city":"Beijing"}"#),
StructuredOutcome::Passed
));
assert!(matches!(
validator.validate(r#"{"city":"Shanghai"}"#),
StructuredOutcome::Passed
));
}
}