1use std::time::{Duration, Instant};
2
3use indexmap::IndexMap;
4
5use crate::backend::{AnswerMeta, Backend, CascadeHop, Evaluated};
6use crate::error::BackendError;
7use crate::ids::QuestionId;
8use crate::wire::{Usage, WireAnswer, WireRequest, WireResponse};
9
10pub struct CascadeRule {
11 pub min: f64,
12 pub first_timeout: Duration,
13 pub always_fallback: Vec<QuestionId>,
14}
15
16impl CascadeRule {
17 pub fn new(min: f64) -> Self {
18 Self {
19 min,
20 first_timeout: Duration::from_millis(400),
21 always_fallback: battery_ids(),
22 }
23 }
24}
25
26pub struct Cascaded<A, B> {
27 pub first: A,
28 pub fallback: B,
29 pub rule: CascadeRule,
30}
31
32impl<A: Backend, B: Backend> Cascaded<A, B> {
33 pub fn new(first: A, fallback: B, rule: CascadeRule) -> Self {
34 Self {
35 first,
36 fallback,
37 rule,
38 }
39 }
40}
41
42pub fn battery_ids() -> Vec<QuestionId> {
43 crate::battery::shipped_questions()
44 .into_iter()
45 .map(|question| question.id().clone())
46 .collect()
47}
48
49pub(crate) fn fallback_still_below(
50 meta: &IndexMap<String, AnswerMeta>,
51 answers: &IndexMap<String, WireAnswer>,
52 min: f64,
53) -> bool {
54 meta.iter().any(|(id, row)| {
55 row.cascade_hop == Some(CascadeHop::Fallback)
56 && answers
57 .get(id)
58 .is_some_and(|answer| !answer_kept(answer, min))
59 })
60}
61
62pub(crate) fn answer_kept(answer: &WireAnswer, min: f64) -> bool {
63 let confidence = match answer {
64 WireAnswer::Choice { confidence, .. } | WireAnswer::Score { confidence, .. } => *confidence,
65 WireAnswer::Noul { noul } => (2.0 * noul - 1.0).abs(),
66 };
67 confidence >= min
68}
69
70impl<A: Backend, B: Backend> Backend for Cascaded<A, B> {
71 fn id(&self) -> &str {
72 "cascade"
73 }
74
75 fn replace_api_key(&self, key: Option<String>) -> Result<Option<String>, crate::error::Error> {
76 let previous_first = self.first.replace_api_key(key.clone())?;
77 match self.fallback.replace_api_key(key) {
78 Ok(_) => Ok(previous_first),
79 Err(err) => {
80 let _ = self.first.replace_api_key(previous_first);
81 Err(err)
82 }
83 }
84 }
85
86 async fn evaluate(
87 &self,
88 req: WireRequest,
89 deadline: Instant,
90 ) -> Result<Evaluated, BackendError> {
91 if Instant::now() >= deadline {
92 return Err(BackendError::Timeout);
93 }
94 let first_budget = deadline
95 .saturating_duration_since(Instant::now())
96 .min(self.rule.first_timeout);
97 let first_deadline = Instant::now() + first_budget;
98 let first_id = self.first.id().to_string();
99 let first = match self.first.evaluate(req.clone(), first_deadline).await {
100 Ok(first) => first,
101 Err(_) => {
102 let fallback = self.fallback.evaluate(req.clone(), deadline).await?;
103 return Ok(mark_fallback(fallback, &first_id, true));
104 }
105 };
106 let retry: Vec<String> = req
107 .questions
108 .keys()
109 .filter(|id| !keep_id(id, &first.wire, &self.rule))
110 .cloned()
111 .collect();
112 if retry.is_empty() {
113 return Ok(mark_first(first));
114 }
115 let mut subset = req.clone();
116 subset
117 .questions
118 .retain(|id, _| retry.iter().any(|retry_id| retry_id == id));
119 let fallback = match self.fallback.evaluate(subset, deadline).await {
120 Ok(fallback) => fallback,
121 Err(_) => return Ok(mark_first(first)),
122 };
123 Ok(merge(first, fallback, &retry))
124 }
125}
126
127fn keep_id(id: &str, wire: &WireResponse, rule: &CascadeRule) -> bool {
128 if rule.always_fallback.iter().any(|forced| forced.0 == id) {
129 return false;
130 }
131 match wire.answers.get(id) {
132 Some(answer) => answer_kept(answer, rule.min),
133 None => false,
134 }
135}
136
137fn mark_first(mut evaluated: Evaluated) -> Evaluated {
138 for id in evaluated.wire.answers.keys() {
139 evaluated.meta.insert(
140 id.clone(),
141 AnswerMeta {
142 cascade_hop: Some(CascadeHop::First),
143 ..AnswerMeta::default()
144 },
145 );
146 }
147 evaluated
148}
149
150fn mark_fallback(mut evaluated: Evaluated, first_id: &str, first_hop_error: bool) -> Evaluated {
151 evaluated.backend_id = format!("cascade:{first_id}+{}", evaluated.backend_id);
152 for id in evaluated.wire.answers.keys() {
153 evaluated.meta.insert(
154 id.clone(),
155 AnswerMeta {
156 cascade_hop: Some(CascadeHop::Fallback),
157 first_hop_error,
158 ..AnswerMeta::default()
159 },
160 );
161 }
162 evaluated
163}
164
165fn merge(first: Evaluated, fallback: Evaluated, retry: &[String]) -> Evaluated {
166 let mut wire = first.wire;
167 for id in retry {
168 if let Some(answer) = fallback.wire.answers.get(id) {
169 wire.answers.insert(id.clone(), answer.clone());
170 }
171 }
172 wire.usage = Usage {
173 input_tokens: wire
174 .usage
175 .input_tokens
176 .saturating_add(fallback.wire.usage.input_tokens),
177 output_tokens: wire
178 .usage
179 .output_tokens
180 .saturating_add(fallback.wire.usage.output_tokens),
181 };
182 wire.model = fallback.wire.model;
183 let mut meta = IndexMap::new();
184 for id in wire.answers.keys() {
185 let hop = if retry.iter().any(|retry_id| retry_id == id) {
186 CascadeHop::Fallback
187 } else {
188 CascadeHop::First
189 };
190 meta.insert(
191 id.clone(),
192 AnswerMeta {
193 cascade_hop: Some(hop),
194 ..AnswerMeta::default()
195 },
196 );
197 }
198 Evaluated {
199 wire,
200 meta,
201 backend_id: format!("cascade:{}+{}", first.backend_id, fallback.backend_id),
202 }
203}