Skip to main content

polyc_judgment/
fallback.rs

1//! [`FallbackJudgment`]: compose a primary judgment backend with a fallback
2//! one.
3
4use async_trait::async_trait;
5
6use crate::{JudgmentError, JudgmentProvider, JudgmentRequest, JudgmentResponse, JudgmentSource};
7
8/// A judgment backend that tries a primary backend first and falls back to a
9/// second backend when the primary fails in a way the fallback can fix.
10///
11/// Both backends answer through the shared [`JudgmentError`] taxonomy, so the
12/// decision of whether to fall back is one match over `P`'s error, not a
13/// per-backend special case.
14pub struct FallbackJudgment<P, F> {
15    primary: P,
16    fallback: F,
17}
18
19impl<P, F> FallbackJudgment<P, F> {
20    /// Composes `primary` with `fallback`.
21    #[must_use]
22    pub const fn new(primary: P, fallback: F) -> Self {
23        Self { primary, fallback }
24    }
25
26    /// The primary backend.
27    #[must_use]
28    pub const fn primary(&self) -> &P {
29        &self.primary
30    }
31
32    /// The fallback backend.
33    #[must_use]
34    pub const fn fallback(&self) -> &F {
35        &self.fallback
36    }
37}
38
39#[async_trait]
40impl<P, F> JudgmentProvider for FallbackJudgment<P, F>
41where
42    P: JudgmentProvider<Error = JudgmentError>,
43    F: JudgmentProvider<Error = JudgmentError>,
44{
45    type Error = JudgmentError;
46
47    async fn judge(&self, request: JudgmentRequest) -> Result<JudgmentResponse, JudgmentError> {
48        match self.primary.judge(request.clone()).await {
49            // `source` is set here, not trusted from the primary backend: a
50            // caller must be able to tell primary from fallback answers by
51            // `FallbackJudgment`'s own routing, not by whether the backend
52            // happened to fill the field correctly.
53            Ok(response) => Ok(JudgmentResponse {
54                source: JudgmentSource::Primary,
55                ..response
56            }),
57            // A rate limit, an outage, a broken transport, a refused
58            // credential, or an exhausted balance are all reasons the
59            // primary backend specifically cannot answer right now — none
60            // say the request itself is bad, so a different backend gets a
61            // real chance. An expired or invalid gateway credential must not
62            // silently take down the gate, so `Unauthorized` falls back here
63            // too, unlike the two arms below; an out-of-credit gateway is
64            // the same shape of failure from the gate's point of view, so
65            // `Exhausted` falls back too.
66            Err(
67                err @ (JudgmentError::RateLimited { .. }
68                | JudgmentError::Unavailable { .. }
69                | JudgmentError::Transport(_)
70                | JudgmentError::Unauthorized
71                | JudgmentError::Exhausted),
72            ) => {
73                tracing::warn!(
74                    kind = primary_error_kind(&err),
75                    "judgment primary backend failed; answering with the fallback backend instead"
76                );
77                self.fallback
78                    .judge(request)
79                    .await
80                    .map(|response| JudgmentResponse {
81                        source: JudgmentSource::Fallback,
82                        ..response
83                    })
84            }
85            // The request shape itself is rejected, or the primary's
86            // response didn't decode — a fallback backend would be asked the
87            // same rejected request or would see the same non-answer, so
88            // falling back cannot fix either.
89            Err(err @ (JudgmentError::Invalid(_) | JudgmentError::Malformed(_))) => Err(err),
90        }
91    }
92}
93
94/// The stable label for one [`JudgmentError`] variant, for the fallback log
95/// line. Never the variant's `Display`: `Transport`'s wraps a boxed transport
96/// error whose text is out of this crate's control.
97const fn primary_error_kind(error: &JudgmentError) -> &'static str {
98    match error {
99        JudgmentError::Invalid(_) => "invalid",
100        JudgmentError::Unauthorized => "unauthorized",
101        JudgmentError::RateLimited { .. } => "rate_limited",
102        JudgmentError::Unavailable { .. } => "unavailable",
103        JudgmentError::Exhausted => "exhausted",
104        JudgmentError::Transport(_) => "transport",
105        JudgmentError::Malformed(_) => "malformed",
106    }
107}
108
109#[cfg(test)]
110mod tests {
111    #![allow(clippy::pedantic, clippy::nursery, missing_docs)]
112
113    use std::sync::atomic::{AtomicUsize, Ordering};
114
115    use async_trait::async_trait;
116
117    use super::super::{
118        Answer, JudgmentError, JudgmentProvider, JudgmentRequest, JudgmentResponse, JudgmentSource,
119        JudgmentUsage,
120    };
121    use super::FallbackJudgment;
122
123    /// A scripted stub: answers with a fixed `Result`, counts its calls.
124    struct Scripted {
125        result: Result<f64, JudgmentError>,
126        calls: AtomicUsize,
127    }
128
129    impl Scripted {
130        fn ok(noul: f64) -> Self {
131            Self {
132                result: Ok(noul),
133                calls: AtomicUsize::new(0),
134            }
135        }
136
137        fn err(build: impl Fn() -> JudgmentError) -> Self {
138            Self {
139                result: Err(build()),
140                calls: AtomicUsize::new(0),
141            }
142        }
143
144        fn calls(&self) -> usize {
145            self.calls.load(Ordering::SeqCst)
146        }
147    }
148
149    #[async_trait]
150    impl JudgmentProvider for Scripted {
151        type Error = JudgmentError;
152
153        async fn judge(&self, _: JudgmentRequest) -> Result<JudgmentResponse, JudgmentError> {
154            self.calls.fetch_add(1, Ordering::SeqCst);
155            match &self.result {
156                Ok(noul) => {
157                    let mut answers = std::collections::BTreeMap::new();
158                    answers.insert("q".to_owned(), Answer::Noul { noul: *noul });
159                    Ok(JudgmentResponse {
160                        model: "scripted".to_owned(),
161                        answers,
162                        usage: JudgmentUsage::default(),
163                        source: JudgmentSource::Primary,
164                    })
165                }
166                Err(_) => Err(clone_err(&self.result)),
167            }
168        }
169    }
170
171    /// `JudgmentError` doesn't derive `Clone` (it wraps a boxed transport
172    /// error) — rebuild an equivalent value from the scripted variant instead.
173    fn clone_err(result: &Result<f64, JudgmentError>) -> JudgmentError {
174        match result {
175            Ok(_) => unreachable!("only called on the Err arm"),
176            Err(JudgmentError::Invalid(s)) => JudgmentError::Invalid(s.clone()),
177            Err(JudgmentError::Unauthorized) => JudgmentError::Unauthorized,
178            Err(JudgmentError::Exhausted) => JudgmentError::Exhausted,
179            Err(JudgmentError::RateLimited { retry_after }) => JudgmentError::RateLimited {
180                retry_after: *retry_after,
181            },
182            Err(JudgmentError::Unavailable { status }) => {
183                JudgmentError::Unavailable { status: *status }
184            }
185            Err(JudgmentError::Transport(_)) => {
186                JudgmentError::Transport(Box::new(std::io::Error::other("scripted transport")))
187            }
188            Err(JudgmentError::Malformed(s)) => JudgmentError::Malformed(s.clone()),
189        }
190    }
191
192    fn request() -> JudgmentRequest {
193        JudgmentRequest {
194            state: serde_json::json!({}),
195            questions: std::collections::BTreeMap::new(),
196        }
197    }
198
199    #[tokio::test]
200    async fn primary_success_never_calls_the_fallback() {
201        let primary = Scripted::ok(0.8);
202        let fallback = Scripted::ok(0.1);
203        let gate = FallbackJudgment::new(primary, fallback);
204        let response = gate.judge(request()).await.expect("primary answered");
205        assert_eq!(response.noul("q"), Some(0.8));
206        assert_eq!(response.source, JudgmentSource::Primary);
207        assert_eq!(gate.primary().calls(), 1);
208        assert_eq!(gate.fallback().calls(), 0);
209    }
210
211    #[tokio::test]
212    async fn primary_rate_limited_falls_back() {
213        let primary = Scripted::err(|| JudgmentError::RateLimited { retry_after: None });
214        let fallback = Scripted::ok(0.6);
215        let gate = FallbackJudgment::new(primary, fallback);
216        let response = gate.judge(request()).await.expect("fallback answered");
217        assert_eq!(response.noul("q"), Some(0.6));
218        assert_eq!(response.source, JudgmentSource::Fallback);
219        assert_eq!(gate.primary().calls(), 1);
220        assert_eq!(gate.fallback().calls(), 1);
221    }
222
223    #[tokio::test]
224    async fn primary_unavailable_falls_back() {
225        let primary = Scripted::err(|| JudgmentError::Unavailable { status: 503 });
226        let fallback = Scripted::ok(0.6);
227        let gate = FallbackJudgment::new(primary, fallback);
228        let response = gate.judge(request()).await.expect("fallback answered");
229        assert_eq!(response.source, JudgmentSource::Fallback);
230        assert_eq!(gate.fallback().calls(), 1);
231    }
232
233    #[tokio::test]
234    async fn primary_transport_failure_falls_back() {
235        let primary = Scripted::err(|| {
236            JudgmentError::Transport(Box::new(std::io::Error::other("connect refused")))
237        });
238        let fallback = Scripted::ok(0.6);
239        let gate = FallbackJudgment::new(primary, fallback);
240        let response = gate.judge(request()).await.expect("fallback answered");
241        assert_eq!(response.source, JudgmentSource::Fallback);
242        assert_eq!(gate.fallback().calls(), 1);
243    }
244
245    /// An expired or invalid gateway credential must not silently take down
246    /// the gate, so `Unauthorized` also triggers fallback.
247    #[tokio::test]
248    async fn primary_unauthorized_falls_back() {
249        let primary = Scripted::err(|| JudgmentError::Unauthorized);
250        let fallback = Scripted::ok(0.6);
251        let gate = FallbackJudgment::new(primary, fallback);
252        let response = gate.judge(request()).await.expect("fallback answered");
253        assert_eq!(response.source, JudgmentSource::Fallback);
254        assert_eq!(gate.fallback().calls(), 1);
255    }
256
257    /// An exhausted gateway balance is an availability failure from the
258    /// gate's point of view, so it falls back like an outage.
259    #[tokio::test]
260    async fn primary_exhausted_falls_back() {
261        let primary = Scripted::err(|| JudgmentError::Exhausted);
262        let fallback = Scripted::ok(0.6);
263        let gate = FallbackJudgment::new(primary, fallback);
264        let response = gate.judge(request()).await.expect("fallback answered");
265        assert_eq!(response.source, JudgmentSource::Fallback);
266        assert_eq!(gate.fallback().calls(), 1);
267    }
268
269    /// A bad-request shape is never fixed by a different backend.
270    #[tokio::test]
271    async fn primary_invalid_returns_as_is_without_a_fallback_call() {
272        let primary = Scripted::err(|| JudgmentError::Invalid("bad state shape".to_owned()));
273        let fallback = Scripted::ok(0.6);
274        let gate = FallbackJudgment::new(primary, fallback);
275        let err = gate
276            .judge(request())
277            .await
278            .expect_err("no fallback for Invalid");
279        assert!(matches!(err, JudgmentError::Invalid(_)), "{err:?}");
280        assert_eq!(gate.fallback().calls(), 0);
281    }
282
283    #[tokio::test]
284    async fn primary_malformed_returns_as_is_without_a_fallback_call() {
285        let primary = Scripted::err(|| JudgmentError::Malformed("not json".to_owned()));
286        let fallback = Scripted::ok(0.6);
287        let gate = FallbackJudgment::new(primary, fallback);
288        let err = gate
289            .judge(request())
290            .await
291            .expect_err("no fallback for Malformed");
292        assert!(matches!(err, JudgmentError::Malformed(_)), "{err:?}");
293        assert_eq!(gate.fallback().calls(), 0);
294    }
295}