use std::fmt;
use std::path::Path;
use serde::de::DeserializeOwned;
use serde::{Deserialize, Serialize};
use crate::schema::validate;
use crate::{
Demonstration, GenerationOptions, Program, Provider, Signature, Strategy, output_schema,
};
pub const FORMAT: u32 = 1;
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CompiledProgram {
pub format: u32,
pub program: String,
pub signature: String,
pub instructions: Option<String>,
pub demonstrations: Vec<Demonstration>,
pub strategy: Option<Strategy>,
pub generation: GenerationOptions,
pub max_repairs: usize,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provenance: Option<Provenance>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Provenance {
pub model: Option<String>,
pub metric: String,
pub score: f64,
pub interval: (f64, f64),
pub examples: usize,
pub dataset: String,
}
#[cfg(feature = "eval")]
impl From<&crate::eval::Report> for Provenance {
fn from(report: &crate::eval::Report) -> Self {
Self {
model: report.label.clone().or_else(|| report.model.clone()),
metric: report.metric.clone(),
score: report.score,
interval: report.interval,
examples: report.examples,
dataset: report.dataset.clone(),
}
}
}
#[derive(Debug)]
#[non_exhaustive]
pub enum CompileError {
Io {
path: String,
message: String,
},
Invalid {
message: String,
},
SignatureMismatch {
program: String,
expected: String,
found: String,
},
InvalidDemonstration {
index: usize,
message: String,
},
}
impl fmt::Display for CompileError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Io { path, message } => write!(f, "{path}: {message}"),
Self::Invalid { message } => write!(f, "not a compiled program: {message}"),
Self::SignatureMismatch {
program,
expected,
found,
} => write!(
f,
"compiled for {found}, not for {program} ({expected}); the output type or \
the name changed — compile the program again"
),
Self::InvalidDemonstration { index, message } => {
write!(f, "demonstration {index}: {message}")
}
}
}
}
impl std::error::Error for CompileError {}
pub fn signature_hash<S: Signature>() -> String {
let fingerprint = serde_json::json!({ "name": S::NAME, "output": output_schema::<S>() });
let bytes = serde_json::to_vec(&fingerprint).expect("JSON values always serialize");
crate::sha256::hex(&crate::sha256::sha256(&bytes))
}
impl CompiledProgram {
pub fn with_provenance(mut self, provenance: impl Into<Provenance>) -> Self {
self.provenance = Some(provenance.into());
self
}
pub fn to_json(&self) -> String {
let mut json = serde_json::to_string_pretty(self).expect("compiled programs serialize");
json.push('\n');
json
}
pub fn from_json(json: &str) -> Result<Self, CompileError> {
let compiled: Self = serde_json::from_str(json).map_err(|e| CompileError::Invalid {
message: e.to_string(),
})?;
if compiled.format > FORMAT {
return Err(CompileError::Invalid {
message: format!(
"format {} is newer than this version of typedlm reads ({FORMAT})",
compiled.format
),
});
}
Ok(compiled)
}
pub fn load(path: impl AsRef<Path>) -> Result<Self, CompileError> {
let path = path.as_ref();
let text = std::fs::read_to_string(path).map_err(|e| io(path, e))?;
Self::from_json(&text)
}
pub fn save(&self, path: impl AsRef<Path>) -> Result<(), CompileError> {
let path = path.as_ref();
std::fs::write(path, self.to_json()).map_err(|e| io(path, e))
}
pub fn program<S, P>(&self, provider: P) -> Result<Program<S, P>, CompileError>
where
S: Signature + DeserializeOwned,
P: Provider,
{
self.check::<S>()?;
Ok(Program::from_parts(provider, self))
}
fn check<S: Signature + DeserializeOwned>(&self) -> Result<(), CompileError> {
let expected = signature_hash::<S>();
if self.program != S::NAME || self.signature != expected {
return Err(CompileError::SignatureMismatch {
program: S::NAME.into(),
expected: short(&expected),
found: format!("{} ({})", self.program, short(&self.signature)),
});
}
let schema = output_schema::<S>();
for (index, demo) in self.demonstrations.iter().enumerate() {
let invalid = |message: String| CompileError::InvalidDemonstration { index, message };
serde_json::from_value::<S>(demo.input.clone())
.map_err(|e| invalid(format!("input does not fit {}: {e}", S::NAME)))?;
let violations = validate(&schema, &demo.output);
if let Some(first) = violations.first() {
return Err(invalid(format!("output violates the schema: {first}")));
}
serde_json::from_value::<S::Output>(demo.output.clone())
.map_err(|e| invalid(format!("output does not deserialize: {e}")))?;
}
Ok(())
}
}
fn short(hash: &str) -> String {
hash.chars().take(12).collect()
}
fn io(path: &Path, error: std::io::Error) -> CompileError {
CompileError::Io {
path: path.display().to_string(),
message: error.to_string(),
}
}