use std::{collections::BTreeMap, sync::Arc};
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
pub mod fallback;
#[cfg(feature = "test-fixtures")]
pub mod stub;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct NoulCriteria {
#[serde(rename = "true")]
pub yes: String,
#[serde(rename = "false")]
pub no: String,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum Question {
Noul {
instructions: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
criteria: Option<NoulCriteria>,
},
Choice {
instructions: String,
criteria: BTreeMap<String, Option<String>>,
},
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct JudgmentRequest {
pub state: serde_json::Value,
pub questions: BTreeMap<String, Question>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum Answer {
Noul {
noul: f64,
},
Choice {
choice: String,
probabilities: BTreeMap<String, f64>,
confidence: f64,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
pub struct JudgmentUsage {
pub input_tokens: u64,
pub output_tokens: u64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum JudgmentSource {
#[default]
Primary,
Fallback,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct JudgmentResponse {
pub model: String,
pub answers: BTreeMap<String, Answer>,
pub usage: JudgmentUsage,
#[serde(default)]
pub source: JudgmentSource,
}
impl JudgmentResponse {
#[must_use]
pub fn noul(&self, id: &str) -> Option<f64> {
match self.answers.get(id) {
Some(Answer::Noul { noul }) => Some(*noul),
_ => None,
}
}
}
#[derive(Debug, thiserror::Error)]
pub enum JudgmentError {
#[error("judgment request rejected: {0}")]
Invalid(String),
#[error("judgment backend refused the credential")]
Unauthorized,
#[error("judgment backend refused for lack of credit or quota")]
Exhausted,
#[error("judgment backend rate limited (retry after {retry_after:?})")]
RateLimited {
retry_after: Option<std::time::Duration>,
},
#[error("judgment backend unavailable (status {status})")]
Unavailable {
status: u16,
},
#[error("judgment transport failed: {0}")]
Transport(Box<dyn std::error::Error + Send + Sync + 'static>),
#[error("judgment response malformed: {0}")]
Malformed(String),
}
#[async_trait]
pub trait JudgmentProvider: Send + Sync + 'static {
type Error: std::error::Error + Send + Sync + 'static;
async fn judge(&self, request: JudgmentRequest) -> Result<JudgmentResponse, Self::Error>;
}
#[derive(Debug)]
pub struct BoxError(Box<dyn std::error::Error + Send + Sync + 'static>);
impl BoxError {
#[must_use]
pub fn new<E: std::error::Error + Send + Sync + 'static>(err: E) -> Self {
Self(Box::new(err))
}
}
impl std::fmt::Display for BoxError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
std::fmt::Display::fmt(&self.0, f)
}
}
impl std::error::Error for BoxError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
Some(&*self.0)
}
}
pub struct ErasedJudgment<P>(P);
#[async_trait]
impl<P: JudgmentProvider> JudgmentProvider for ErasedJudgment<P> {
type Error = BoxError;
async fn judge(&self, request: JudgmentRequest) -> Result<JudgmentResponse, Self::Error> {
self.0.judge(request).await.map_err(BoxError::new)
}
}
pub type DynJudgment = dyn JudgmentProvider<Error = BoxError>;
#[must_use]
pub fn into_dyn<P: JudgmentProvider>(provider: P) -> Arc<DynJudgment> {
Arc::new(ErasedJudgment(provider))
}
#[cfg(test)]
mod tests {
#![allow(clippy::pedantic, clippy::nursery, missing_docs)]
use super::*;
#[test]
fn noul_question_serializes_to_the_wire_shape() {
let q = Question::Noul {
instructions: "Is it urgent?".to_owned(),
criteria: Some(NoulCriteria {
yes: "Time-sensitive".to_owned(),
no: "No urgency".to_owned(),
}),
};
let json = serde_json::to_value(&q).expect("serializable");
assert_eq!(
json,
serde_json::json!({
"type": "noul",
"instructions": "Is it urgent?",
"criteria": {"true": "Time-sensitive", "false": "No urgency"}
})
);
}
#[test]
fn noul_question_without_criteria_omits_the_field() {
let q = Question::Noul {
instructions: "Is it urgent?".to_owned(),
criteria: None,
};
let json = serde_json::to_value(&q).expect("serializable");
assert!(json.get("criteria").is_none(), "{json}");
}
#[test]
fn choice_answer_round_trips() {
let raw = serde_json::json!({
"type": "choice",
"choice": "billing",
"probabilities": {"billing": 0.9, "sales": 0.1},
"confidence": 0.85
});
let answer: Answer = serde_json::from_value(raw.clone()).expect("decodes");
assert!(matches!(&answer, Answer::Choice { choice, .. } if choice == "billing"));
assert_eq!(serde_json::to_value(&answer).expect("encodes"), raw);
}
#[test]
fn response_noul_accessor_is_none_for_other_shapes() {
let mut answers = BTreeMap::new();
answers.insert("yes".to_owned(), Answer::Noul { noul: 0.7 });
answers.insert(
"pick".to_owned(),
Answer::Choice {
choice: "a".to_owned(),
probabilities: BTreeMap::new(),
confidence: 1.0,
},
);
let response = JudgmentResponse {
model: "m".to_owned(),
answers,
usage: JudgmentUsage::default(),
source: JudgmentSource::Primary,
};
assert_eq!(response.noul("yes"), Some(0.7));
assert_eq!(response.noul("pick"), None);
assert_eq!(response.noul("absent"), None);
}
}