use crate::error::{Error, Result};
use serde::{Deserialize, Serialize};
use serde_json::{json, Map, Value};
use std::collections::BTreeMap;
use std::sync::Mutex;
pub trait DecideBackend: Send + Sync {
fn decide(&self, request_json: &str) -> Result<String>;
fn calibrated(&self) -> bool;
fn describe(&self) -> String;
}
impl<T: DecideBackend + ?Sized> DecideBackend for Box<T> {
fn decide(&self, request_json: &str) -> Result<String> {
(**self).decide(request_json)
}
fn calibrated(&self) -> bool {
(**self).calibrated()
}
fn describe(&self) -> String {
(**self).describe()
}
}
pub const DECIDE_MIN_P: f64 = 0.75;
pub const CAUSE_MIN_P: f64 = 0.6;
pub const DEFAULT_PAIR_CAP: usize = 200;
pub const QUESTIONS_PER_REQUEST: usize = 16;
pub const TOOL_CAUSES: &[(&str, &str)] = &[
("timeout", "The call ran out of time: a deadline, timeout or no response in time."),
(
"executor_error",
"The executor or the remote service failed while running the call: a transport fault, a crash, a 5xx, or an error the tool raised.",
),
(
"schema_validation_failed",
"The call's input or output did not match the tool's schema: a missing, extra or wrongly typed field.",
),
("user_aborted", "A person or the calling agent cancelled or aborted the call."),
(
"context_overflow",
"The model refused the request because the prompt was too long for its context window.",
),
("unknown", "None of the above, or the text does not say why the call failed."),
];
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct JudgedBy {
pub backend: String,
pub provider: String,
pub model: String,
pub calibrated: bool,
pub latency_ms: u64,
pub stage: String,
#[serde(default)]
pub answers: BTreeMap<String, f64>,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct DeciderReport {
pub backend: String,
pub calibrated: bool,
pub calls: u64,
pub failed_calls: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub last_error: Option<String>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct Answered {
pub noul: BTreeMap<String, f64>,
pub choices: BTreeMap<String, BTreeMap<String, f64>>,
pub provider: String,
pub model: String,
pub calibrated: bool,
pub latency_ms: u64,
}
impl Answered {
pub fn judged_by(&self, backend: &str, stage: &str, answers: BTreeMap<String, f64>) -> JudgedBy {
JudgedBy {
backend: backend.to_string(),
provider: self.provider.clone(),
model: self.model.clone(),
calibrated: self.calibrated,
latency_ms: self.latency_ms,
stage: stage.to_string(),
answers,
}
}
}
#[derive(Debug, Clone)]
pub enum Ask {
Noul { id: String, instructions: String },
Choice { id: String, instructions: String, options: Vec<(String, String)> },
}
impl Ask {
fn id(&self) -> &str {
match self {
Ask::Noul { id, .. } | Ask::Choice { id, .. } => id,
}
}
fn to_wire(&self) -> Value {
match self {
Ask::Noul { instructions, .. } => json!({"type": "noul", "instructions": instructions}),
Ask::Choice { instructions, options, .. } => {
let criteria: Map<String, Value> =
options.iter().map(|(k, d)| (k.clone(), Value::from(d.clone()))).collect();
json!({"type": "choice", "instructions": instructions, "criteria": criteria})
}
}
}
}
pub type CauseVerdict = Option<(String, f64, JudgedBy)>;
pub struct Decider {
backend: Box<dyn DecideBackend>,
pair_cap: usize,
stats: Mutex<DeciderReport>,
cause_cache: Mutex<BTreeMap<String, CauseVerdict>>,
}
impl Decider {
pub fn new(backend: Box<dyn DecideBackend>) -> Self {
Decider {
backend,
pair_cap: DEFAULT_PAIR_CAP,
stats: Mutex::new(DeciderReport::default()),
cause_cache: Mutex::new(BTreeMap::new()),
}
}
pub fn with_pair_cap(mut self, cap: usize) -> Self {
self.pair_cap = cap;
self
}
pub fn pair_cap(&self) -> usize {
self.pair_cap
}
pub fn calibrated(&self) -> bool {
self.backend.calibrated()
}
pub fn describe(&self) -> String {
self.backend.describe()
}
pub(crate) fn reset(&self) {
if let Ok(mut s) = self.stats.lock() {
*s = DeciderReport::default();
}
if let Ok(mut c) = self.cause_cache.lock() {
c.clear();
}
}
pub fn report(&self) -> DeciderReport {
let mut r = self.stats.lock().map(|s| s.clone()).unwrap_or_default();
r.backend = self.describe();
r.calibrated = self.calibrated();
r
}
pub fn ask(&self, state: Value, questions: &[Ask]) -> Result<Answered> {
let out = self.ask_inner(state, questions);
if let Ok(mut s) = self.stats.lock() {
s.calls += 1;
if let Err(e) = &out {
s.failed_calls += 1;
s.last_error = Some(e.to_string());
}
}
out
}
fn ask_inner(&self, state: Value, questions: &[Ask]) -> Result<Answered> {
if questions.is_empty() {
return Err(Error::DecideBackend("no questions".into()));
}
let qs: Map<String, Value> =
questions.iter().map(|q| (q.id().to_string(), q.to_wire())).collect();
let body = json!({"state": state, "questions": qs}).to_string();
let raw = self.backend.decide(&body).map_err(|e| match e {
Error::DecideBackend(_) => e,
other => Error::DecideBackend(other.to_string()),
})?;
parse_response(&raw, questions, self.backend.calibrated())
}
pub fn classify_causes(&self, texts: &[String]) -> BTreeMap<String, CauseVerdict> {
let mut out = BTreeMap::new();
if !self.calibrated() {
for t in texts {
out.insert(t.clone(), None);
}
return out;
}
let mut todo: Vec<String> = Vec::new();
{
let cache = self.cause_cache.lock().ok();
for t in texts {
match cache.as_ref().and_then(|c| c.get(t)) {
Some(hit) => {
out.insert(t.clone(), hit.clone());
}
None if !todo.contains(t) => todo.push(t.clone()),
None => {}
}
}
}
let asked_before = self.cause_cache.lock().map(|c| c.len()).unwrap_or(0);
let budget = self.pair_cap.saturating_sub(asked_before);
let (ask, skip) = todo.split_at(todo.len().min(budget));
for t in skip {
out.insert(t.clone(), None);
}
let backend = self.describe();
for chunk in ask.chunks(QUESTIONS_PER_REQUEST) {
let mut failures = Map::new();
let mut questions = Vec::new();
for (i, t) in chunk.iter().enumerate() {
let id = format!("c{i}");
failures.insert(id.clone(), Value::from(t.clone()));
questions.push(Ask::Choice {
instructions: format!(
"Which cause best explains the tool failure described by item \"{id}\" (in state.failures)?"
),
id,
options: TOOL_CAUSES.iter().map(|(k, d)| (k.to_string(), d.to_string())).collect(),
});
}
let answered = self.ask(json!({"failures": failures}), &questions);
for (i, t) in chunk.iter().enumerate() {
let verdict = answered.as_ref().ok().and_then(|a| {
if !a.calibrated {
return None;
}
let probs = a.choices.get(&format!("c{i}"))?;
let (best, p) = argmax(probs)?;
(p >= CAUSE_MIN_P).then(|| {
let judged = a.judged_by(&backend, "tool_cause", probs.clone());
(best, p, judged)
})
});
if let Ok(mut c) = self.cause_cache.lock() {
c.insert(t.clone(), verdict.clone());
}
out.insert(t.clone(), verdict);
}
}
out
}
}
fn argmax(probs: &BTreeMap<String, f64>) -> Option<(String, f64)> {
let mut best: Option<(&String, f64)> = None;
for (k, &p) in probs {
if best.is_none_or(|(_, bp)| p > bp) {
best = Some((k, p));
}
}
best.map(|(k, p)| (k.clone(), p))
}
fn prob(v: &Value) -> Option<f64> {
v.as_f64().filter(|p| p.is_finite() && (0.0..=1.0).contains(p))
}
pub fn parse_response(raw: &str, asked: &[Ask], backend_calibrated: bool) -> Result<Answered> {
let bad = |m: String| Error::DecideBackend(format!("malformed answer: {m}"));
let v: Value = serde_json::from_str(raw.trim()).map_err(|e| bad(format!("not JSON: {e}")))?;
let answers = v
.get("answers")
.and_then(Value::as_object)
.ok_or_else(|| bad("no `answers` object".into()))?;
let mut noul = BTreeMap::new();
let mut choices = BTreeMap::new();
for q in asked {
let a = answers.get(q.id()).ok_or_else(|| bad(format!("no answer for {:?}", q.id())))?;
match q {
Ask::Noul { id, .. } => {
if a.get("type").and_then(Value::as_str).is_some_and(|t| t != "noul") {
return Err(bad(format!("{id:?} is not a noul answer")));
}
let p = a
.get("noul")
.and_then(prob)
.ok_or_else(|| bad(format!("{id:?} has no probability in [0, 1]")))?;
noul.insert(id.clone(), p);
}
Ask::Choice { id, options, .. } => {
if a.get("type").and_then(Value::as_str).is_some_and(|t| t != "choice") {
return Err(bad(format!("{id:?} is not a choice answer")));
}
let probs = a
.get("probabilities")
.and_then(Value::as_object)
.ok_or_else(|| bad(format!("{id:?} has no probabilities")))?;
let mut m = BTreeMap::new();
for (k, p) in probs {
if !options.iter().any(|(o, _)| o == k) {
return Err(bad(format!("{id:?} answered an option nobody offered: {k:?}")));
}
let p = prob(p).ok_or_else(|| bad(format!("{id:?} option {k:?} is not in [0, 1]")))?;
m.insert(k.clone(), p);
}
if m.is_empty() {
return Err(bad(format!("{id:?} has no probabilities")));
}
choices.insert(id.clone(), m);
}
}
}
let s = |k: &str| v.get(k).and_then(Value::as_str).unwrap_or("").to_string();
Ok(Answered {
noul,
choices,
provider: s("provider"),
model: s("model"),
calibrated: backend_calibrated && v.get("calibrated").and_then(Value::as_bool).unwrap_or(true),
latency_ms: v.get("latency_ms").and_then(Value::as_u64).unwrap_or(0),
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_response_must_answer_every_question_with_a_probability() {
let asked = vec![
Ask::Noul { id: "a".into(), instructions: "?".into() },
Ask::Choice {
id: "b".into(),
instructions: "?".into(),
options: vec![("x".into(), "".into()), ("y".into(), "".into())],
},
];
let ok = r#"{"answers":{"a":{"type":"noul","noul":0.9},
"b":{"type":"choice","choice":"x","probabilities":{"x":0.7,"y":0.3}}},
"provider":"fake","model":"m","calibrated":true,"latency_ms":3}"#;
let a = parse_response(ok, &asked, true).unwrap();
assert_eq!(a.noul["a"], 0.9);
assert_eq!(a.choices["b"]["x"], 0.7);
assert!(a.calibrated);
assert_eq!(a.latency_ms, 3);
let unc = ok.replace("\"calibrated\":true", "\"calibrated\":false");
assert!(!parse_response(&unc, &asked, true).unwrap().calibrated);
assert!(!parse_response(ok, &asked, false).unwrap().calibrated);
for broken in [
"not json",
r#"{"answers":{"a":{"type":"noul","noul":0.9}}}"#,
r#"{"answers":{"a":{"type":"noul","noul":1.5},"b":{"probabilities":{"x":1}}}}"#,
r#"{"answers":{"a":{"type":"noul","noul":0.5},"b":{"probabilities":{"z":1}}}}"#,
r#"{"answers":{"a":{"type":"choice","noul":0.5},"b":{"probabilities":{"x":1}}}}"#,
] {
let e = parse_response(broken, &asked, true).unwrap_err();
assert_eq!(e.code(), "LOP-E051", "{broken}");
}
}
#[test]
fn argmax_breaks_ties_deterministically() {
let m: BTreeMap<String, f64> = [("b".to_string(), 0.5), ("a".to_string(), 0.5)].into();
assert_eq!(argmax(&m), Some(("a".to_string(), 0.5)));
assert_eq!(argmax(&BTreeMap::new()), None);
}
}