Skip to main content

snapif/backends/
cascade.rs

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}