use serde::{Deserialize, Serialize};
use crate::{
error::{PredictError, Result},
field::{FieldDef, FieldKind, FieldType, OneOfDiscriminator, VariantArm},
};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Signature {
instructions: String,
fields: Vec<FieldDef>,
}
impl Signature {
pub fn builder(instructions: impl Into<String>) -> SignatureBuilder {
SignatureBuilder {
instructions: instructions.into(),
fields: Vec::new(),
}
}
#[must_use]
pub fn instructions(&self) -> &str {
&self.instructions
}
pub fn input_fields(&self) -> impl Iterator<Item = &FieldDef> {
self.fields.iter().filter(|f| f.kind == FieldKind::Input)
}
pub fn output_fields(&self) -> impl Iterator<Item = &FieldDef> {
self.fields.iter().filter(|f| f.kind == FieldKind::Output)
}
#[must_use]
pub fn fields(&self) -> &[FieldDef] {
&self.fields
}
pub fn dump_state(&self) -> Result<String> {
serde_json::to_string(self).map_err(Into::into)
}
pub fn set_instructions(&mut self, instructions: impl Into<String>) {
self.instructions = instructions.into();
}
#[must_use]
pub fn with_prepended_output(&self, field: FieldDef) -> Self {
let mut new_fields = Vec::with_capacity(self.fields.len() + 1);
let mut inserted = false;
for f in &self.fields {
if f.kind == FieldKind::Output && !inserted {
new_fields.push(FieldDef {
kind: FieldKind::Output,
..field.clone()
});
inserted = true;
}
new_fields.push(f.clone());
}
if !inserted {
new_fields.push(FieldDef {
kind: FieldKind::Output,
..field
});
}
Self {
instructions: self.instructions.clone(),
fields: new_fields,
}
}
pub fn load_state(json: &str) -> Result<Self> {
let sig: Self = serde_json::from_str(json)?;
if sig.input_fields().next().is_none() {
return Err(PredictError::invalid_signature(
"signature must have at least one input field",
));
}
if sig.output_fields().next().is_none() {
return Err(PredictError::invalid_signature(
"signature must have at least one output field",
));
}
Ok(sig)
}
}
pub struct SignatureBuilder {
instructions: String,
fields: Vec<FieldDef>,
}
impl SignatureBuilder {
#[must_use]
pub fn input(mut self, mut def: FieldDef) -> Self {
def.kind = FieldKind::Input;
self.fields.push(def);
self
}
#[must_use]
pub fn output(mut self, mut def: FieldDef) -> Self {
def.kind = FieldKind::Output;
self.fields.push(def);
self
}
pub fn build(self) -> Result<Signature> {
let has_input = self.fields.iter().any(|f| f.kind == FieldKind::Input);
let has_output = self.fields.iter().any(|f| f.kind == FieldKind::Output);
if !has_input {
return Err(PredictError::invalid_signature(
"signature must have at least one input field",
));
}
if !has_output {
return Err(PredictError::invalid_signature(
"signature must have at least one output field",
));
}
for field in &self.fields {
validate_variant_shapes(&field.name, &field.field_type)?;
}
Ok(Signature {
instructions: self.instructions,
fields: self.fields,
})
}
}
fn validate_variant_shapes(path: &str, ft: &FieldType) -> Result<()> {
match ft {
FieldType::OneOf {
arms,
discriminator,
} => {
validate_variant_arms(path, arms, "OneOf")?;
if let Some(disc) = discriminator {
validate_one_of_discriminator(path, arms, disc)?;
}
for (i, arm) in arms.iter().enumerate() {
let arm_path = format!("{path}#arm{i}");
validate_variant_shapes(&arm_path, &arm.field_type)?;
}
Ok(())
}
FieldType::AnyOf { arms } => {
validate_variant_arms(path, arms, "AnyOf")?;
for (i, arm) in arms.iter().enumerate() {
let arm_path = format!("{path}#arm{i}");
validate_variant_shapes(&arm_path, &arm.field_type)?;
}
Ok(())
}
FieldType::List(inner) | FieldType::Nullable(inner) | FieldType::Map(inner) => {
validate_variant_shapes(path, inner)
}
FieldType::Object(fields) => {
for f in fields {
let nested = format!("{path}.{}", f.name);
validate_variant_shapes(&nested, &f.field_type)?;
}
Ok(())
}
FieldType::String
| FieldType::Int
| FieldType::Float
| FieldType::Bool
| FieldType::Enum(_)
| FieldType::Media { .. } => Ok(()),
}
}
fn validate_variant_arms(path: &str, arms: &[VariantArm], kind: &str) -> Result<()> {
if arms.is_empty() {
return Err(PredictError::invalid_signature(format!(
"{kind} field at `{path}` must declare at least one arm"
)));
}
Ok(())
}
fn validate_one_of_discriminator(
path: &str,
arms: &[VariantArm],
discriminator: &OneOfDiscriminator,
) -> Result<()> {
if discriminator.tags.len() != arms.len() {
return Err(PredictError::invalid_signature(format!(
"OneOf at `{path}` discriminator has {} tags but {} arms; the \
vectors must be parallel",
discriminator.tags.len(),
arms.len()
)));
}
let mut seen = std::collections::HashSet::with_capacity(discriminator.tags.len());
for tag in &discriminator.tags {
if !seen.insert(tag.as_str()) {
return Err(PredictError::invalid_signature(format!(
"OneOf at `{path}` discriminator tag `{tag}` appears more than \
once; tags must be unique"
)));
}
}
for (i, arm) in arms.iter().enumerate() {
let arm_path = format!("{path}#{tag}", tag = discriminator.tags[i]);
validate_arm_carries_discriminator(
&arm_path,
&arm.field_type,
&discriminator.property,
&discriminator.tags[i],
)?;
}
Ok(())
}
fn validate_arm_carries_discriminator(
arm_path: &str,
arm_type: &FieldType,
property: &str,
tag: &str,
) -> Result<()> {
let inner = match arm_type {
FieldType::Nullable(inner) => inner.as_ref(),
other => other,
};
let FieldType::Object(fields) = inner else {
return Err(PredictError::invalid_signature(format!(
"tagged OneOf arm at `{arm_path}` must be an Object (or \
Nullable<Object>); got `{}`",
arm_type.type_label()
)));
};
let prop = fields.iter().find(|f| f.name == property).ok_or_else(|| {
PredictError::invalid_signature(format!(
"tagged OneOf arm at `{arm_path}` is missing the discriminator \
property `{property}`"
))
})?;
match &prop.field_type {
FieldType::Enum(variants) if variants.iter().any(|v| v == tag) => Ok(()),
other => Err(PredictError::invalid_signature(format!(
"tagged OneOf arm at `{arm_path}` declares discriminator \
property `{property}` as `{}`, but the discriminator's tag \
`{tag}` requires an Enum variant containing it (typical: \
single-element Enum from a const-restricted schema)",
other.type_label()
))),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::field::FieldType;
fn simple_signature() -> Signature {
Signature::builder("Answer the question.")
.input(FieldDef::input(
"question",
FieldType::String,
"The user question",
))
.output(FieldDef::output("answer", FieldType::String, "The answer"))
.build()
.unwrap()
}
#[test]
fn builder_produces_valid_signature() {
let sig = simple_signature();
assert_eq!(sig.instructions(), "Answer the question.");
assert_eq!(sig.input_fields().count(), 1);
assert_eq!(sig.output_fields().count(), 1);
}
#[test]
fn builder_no_inputs_fails() {
let result = Signature::builder("test")
.output(FieldDef::output("answer", FieldType::String, "answer"))
.build();
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(err.contains("input"));
}
#[test]
fn builder_no_outputs_fails() {
let result = Signature::builder("test")
.input(FieldDef::input("question", FieldType::String, "question"))
.build();
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(err.contains("output"));
}
#[test]
fn field_ordering_preserved() {
let sig = Signature::builder("test")
.input(FieldDef::input("a", FieldType::String, "first"))
.input(FieldDef::input("b", FieldType::Int, "second"))
.output(FieldDef::output("x", FieldType::String, "third"))
.output(FieldDef::output("y", FieldType::Bool, "fourth"))
.build()
.unwrap();
let input_names: Vec<_> = sig.input_fields().map(|f| f.name.as_str()).collect();
assert_eq!(input_names, vec!["a", "b"]);
let output_names: Vec<_> = sig.output_fields().map(|f| f.name.as_str()).collect();
assert_eq!(output_names, vec!["x", "y"]);
}
#[test]
fn dump_load_round_trip() {
let sig = Signature::builder("Classify sentiment.")
.input(FieldDef::input("text", FieldType::String, "Input text"))
.output(FieldDef::output(
"sentiment",
FieldType::Enum(vec!["positive".into(), "negative".into(), "neutral".into()]),
"The sentiment",
))
.build()
.unwrap();
let json = sig.dump_state().unwrap();
let restored = Signature::load_state(&json).unwrap();
assert_eq!(sig.instructions(), restored.instructions());
assert_eq!(sig.fields().len(), restored.fields().len());
for (a, b) in sig.fields().iter().zip(restored.fields().iter()) {
assert_eq!(a, b);
}
}
#[test]
fn dump_load_with_nested_types() {
let sig = Signature::builder("Extract info.")
.input(FieldDef::input("doc", FieldType::String, "The document"))
.output(FieldDef::output(
"entities",
FieldType::List(Box::new(FieldType::Object(vec![
crate::field::ObjectField {
name: "name".into(),
description: "Entity name".into(),
field_type: FieldType::String,
},
crate::field::ObjectField {
name: "type".into(),
description: "Entity type".into(),
field_type: FieldType::String,
},
]))),
"Extracted entities",
))
.build()
.unwrap();
let json = sig.dump_state().unwrap();
let restored = Signature::load_state(&json).unwrap();
assert_eq!(sig.fields(), restored.fields());
}
use crate::field::{ObjectField, OneOfDiscriminator, VariantArm};
fn tagged_oneof_object_arm(tag: &str, extra_field: (&str, FieldType)) -> VariantArm {
VariantArm {
description: tag.into(),
field_type: FieldType::Object(vec![
ObjectField {
name: "kind".into(),
description: String::new(),
field_type: FieldType::Enum(vec![tag.into()]),
},
ObjectField {
name: extra_field.0.into(),
description: String::new(),
field_type: extra_field.1,
},
]),
}
}
#[test]
fn build_rejects_oneof_with_empty_arms() {
let result = Signature::builder("inst")
.input(FieldDef::input("q", FieldType::String, ""))
.output(FieldDef::output(
"out",
FieldType::OneOf {
arms: vec![],
discriminator: None,
},
"",
))
.build();
let err = result.unwrap_err().to_string();
assert!(err.contains("at least one arm"), "got: {err}");
assert!(err.contains("OneOf"), "got: {err}");
}
#[test]
fn build_rejects_anyof_with_empty_arms() {
let result = Signature::builder("inst")
.input(FieldDef::input("q", FieldType::String, ""))
.output(FieldDef::output(
"out",
FieldType::AnyOf { arms: vec![] },
"",
))
.build();
let err = result.unwrap_err().to_string();
assert!(err.contains("at least one arm"), "got: {err}");
assert!(err.contains("AnyOf"), "got: {err}");
}
#[test]
fn build_rejects_discriminator_with_mismatched_tag_count() {
let result = Signature::builder("inst")
.input(FieldDef::input("q", FieldType::String, ""))
.output(FieldDef::output(
"out",
FieldType::OneOf {
arms: vec![
tagged_oneof_object_arm("a", ("x", FieldType::Int)),
tagged_oneof_object_arm("b", ("y", FieldType::String)),
],
discriminator: Some(OneOfDiscriminator {
property: "kind".into(),
tags: vec!["a".into()],
}),
},
"",
))
.build();
let err = result.unwrap_err().to_string();
assert!(err.contains("1 tags but 2 arms"), "got: {err}");
assert!(err.contains("parallel"), "got: {err}");
}
#[test]
fn build_rejects_duplicate_discriminator_tags() {
let result = Signature::builder("inst")
.input(FieldDef::input("q", FieldType::String, ""))
.output(FieldDef::output(
"out",
FieldType::OneOf {
arms: vec![
tagged_oneof_object_arm("dup", ("x", FieldType::Int)),
tagged_oneof_object_arm("dup", ("y", FieldType::String)),
],
discriminator: Some(OneOfDiscriminator {
property: "kind".into(),
tags: vec!["dup".into(), "dup".into()],
}),
},
"",
))
.build();
let err = result.unwrap_err().to_string();
assert!(err.contains("dup"), "got: {err}");
assert!(err.contains("unique"), "got: {err}");
}
#[test]
fn build_rejects_tagged_arm_that_is_not_an_object() {
let result = Signature::builder("inst")
.input(FieldDef::input("q", FieldType::String, ""))
.output(FieldDef::output(
"out",
FieldType::OneOf {
arms: vec![VariantArm {
description: "scalar arm".into(),
field_type: FieldType::Int,
}],
discriminator: Some(OneOfDiscriminator {
property: "kind".into(),
tags: vec!["a".into()],
}),
},
"",
))
.build();
let err = result.unwrap_err().to_string();
assert!(err.contains("must be an Object"), "got: {err}");
}
#[test]
fn build_rejects_tagged_arm_missing_discriminator_property() {
let result = Signature::builder("inst")
.input(FieldDef::input("q", FieldType::String, ""))
.output(FieldDef::output(
"out",
FieldType::OneOf {
arms: vec![VariantArm {
description: "no discriminator field".into(),
field_type: FieldType::Object(vec![ObjectField {
name: "x".into(),
description: String::new(),
field_type: FieldType::Int,
}]),
}],
discriminator: Some(OneOfDiscriminator {
property: "kind".into(),
tags: vec!["a".into()],
}),
},
"",
))
.build();
let err = result.unwrap_err().to_string();
assert!(
err.contains("missing the discriminator property"),
"got: {err}"
);
assert!(err.contains("kind"), "got: {err}");
}
#[test]
fn build_rejects_tagged_arm_with_wrong_enum_for_discriminator() {
let result = Signature::builder("inst")
.input(FieldDef::input("q", FieldType::String, ""))
.output(FieldDef::output(
"out",
FieldType::OneOf {
arms: vec![VariantArm {
description: "wrong tag enum".into(),
field_type: FieldType::Object(vec![
ObjectField {
name: "kind".into(),
description: String::new(),
field_type: FieldType::Enum(vec!["other".into()]),
},
ObjectField {
name: "x".into(),
description: String::new(),
field_type: FieldType::Int,
},
]),
}],
discriminator: Some(OneOfDiscriminator {
property: "kind".into(),
tags: vec!["a".into()],
}),
},
"",
))
.build();
let err = result.unwrap_err().to_string();
assert!(err.contains("Enum variant containing"), "got: {err}");
}
#[test]
fn build_accepts_well_formed_tagged_oneof() {
Signature::builder("inst")
.input(FieldDef::input("q", FieldType::String, ""))
.output(FieldDef::output(
"out",
FieldType::OneOf {
arms: vec![
tagged_oneof_object_arm("a", ("x", FieldType::Int)),
tagged_oneof_object_arm("b", ("y", FieldType::String)),
],
discriminator: Some(OneOfDiscriminator {
property: "kind".into(),
tags: vec!["a".into(), "b".into()],
}),
},
"",
))
.build()
.expect("well-formed tagged OneOf must build");
}
#[test]
fn build_validates_variant_nested_inside_list_and_nullable() {
let result = Signature::builder("inst")
.input(FieldDef::input("q", FieldType::String, ""))
.output(FieldDef::output(
"out",
FieldType::List(Box::new(FieldType::Nullable(Box::new(FieldType::OneOf {
arms: vec![],
discriminator: None,
})))),
"",
))
.build();
let err = result.unwrap_err().to_string();
assert!(err.contains("at least one arm"), "got: {err}");
}
#[test]
fn load_state_validates_fields() {
let json = serde_json::json!({
"instructions": "test",
"fields": [{
"name": "answer",
"description": "answer",
"field_type": {"kind": "string"},
"kind": "output"
}]
});
let result = Signature::load_state(&json.to_string());
assert!(result.is_err());
}
#[test]
fn builder_forces_kind() {
let wrong_kind = FieldDef::output("question", FieldType::String, "A question");
let sig = Signature::builder("test")
.input(wrong_kind) .output(FieldDef::output("answer", FieldType::String, "answer"))
.build()
.unwrap();
let input = sig.input_fields().next().unwrap();
assert_eq!(input.kind, FieldKind::Input);
}
}