use std::fmt;
use crate::ProviderError;
use crate::schema::Violation;
#[derive(Debug)]
#[non_exhaustive]
pub enum Error {
Provider(ProviderError),
InvalidInput {
message: String,
},
Refusal {
model: String,
message: String,
},
Truncated {
model: String,
},
InvalidJson {
reason: String,
response: String,
},
SchemaViolation {
violations: Vec<Violation>,
response: String,
},
Deserialize {
message: String,
response: String,
},
Validation {
errors: Vec<String>,
},
RetryLimitExceeded {
attempts: Vec<Error>,
},
}
impl Error {
pub fn kind(&self) -> &'static str {
match self {
Self::Provider(_) => "provider",
Self::InvalidInput { .. } => "invalid_input",
Self::Refusal { .. } => "refusal",
Self::Truncated { .. } => "truncated",
Self::InvalidJson { .. } => "invalid_json",
Self::SchemaViolation { .. } => "schema_violation",
Self::Deserialize { .. } => "deserialize",
Self::Validation { .. } => "validation",
Self::RetryLimitExceeded { .. } => "retry_limit_exceeded",
}
}
pub fn is_repairable(&self) -> bool {
matches!(
self,
Self::InvalidJson { .. }
| Self::SchemaViolation { .. }
| Self::Deserialize { .. }
| Self::Validation { .. }
)
}
pub fn repair_feedback(&self) -> Option<String> {
let details = match self {
Self::InvalidJson { reason, .. } => {
format!("It is not valid JSON ({reason}). Respond with a single JSON value only.")
}
Self::SchemaViolation { violations, .. } => {
let lines: Vec<String> = violations.iter().map(|v| format!("- {v}")).collect();
format!("It violates the output schema:\n{}", lines.join("\n"))
}
Self::Deserialize { message, .. } => {
format!("It does not match the expected structure: {message}")
}
Self::Validation { errors } => {
let lines: Vec<String> = errors.iter().map(|e| format!("- {e}")).collect();
format!("It fails validation:\n{}", lines.join("\n"))
}
_ => return None,
};
Some(format!(
"Your previous response does not satisfy the output contract. {details}\n\
Return a corrected response."
))
}
}
impl fmt::Display for Error {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Provider(e) => write!(f, "{e}"),
Self::InvalidInput { message } => write!(f, "input cannot be serialized: {message}"),
Self::Refusal { model, .. } => write!(f, "model {model} refused to answer"),
Self::Truncated { model } => write!(f, "answer of model {model} was truncated"),
Self::InvalidJson { reason, .. } => write!(f, "answer is not valid JSON: {reason}"),
Self::SchemaViolation { violations, .. } => {
let paths: Vec<&str> = violations.iter().map(|v| v.path.as_str()).collect();
write!(
f,
"answer violates the output schema at {}",
paths.join(", ")
)
}
Self::Deserialize { .. } => {
write!(f, "answer does not deserialize into the output type")
}
Self::Validation { errors } => {
write!(f, "answer failed validation: {}", errors.join("; "))
}
Self::RetryLimitExceeded { attempts } => {
write!(f, "no valid answer after {} attempts", attempts.len())?;
if let Some(last) = attempts.last() {
write!(f, ", last error: {last}")?;
}
Ok(())
}
}
}
}
impl std::error::Error for Error {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Provider(e) => Some(e),
Self::RetryLimitExceeded { attempts } => attempts
.last()
.map(|e| e as &(dyn std::error::Error + 'static)),
_ => None,
}
}
}
impl From<ProviderError> for Error {
fn from(e: ProviderError) -> Self {
Self::Provider(e)
}
}