1use std::collections::HashMap;
39use std::fmt;
40use std::sync::{Arc, Mutex, PoisonError};
41
42use serde::de::DeserializeOwned;
43use serde::{Deserialize, Serialize};
44
45use crate::response::ModelResponse;
46
47pub const MAX_FIELD_LEN: usize = 64;
49
50#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, thiserror::Error)]
55#[serde(tag = "kind", rename_all = "snake_case")]
56#[non_exhaustive]
57pub enum StructuredOutputError {
58 #[error("model produced no complete output")]
61 NoOutput,
62 #[error("model output is not JSON: {detail}")]
64 NotJson {
65 detail: String,
68 },
69 #[error("schema violation at {pointer}: {keyword}")]
71 SchemaViolation {
72 pointer: String,
74 keyword: String,
76 },
77 #[error("unknown field {field} at {pointer}")]
80 UnknownField {
81 pointer: String,
83 field: String,
85 },
86 #[error("missing field {field} at {pointer}")]
88 MissingField {
89 pointer: String,
91 field: String,
93 },
94 #[error("model produced {candidates} candidate documents, expected one")]
96 MultipleCandidates {
97 candidates: usize,
99 },
100 #[error("model refused to answer")]
102 Refusal,
103}
104
105impl StructuredOutputError {
106 #[must_use]
108 pub fn not_json(error: serde_json::Error) -> Self {
109 Self::NotJson {
110 detail: error.to_string(),
111 }
112 }
113
114 #[must_use]
116 pub const fn as_str(&self) -> &'static str {
117 match self {
118 Self::NoOutput => "no_output",
119 Self::NotJson { .. } => "not_json",
120 Self::SchemaViolation { .. } => "schema_violation",
121 Self::UnknownField { .. } => "unknown_field",
122 Self::MissingField { .. } => "missing_field",
123 Self::MultipleCandidates { .. } => "multiple_candidates",
124 Self::Refusal => "refusal",
125 }
126 }
127}
128
129impl From<StructuredOutputError> for crate::error::ProviderError {
130 fn from(value: StructuredOutputError) -> Self {
135 match value {
136 StructuredOutputError::Refusal => Self::refusal(),
137 other => Self::malformed(other.as_str()),
138 }
139 }
140}
141
142fn sanitize(raw: &str) -> String {
144 let mut out = String::with_capacity(raw.len().min(MAX_FIELD_LEN));
145 for ch in raw.chars() {
146 if out.len() >= MAX_FIELD_LEN {
147 break;
148 }
149 if ch.is_ascii_alphanumeric() || matches!(ch, '_' | '-' | '.' | '/' | '[' | ']') {
150 out.push(ch);
151 } else {
152 out.push('?');
153 }
154 }
155 out
156}
157
158const ROOT_POINTER: &str = "/";
160
161fn pointer_or_root(location: &str) -> String {
162 if location.is_empty() {
163 ROOT_POINTER.to_owned()
164 } else {
165 sanitize(location)
166 }
167}
168
169#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
174#[error("invalid JSON Schema at {pointer}: {keyword}")]
175pub struct SchemaCompileError {
176 pub pointer: String,
178 pub keyword: String,
180}
181
182#[derive(Clone)]
186pub struct CompiledSchema {
187 schema: Arc<serde_json::Value>,
188 validator: Arc<jsonschema::Validator>,
189 fingerprint: turnframe_core::hash::Digest,
190}
191
192impl CompiledSchema {
193 pub fn compile(schema: &serde_json::Value) -> Result<Self, SchemaCompileError> {
214 let validator = jsonschema::validator_for(schema).map_err(|error| SchemaCompileError {
215 pointer: pointer_or_root(&error.schema_path().to_string()),
216 keyword: sanitize(keyword_of(error.kind())),
217 })?;
218 let fingerprint =
219 turnframe_core::hash::Digest::of_canonical(schema).map_err(|_| SchemaCompileError {
220 pointer: ROOT_POINTER.to_owned(),
221 keyword: "not_serializable".to_owned(),
222 })?;
223 Ok(Self {
224 schema: Arc::new(schema.clone()),
225 validator: Arc::new(validator),
226 fingerprint,
227 })
228 }
229
230 #[must_use]
232 pub fn schema(&self) -> &serde_json::Value {
233 &self.schema
234 }
235
236 #[must_use]
239 pub fn fingerprint(&self) -> &turnframe_core::hash::Digest {
240 &self.fingerprint
241 }
242
243 pub fn validate(&self, instance: &serde_json::Value) -> Result<(), StructuredOutputError> {
254 let mut unknown = None;
255 let mut missing = None;
256 let mut other = None;
257 for error in self.validator.iter_errors(instance) {
258 let pointer = pointer_or_root(&error.instance_path().to_string());
259 match error.kind() {
260 jsonschema::error::ValidationErrorKind::AdditionalProperties { unexpected }
261 | jsonschema::error::ValidationErrorKind::UnevaluatedProperties { unexpected } => {
262 if unknown.is_none() {
263 let field = unexpected
264 .first()
265 .map_or_else(String::new, |name| sanitize(name));
266 unknown = Some(StructuredOutputError::UnknownField { pointer, field });
267 }
268 }
269 jsonschema::error::ValidationErrorKind::Required { property } => {
270 if missing.is_none() {
271 let field = property.as_str().map_or_else(String::new, sanitize);
272 missing = Some(StructuredOutputError::MissingField { pointer, field });
273 }
274 }
275 kind => {
276 if other.is_none() {
277 other = Some(StructuredOutputError::SchemaViolation {
278 pointer,
279 keyword: sanitize(keyword_of(kind)),
280 });
281 }
282 }
283 }
284 if unknown.is_some() && missing.is_some() && other.is_some() {
285 break;
286 }
287 }
288 match unknown.or(missing).or(other) {
289 Some(error) => Err(error),
290 None => Ok(()),
291 }
292 }
293}
294
295impl fmt::Debug for CompiledSchema {
296 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
297 f.debug_struct("CompiledSchema")
298 .field("fingerprint", &self.fingerprint.as_str())
299 .finish_non_exhaustive()
300 }
301}
302
303impl PartialEq for CompiledSchema {
304 fn eq(&self, other: &Self) -> bool {
305 self.fingerprint == other.fingerprint
306 }
307}
308
309impl Eq for CompiledSchema {}
310
311fn keyword_of(kind: &jsonschema::error::ValidationErrorKind) -> &'static str {
313 use jsonschema::error::ValidationErrorKind as K;
314 match kind {
315 K::AdditionalItems { .. } => "additionalItems",
316 K::AdditionalProperties { .. } => "additionalProperties",
317 K::AnyOf { .. } => "anyOf",
318 K::BacktrackLimitExceeded { .. } | K::RegexEngineFailure { .. } => "pattern",
319 K::Constant { .. } => "const",
320 K::Contains => "contains",
321 K::ContentEncoding { .. } => "contentEncoding",
322 K::ContentMediaType { .. } => "contentMediaType",
323 K::Custom { .. } => "custom",
324 K::Enum { .. } => "enum",
325 K::ExclusiveMaximum { .. } => "exclusiveMaximum",
326 K::ExclusiveMinimum { .. } => "exclusiveMinimum",
327 K::FalseSchema => "false",
328 K::Format { .. } => "format",
329 K::MaxItems { .. } => "maxItems",
330 K::Maximum { .. } => "maximum",
331 K::MaxLength { .. } => "maxLength",
332 K::MaxProperties { .. } => "maxProperties",
333 K::MinItems { .. } => "minItems",
334 K::Minimum { .. } => "minimum",
335 K::MinLength { .. } => "minLength",
336 K::MinProperties { .. } => "minProperties",
337 K::MultipleOf { .. } => "multipleOf",
338 K::Not { .. } => "not",
339 K::OneOfMultipleValid { .. } | K::OneOfNotValid { .. } => "oneOf",
340 K::Pattern { .. } => "pattern",
341 K::PropertyNames { .. } => "propertyNames",
342 K::Required { .. } => "required",
343 K::Type { .. } => "type",
344 K::UnevaluatedItems { .. } => "unevaluatedItems",
345 K::UnevaluatedProperties { .. } => "unevaluatedProperties",
346 K::UniqueItems => "uniqueItems",
347 _ => "schema",
348 }
349}
350
351#[derive(Clone, Default)]
357pub struct SchemaCache {
358 entries: Arc<Mutex<HashMap<String, CompiledSchema>>>,
359}
360
361impl SchemaCache {
362 #[must_use]
364 pub fn new() -> Self {
365 Self::default()
366 }
367
368 pub fn compile(
388 &self,
389 schema: &serde_json::Value,
390 ) -> Result<CompiledSchema, SchemaCompileError> {
391 let key = turnframe_core::hash::Digest::of_canonical(schema)
392 .map_err(|_| SchemaCompileError {
393 pointer: ROOT_POINTER.to_owned(),
394 keyword: "not_serializable".to_owned(),
395 })?
396 .into();
397 {
398 let entries = self.entries.lock().unwrap_or_else(PoisonError::into_inner);
399 if let Some(found) = entries.get(&key) {
400 return Ok(found.clone());
401 }
402 }
403 let compiled = CompiledSchema::compile(schema)?;
404 let mut entries = self.entries.lock().unwrap_or_else(PoisonError::into_inner);
405 Ok(entries.entry(key).or_insert(compiled).clone())
406 }
407
408 #[must_use]
410 pub fn len(&self) -> usize {
411 self.entries
412 .lock()
413 .unwrap_or_else(PoisonError::into_inner)
414 .len()
415 }
416
417 #[must_use]
419 pub fn is_empty(&self) -> bool {
420 self.len() == 0
421 }
422
423 pub fn clear(&self) {
425 self.entries
426 .lock()
427 .unwrap_or_else(PoisonError::into_inner)
428 .clear();
429 }
430}
431
432impl fmt::Debug for SchemaCache {
433 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
434 f.debug_struct("SchemaCache")
435 .field("len", &self.len())
436 .finish()
437 }
438}
439
440pub fn parse_structured<T: DeserializeOwned>(
473 response: &ModelResponse,
474 schema: &CompiledSchema,
475) -> Result<T, StructuredOutputError> {
476 let value = response.single_json()?;
477 parse_structured_value(&value, schema)
478}
479
480pub fn parse_structured_value<T: DeserializeOwned>(
488 value: &serde_json::Value,
489 schema: &CompiledSchema,
490) -> Result<T, StructuredOutputError> {
491 schema.validate(value)?;
492 serde_json::from_value(value.clone()).map_err(classify_serde_error)
493}
494
495fn classify_serde_error(error: serde_json::Error) -> StructuredOutputError {
501 let message = error.to_string();
502 if let Some(field) = quoted_name(&message, "unknown field ") {
503 return StructuredOutputError::UnknownField {
504 pointer: ROOT_POINTER.to_owned(),
505 field,
506 };
507 }
508 if let Some(field) = quoted_name(&message, "missing field ") {
509 return StructuredOutputError::MissingField {
510 pointer: ROOT_POINTER.to_owned(),
511 field,
512 };
513 }
514 StructuredOutputError::SchemaViolation {
515 pointer: ROOT_POINTER.to_owned(),
516 keyword: "type".to_owned(),
517 }
518}
519
520fn quoted_name(message: &str, prefix: &str) -> Option<String> {
523 let rest = message.strip_prefix(prefix)?;
524 let inner = rest.strip_prefix('`')?;
525 let end = inner.find('`')?;
526 Some(sanitize(&inner[..end]))
527}
528
529#[cfg(test)]
530mod tests {
531 use super::*;
532 use crate::ids::RequestId;
533 use crate::request::ToolCall;
534 use crate::response::{FinishReason, ModelResponse};
535 use serde_json::json;
536
537 #[derive(Debug, Deserialize, PartialEq)]
538 #[serde(deny_unknown_fields)]
539 struct Act {
540 operation: String,
541 target: String,
542 }
543
544 #[derive(Debug, Deserialize, PartialEq)]
545 #[serde(deny_unknown_fields)]
546 struct Plan {
547 acts: Vec<Act>,
548 }
549
550 fn act_schema() -> serde_json::Value {
551 json!({
552 "type": "object",
553 "properties": {
554 "operation": {"type": "string"},
555 "target": {"type": "string"}
556 },
557 "required": ["operation", "target"],
558 "additionalProperties": false
559 })
560 }
561
562 fn plan_schema() -> CompiledSchema {
563 CompiledSchema::compile(&json!({
564 "type": "object",
565 "properties": {"acts": {"type": "array", "items": act_schema()}},
566 "required": ["acts"],
567 "additionalProperties": false
568 }))
569 .unwrap()
570 }
571
572 fn text_response(text: &str) -> ModelResponse {
573 ModelResponse::new(RequestId::nil(), "p", "m").with_text(text)
574 }
575
576 #[test]
577 fn a_valid_two_act_plan_parses() {
578 let response = text_response(
579 r#"{"acts": [
580 {"operation": "set_travel_date", "target": "tok_1"},
581 {"operation": "set_amount", "target": "tok_2"}
582 ]}"#,
583 );
584 let plan: Plan = parse_structured(&response, &plan_schema()).unwrap();
585 assert_eq!(plan.acts.len(), 2);
586 assert_eq!(plan.acts[1].operation, "set_amount");
587 }
588
589 #[test]
590 fn one_malformed_act_rejects_the_whole_response() {
591 let response = text_response(
593 r#"{"acts": [
594 {"operation": "set_travel_date", "target": "tok_1"},
595 {"operation": "set_amount"}
596 ]}"#,
597 );
598 let error = parse_structured::<Plan>(&response, &plan_schema()).unwrap_err();
599 assert!(
600 matches!(
601 &error,
602 StructuredOutputError::MissingField { pointer, field }
603 if pointer == "/acts/1" && field == "target"
604 ),
605 "{error:?}"
606 );
607 }
608
609 #[test]
610 fn one_act_with_an_unknown_field_rejects_the_whole_response() {
611 let response = text_response(
612 r#"{"acts": [
613 {"operation": "set_travel_date", "target": "tok_1"},
614 {"operation": "set_amount", "target": "tok_2", "force": true}
615 ]}"#,
616 );
617 let error = parse_structured::<Plan>(&response, &plan_schema()).unwrap_err();
618 assert!(
619 matches!(
620 &error,
621 StructuredOutputError::UnknownField { pointer, field }
622 if pointer == "/acts/1" && field == "force"
623 ),
624 "{error:?}"
625 );
626 }
627
628 #[test]
629 fn unknown_field_wins_over_missing_field_deterministically() {
630 let response = text_response(r#"{"acts": [{"operation": "x", "force": true}]}"#);
632 let schema = plan_schema();
633 let first = parse_structured::<Plan>(&response, &schema).unwrap_err();
634 let second = parse_structured::<Plan>(&response, &schema).unwrap_err();
635 assert_eq!(first, second);
636 assert!(matches!(first, StructuredOutputError::UnknownField { .. }));
637 }
638
639 #[test]
640 fn not_json_is_reported_without_the_payload() {
641 let response = text_response("Certo! Ecco il piano: primo, secondo.");
642 let error = parse_structured::<Plan>(&response, &plan_schema()).unwrap_err();
643 let StructuredOutputError::NotJson { detail } = &error else {
644 panic!("{error:?}");
645 };
646 assert!(detail.contains("expected"), "{detail}");
647 assert!(!detail.contains("Certo"), "the payload leaked: {detail}");
648 assert_eq!(error.as_str(), "not_json");
649 }
650
651 #[test]
652 fn a_schema_violation_names_the_pointer_and_keyword_only() {
653 let response = text_response(r#"{"acts": [{"operation": 7, "target": "t"}]}"#);
654 let error = parse_structured::<Plan>(&response, &plan_schema()).unwrap_err();
655 assert!(
656 matches!(
657 &error,
658 StructuredOutputError::SchemaViolation { pointer, keyword }
659 if pointer == "/acts/0/operation" && keyword == "type"
660 ),
661 "{error:?}"
662 );
663 assert!(!error.to_string().contains('7'));
664 }
665
666 #[test]
667 fn no_output_empty_and_refusal_are_distinct() {
668 let empty = ModelResponse::new(RequestId::nil(), "p", "m");
669 assert_eq!(
670 parse_structured::<Plan>(&empty, &plan_schema()).unwrap_err(),
671 StructuredOutputError::NoOutput
672 );
673
674 let refused = text_response("no").with_finish(FinishReason::Refusal);
675 assert_eq!(
676 parse_structured::<Plan>(&refused, &plan_schema()).unwrap_err(),
677 StructuredOutputError::Refusal
678 );
679 }
680
681 #[test]
682 fn two_tool_calls_are_multiple_candidates() {
683 let response = ModelResponse::new(RequestId::nil(), "p", "m")
684 .with_tool_call(ToolCall::new("a", "plan", json!({"acts": []})))
685 .with_tool_call(ToolCall::new("b", "plan", json!({"acts": []})))
686 .with_finish(FinishReason::ToolCalls);
687 assert_eq!(
688 parse_structured::<Plan>(&response, &plan_schema()).unwrap_err(),
689 StructuredOutputError::MultipleCandidates { candidates: 2 }
690 );
691 }
692
693 #[test]
694 fn serde_catches_what_a_loose_schema_lets_through() {
695 let loose = CompiledSchema::compile(&json!({"type": "object"})).unwrap();
697 let response = text_response(r#"{"acts": [], "extra": 1}"#);
698 let error = parse_structured::<Plan>(&response, &loose).unwrap_err();
699 assert!(
700 matches!(&error, StructuredOutputError::UnknownField { field, .. } if field == "extra"),
701 "{error:?}"
702 );
703
704 let missing = text_response(r#"{}"#);
705 let error = parse_structured::<Plan>(&missing, &loose).unwrap_err();
706 assert!(
707 matches!(&error, StructuredOutputError::MissingField { field, .. } if field == "acts"),
708 "{error:?}"
709 );
710
711 let wrong_type = text_response(r#"{"acts": "no"}"#);
712 let error = parse_structured::<Plan>(&wrong_type, &loose).unwrap_err();
713 assert!(
714 matches!(error, StructuredOutputError::SchemaViolation { .. }),
715 "{error:?}"
716 );
717 }
718
719 #[test]
720 fn field_names_from_the_model_are_sanitized() {
721 let loose = CompiledSchema::compile(&json!({"type": "object"})).unwrap();
722 let response = text_response(r#"{"acts": [], "a field with spaces": 1}"#);
723 let error = parse_structured::<Plan>(&response, &loose).unwrap_err();
724 let StructuredOutputError::UnknownField { field, .. } = &error else {
725 panic!("{error:?}");
726 };
727 assert!(!field.contains(' '), "{field}");
728 assert_eq!(field, "a?field?with?spaces");
729 }
730
731 #[test]
732 fn the_cache_compiles_once_and_rejects_bad_schemas() {
733 let cache = SchemaCache::new();
734 assert!(cache.is_empty());
735 let schema = json!({"type": "object", "properties": {"a": {"type": "string"}}});
736 let first = cache.compile(&schema).unwrap();
737 let second = cache.compile(&schema).unwrap();
738 assert_eq!(first.fingerprint(), second.fingerprint());
739 assert_eq!(first, second);
740 assert_eq!(cache.len(), 1);
741
742 let reordered = json!({"properties": {"a": {"type": "string"}}, "type": "object"});
744 let third = cache.compile(&reordered).unwrap();
745 assert_eq!(third.fingerprint(), first.fingerprint());
746 assert_eq!(cache.len(), 1);
747
748 let invalid = cache.compile(&json!({"type": "not-a-type"}));
749 assert!(invalid.is_err());
750 assert_eq!(cache.len(), 1, "invalid schemas are not cached");
751
752 cache.clear();
753 assert!(cache.is_empty());
754 assert!(format!("{cache:?}").contains("SchemaCache"));
755 }
756
757 #[test]
758 fn structured_errors_map_onto_provider_errors() {
759 use crate::error::{ProviderErrorKind, RetryClass};
760 let malformed = crate::error::ProviderError::from(StructuredOutputError::NoOutput);
761 assert!(matches!(malformed.kind(), ProviderErrorKind::Malformed));
762 assert_eq!(malformed.retry_class(), RetryClass::Retry);
763 assert_eq!(
764 malformed.code().map(|c| c.as_str().to_owned()),
765 Some("no_output".to_owned())
766 );
767
768 let refusal = crate::error::ProviderError::from(StructuredOutputError::Refusal);
769 assert!(matches!(refusal.kind(), ProviderErrorKind::Refusal));
770 assert_eq!(refusal.retry_class(), RetryClass::Fatal);
771 }
772}