1use crate::error::{Error, Result};
36use serde::{Deserialize, Serialize};
37use serde_json::{json, Map, Value};
38use std::collections::BTreeMap;
39use std::sync::Mutex;
40
41pub trait DecideBackend: Send + Sync {
43 fn decide(&self, request_json: &str) -> Result<String>;
48 fn calibrated(&self) -> bool;
51 fn describe(&self) -> String;
53}
54
55impl<T: DecideBackend + ?Sized> DecideBackend for Box<T> {
56 fn decide(&self, request_json: &str) -> Result<String> {
57 (**self).decide(request_json)
58 }
59 fn calibrated(&self) -> bool {
60 (**self).calibrated()
61 }
62 fn describe(&self) -> String {
63 (**self).describe()
64 }
65}
66
67pub const DECIDE_MIN_P: f64 = 0.75;
71pub const CAUSE_MIN_P: f64 = 0.6;
74pub const DEFAULT_PAIR_CAP: usize = 200;
76pub const QUESTIONS_PER_REQUEST: usize = 16;
78
79pub const TOOL_CAUSES: &[(&str, &str)] = &[
83 ("timeout", "The call ran out of time: a deadline, timeout or no response in time."),
84 (
85 "executor_error",
86 "The executor or the remote service failed while running the call: a transport fault, a crash, a 5xx, or an error the tool raised.",
87 ),
88 (
89 "schema_validation_failed",
90 "The call's input or output did not match the tool's schema: a missing, extra or wrongly typed field.",
91 ),
92 ("user_aborted", "A person or the calling agent cancelled or aborted the call."),
93 (
94 "context_overflow",
95 "The model refused the request because the prompt was too long for its context window.",
96 ),
97 ("unknown", "None of the above, or the text does not say why the call failed."),
98];
99
100#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
104pub struct JudgedBy {
105 pub backend: String,
107 pub provider: String,
109 pub model: String,
110 pub calibrated: bool,
111 pub latency_ms: u64,
112 pub stage: String,
115 #[serde(default)]
117 pub answers: BTreeMap<String, f64>,
118}
119
120#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
123pub struct DeciderReport {
124 pub backend: String,
126 pub calibrated: bool,
129 pub calls: u64,
131 pub failed_calls: u64,
134 #[serde(default, skip_serializing_if = "Option::is_none")]
136 pub last_error: Option<String>,
137}
138
139#[derive(Debug, Clone, PartialEq)]
141pub struct Answered {
142 pub noul: BTreeMap<String, f64>,
144 pub choices: BTreeMap<String, BTreeMap<String, f64>>,
146 pub provider: String,
147 pub model: String,
148 pub calibrated: bool,
150 pub latency_ms: u64,
151}
152
153impl Answered {
154 pub fn judged_by(&self, backend: &str, stage: &str, answers: BTreeMap<String, f64>) -> JudgedBy {
156 JudgedBy {
157 backend: backend.to_string(),
158 provider: self.provider.clone(),
159 model: self.model.clone(),
160 calibrated: self.calibrated,
161 latency_ms: self.latency_ms,
162 stage: stage.to_string(),
163 answers,
164 }
165 }
166}
167
168#[derive(Debug, Clone)]
171pub enum Ask {
172 Noul { id: String, instructions: String },
173 Choice { id: String, instructions: String, options: Vec<(String, String)> },
174}
175
176impl Ask {
177 fn id(&self) -> &str {
178 match self {
179 Ask::Noul { id, .. } | Ask::Choice { id, .. } => id,
180 }
181 }
182 fn to_wire(&self) -> Value {
183 match self {
184 Ask::Noul { instructions, .. } => json!({"type": "noul", "instructions": instructions}),
185 Ask::Choice { instructions, options, .. } => {
186 let criteria: Map<String, Value> =
187 options.iter().map(|(k, d)| (k.clone(), Value::from(d.clone()))).collect();
188 json!({"type": "choice", "instructions": instructions, "criteria": criteria})
189 }
190 }
191 }
192}
193
194pub type CauseVerdict = Option<(String, f64, JudgedBy)>;
197
198pub struct Decider {
202 backend: Box<dyn DecideBackend>,
203 pair_cap: usize,
204 stats: Mutex<DeciderReport>,
205 cause_cache: Mutex<BTreeMap<String, CauseVerdict>>,
206}
207
208impl Decider {
209 pub fn new(backend: Box<dyn DecideBackend>) -> Self {
210 Decider {
211 backend,
212 pair_cap: DEFAULT_PAIR_CAP,
213 stats: Mutex::new(DeciderReport::default()),
214 cause_cache: Mutex::new(BTreeMap::new()),
215 }
216 }
217
218 pub fn with_pair_cap(mut self, cap: usize) -> Self {
220 self.pair_cap = cap;
221 self
222 }
223
224 pub fn pair_cap(&self) -> usize {
225 self.pair_cap
226 }
227 pub fn calibrated(&self) -> bool {
228 self.backend.calibrated()
229 }
230 pub fn describe(&self) -> String {
231 self.backend.describe()
232 }
233
234 pub(crate) fn reset(&self) {
237 if let Ok(mut s) = self.stats.lock() {
238 *s = DeciderReport::default();
239 }
240 if let Ok(mut c) = self.cause_cache.lock() {
241 c.clear();
242 }
243 }
244
245 pub fn report(&self) -> DeciderReport {
247 let mut r = self.stats.lock().map(|s| s.clone()).unwrap_or_default();
248 r.backend = self.describe();
249 r.calibrated = self.calibrated();
250 r
251 }
252
253 pub fn ask(&self, state: Value, questions: &[Ask]) -> Result<Answered> {
257 let out = self.ask_inner(state, questions);
258 if let Ok(mut s) = self.stats.lock() {
259 s.calls += 1;
260 if let Err(e) = &out {
261 s.failed_calls += 1;
262 s.last_error = Some(e.to_string());
263 }
264 }
265 out
266 }
267
268 fn ask_inner(&self, state: Value, questions: &[Ask]) -> Result<Answered> {
269 if questions.is_empty() {
270 return Err(Error::DecideBackend("no questions".into()));
271 }
272 let qs: Map<String, Value> =
273 questions.iter().map(|q| (q.id().to_string(), q.to_wire())).collect();
274 let body = json!({"state": state, "questions": qs}).to_string();
275 let raw = self.backend.decide(&body).map_err(|e| match e {
276 Error::DecideBackend(_) => e,
277 other => Error::DecideBackend(other.to_string()),
278 })?;
279 parse_response(&raw, questions, self.backend.calibrated())
280 }
281
282 pub fn classify_causes(&self, texts: &[String]) -> BTreeMap<String, CauseVerdict> {
290 let mut out = BTreeMap::new();
291 if !self.calibrated() {
292 for t in texts {
293 out.insert(t.clone(), None);
294 }
295 return out;
296 }
297 let mut todo: Vec<String> = Vec::new();
298 {
299 let cache = self.cause_cache.lock().ok();
300 for t in texts {
301 match cache.as_ref().and_then(|c| c.get(t)) {
302 Some(hit) => {
303 out.insert(t.clone(), hit.clone());
304 }
305 None if !todo.contains(t) => todo.push(t.clone()),
306 None => {}
307 }
308 }
309 }
310 let asked_before = self.cause_cache.lock().map(|c| c.len()).unwrap_or(0);
311 let budget = self.pair_cap.saturating_sub(asked_before);
312 let (ask, skip) = todo.split_at(todo.len().min(budget));
313 for t in skip {
314 out.insert(t.clone(), None);
315 }
316 let backend = self.describe();
317 for chunk in ask.chunks(QUESTIONS_PER_REQUEST) {
318 let mut failures = Map::new();
319 let mut questions = Vec::new();
320 for (i, t) in chunk.iter().enumerate() {
321 let id = format!("c{i}");
322 failures.insert(id.clone(), Value::from(t.clone()));
323 questions.push(Ask::Choice {
324 instructions: format!(
325 "Which cause best explains the tool failure described by item \"{id}\" (in state.failures)?"
326 ),
327 id,
328 options: TOOL_CAUSES.iter().map(|(k, d)| (k.to_string(), d.to_string())).collect(),
329 });
330 }
331 let answered = self.ask(json!({"failures": failures}), &questions);
332 for (i, t) in chunk.iter().enumerate() {
333 let verdict = answered.as_ref().ok().and_then(|a| {
334 if !a.calibrated {
335 return None;
336 }
337 let probs = a.choices.get(&format!("c{i}"))?;
338 let (best, p) = argmax(probs)?;
339 (p >= CAUSE_MIN_P).then(|| {
340 let judged = a.judged_by(&backend, "tool_cause", probs.clone());
341 (best, p, judged)
342 })
343 });
344 if let Ok(mut c) = self.cause_cache.lock() {
345 c.insert(t.clone(), verdict.clone());
346 }
347 out.insert(t.clone(), verdict);
348 }
349 }
350 out
351 }
352}
353
354fn argmax(probs: &BTreeMap<String, f64>) -> Option<(String, f64)> {
357 let mut best: Option<(&String, f64)> = None;
358 for (k, &p) in probs {
359 if best.is_none_or(|(_, bp)| p > bp) {
360 best = Some((k, p));
361 }
362 }
363 best.map(|(k, p)| (k.clone(), p))
364}
365
366fn prob(v: &Value) -> Option<f64> {
367 v.as_f64().filter(|p| p.is_finite() && (0.0..=1.0).contains(p))
368}
369
370pub fn parse_response(raw: &str, asked: &[Ask], backend_calibrated: bool) -> Result<Answered> {
374 let bad = |m: String| Error::DecideBackend(format!("malformed answer: {m}"));
375 let v: Value = serde_json::from_str(raw.trim()).map_err(|e| bad(format!("not JSON: {e}")))?;
376 let answers = v
377 .get("answers")
378 .and_then(Value::as_object)
379 .ok_or_else(|| bad("no `answers` object".into()))?;
380 let mut noul = BTreeMap::new();
381 let mut choices = BTreeMap::new();
382 for q in asked {
383 let a = answers.get(q.id()).ok_or_else(|| bad(format!("no answer for {:?}", q.id())))?;
384 match q {
385 Ask::Noul { id, .. } => {
386 if a.get("type").and_then(Value::as_str).is_some_and(|t| t != "noul") {
387 return Err(bad(format!("{id:?} is not a noul answer")));
388 }
389 let p = a
390 .get("noul")
391 .and_then(prob)
392 .ok_or_else(|| bad(format!("{id:?} has no probability in [0, 1]")))?;
393 noul.insert(id.clone(), p);
394 }
395 Ask::Choice { id, options, .. } => {
396 if a.get("type").and_then(Value::as_str).is_some_and(|t| t != "choice") {
397 return Err(bad(format!("{id:?} is not a choice answer")));
398 }
399 let probs = a
400 .get("probabilities")
401 .and_then(Value::as_object)
402 .ok_or_else(|| bad(format!("{id:?} has no probabilities")))?;
403 let mut m = BTreeMap::new();
404 for (k, p) in probs {
405 if !options.iter().any(|(o, _)| o == k) {
406 return Err(bad(format!("{id:?} answered an option nobody offered: {k:?}")));
407 }
408 let p = prob(p).ok_or_else(|| bad(format!("{id:?} option {k:?} is not in [0, 1]")))?;
409 m.insert(k.clone(), p);
410 }
411 if m.is_empty() {
412 return Err(bad(format!("{id:?} has no probabilities")));
413 }
414 choices.insert(id.clone(), m);
415 }
416 }
417 }
418 let s = |k: &str| v.get(k).and_then(Value::as_str).unwrap_or("").to_string();
419 Ok(Answered {
420 noul,
421 choices,
422 provider: s("provider"),
423 model: s("model"),
424 calibrated: backend_calibrated && v.get("calibrated").and_then(Value::as_bool).unwrap_or(true),
425 latency_ms: v.get("latency_ms").and_then(Value::as_u64).unwrap_or(0),
426 })
427}
428
429#[cfg(test)]
430mod tests {
431 use super::*;
432
433 #[test]
434 fn a_response_must_answer_every_question_with_a_probability() {
435 let asked = vec![
436 Ask::Noul { id: "a".into(), instructions: "?".into() },
437 Ask::Choice {
438 id: "b".into(),
439 instructions: "?".into(),
440 options: vec![("x".into(), "".into()), ("y".into(), "".into())],
441 },
442 ];
443 let ok = r#"{"answers":{"a":{"type":"noul","noul":0.9},
444 "b":{"type":"choice","choice":"x","probabilities":{"x":0.7,"y":0.3}}},
445 "provider":"fake","model":"m","calibrated":true,"latency_ms":3}"#;
446 let a = parse_response(ok, &asked, true).unwrap();
447 assert_eq!(a.noul["a"], 0.9);
448 assert_eq!(a.choices["b"]["x"], 0.7);
449 assert!(a.calibrated);
450 assert_eq!(a.latency_ms, 3);
451 let unc = ok.replace("\"calibrated\":true", "\"calibrated\":false");
453 assert!(!parse_response(&unc, &asked, true).unwrap().calibrated);
454 assert!(!parse_response(ok, &asked, false).unwrap().calibrated);
455
456 for broken in [
457 "not json",
458 r#"{"answers":{"a":{"type":"noul","noul":0.9}}}"#,
459 r#"{"answers":{"a":{"type":"noul","noul":1.5},"b":{"probabilities":{"x":1}}}}"#,
460 r#"{"answers":{"a":{"type":"noul","noul":0.5},"b":{"probabilities":{"z":1}}}}"#,
461 r#"{"answers":{"a":{"type":"choice","noul":0.5},"b":{"probabilities":{"x":1}}}}"#,
462 ] {
463 let e = parse_response(broken, &asked, true).unwrap_err();
464 assert_eq!(e.code(), "LOP-E051", "{broken}");
465 }
466 }
467
468 #[test]
469 fn argmax_breaks_ties_deterministically() {
470 let m: BTreeMap<String, f64> = [("b".to_string(), 0.5), ("a".to_string(), 0.5)].into();
471 assert_eq!(argmax(&m), Some(("a".to_string(), 0.5)));
472 assert_eq!(argmax(&BTreeMap::new()), None);
473 }
474}