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                JudgmentError::RateLimited { .. }
68                | JudgmentError::Unavailable { .. }
69                | JudgmentError::Transport(_)
70                | JudgmentError::Unauthorized
71                | JudgmentError::Exhausted,
72            ) => self
73                .fallback
74                .judge(request)
75                .await
76                .map(|response| JudgmentResponse {
77                    source: JudgmentSource::Fallback,
78                    ..response
79                }),
80            // The request shape itself is rejected, or the primary's
81            // response didn't decode — a fallback backend would be asked the
82            // same rejected request or would see the same non-answer, so
83            // falling back cannot fix either.
84            Err(err @ (JudgmentError::Invalid(_) | JudgmentError::Malformed(_))) => Err(err),
85        }
86    }
87}
88
89#[cfg(test)]
90mod tests {
91    #![allow(clippy::pedantic, clippy::nursery, missing_docs)]
92
93    use std::sync::atomic::{AtomicUsize, Ordering};
94
95    use async_trait::async_trait;
96
97    use super::super::{
98        Answer, JudgmentError, JudgmentProvider, JudgmentRequest, JudgmentResponse, JudgmentSource,
99        JudgmentUsage,
100    };
101    use super::FallbackJudgment;
102
103    /// A scripted stub: answers with a fixed `Result`, counts its calls.
104    struct Scripted {
105        result: Result<f64, JudgmentError>,
106        calls: AtomicUsize,
107    }
108
109    impl Scripted {
110        fn ok(noul: f64) -> Self {
111            Self {
112                result: Ok(noul),
113                calls: AtomicUsize::new(0),
114            }
115        }
116
117        fn err(build: impl Fn() -> JudgmentError) -> Self {
118            Self {
119                result: Err(build()),
120                calls: AtomicUsize::new(0),
121            }
122        }
123
124        fn calls(&self) -> usize {
125            self.calls.load(Ordering::SeqCst)
126        }
127    }
128
129    #[async_trait]
130    impl JudgmentProvider for Scripted {
131        type Error = JudgmentError;
132
133        async fn judge(&self, _: JudgmentRequest) -> Result<JudgmentResponse, JudgmentError> {
134            self.calls.fetch_add(1, Ordering::SeqCst);
135            match &self.result {
136                Ok(noul) => {
137                    let mut answers = std::collections::BTreeMap::new();
138                    answers.insert("q".to_owned(), Answer::Noul { noul: *noul });
139                    Ok(JudgmentResponse {
140                        model: "scripted".to_owned(),
141                        answers,
142                        usage: JudgmentUsage::default(),
143                        source: JudgmentSource::Primary,
144                    })
145                }
146                Err(_) => Err(clone_err(&self.result)),
147            }
148        }
149    }
150
151    /// `JudgmentError` doesn't derive `Clone` (it wraps a boxed transport
152    /// error) — rebuild an equivalent value from the scripted variant instead.
153    fn clone_err(result: &Result<f64, JudgmentError>) -> JudgmentError {
154        match result {
155            Ok(_) => unreachable!("only called on the Err arm"),
156            Err(JudgmentError::Invalid(s)) => JudgmentError::Invalid(s.clone()),
157            Err(JudgmentError::Unauthorized) => JudgmentError::Unauthorized,
158            Err(JudgmentError::Exhausted) => JudgmentError::Exhausted,
159            Err(JudgmentError::RateLimited { retry_after }) => JudgmentError::RateLimited {
160                retry_after: *retry_after,
161            },
162            Err(JudgmentError::Unavailable { status }) => {
163                JudgmentError::Unavailable { status: *status }
164            }
165            Err(JudgmentError::Transport(_)) => {
166                JudgmentError::Transport(Box::new(std::io::Error::other("scripted transport")))
167            }
168            Err(JudgmentError::Malformed(s)) => JudgmentError::Malformed(s.clone()),
169        }
170    }
171
172    fn request() -> JudgmentRequest {
173        JudgmentRequest {
174            state: serde_json::json!({}),
175            questions: std::collections::BTreeMap::new(),
176        }
177    }
178
179    #[tokio::test]
180    async fn primary_success_never_calls_the_fallback() {
181        let primary = Scripted::ok(0.8);
182        let fallback = Scripted::ok(0.1);
183        let gate = FallbackJudgment::new(primary, fallback);
184        let response = gate.judge(request()).await.expect("primary answered");
185        assert_eq!(response.noul("q"), Some(0.8));
186        assert_eq!(response.source, JudgmentSource::Primary);
187        assert_eq!(gate.primary().calls(), 1);
188        assert_eq!(gate.fallback().calls(), 0);
189    }
190
191    #[tokio::test]
192    async fn primary_rate_limited_falls_back() {
193        let primary = Scripted::err(|| JudgmentError::RateLimited { retry_after: None });
194        let fallback = Scripted::ok(0.6);
195        let gate = FallbackJudgment::new(primary, fallback);
196        let response = gate.judge(request()).await.expect("fallback answered");
197        assert_eq!(response.noul("q"), Some(0.6));
198        assert_eq!(response.source, JudgmentSource::Fallback);
199        assert_eq!(gate.primary().calls(), 1);
200        assert_eq!(gate.fallback().calls(), 1);
201    }
202
203    #[tokio::test]
204    async fn primary_unavailable_falls_back() {
205        let primary = Scripted::err(|| JudgmentError::Unavailable { status: 503 });
206        let fallback = Scripted::ok(0.6);
207        let gate = FallbackJudgment::new(primary, fallback);
208        let response = gate.judge(request()).await.expect("fallback answered");
209        assert_eq!(response.source, JudgmentSource::Fallback);
210        assert_eq!(gate.fallback().calls(), 1);
211    }
212
213    #[tokio::test]
214    async fn primary_transport_failure_falls_back() {
215        let primary = Scripted::err(|| {
216            JudgmentError::Transport(Box::new(std::io::Error::other("connect refused")))
217        });
218        let fallback = Scripted::ok(0.6);
219        let gate = FallbackJudgment::new(primary, fallback);
220        let response = gate.judge(request()).await.expect("fallback answered");
221        assert_eq!(response.source, JudgmentSource::Fallback);
222        assert_eq!(gate.fallback().calls(), 1);
223    }
224
225    /// An expired or invalid gateway credential must not silently take down
226    /// the gate, so `Unauthorized` also triggers fallback.
227    #[tokio::test]
228    async fn primary_unauthorized_falls_back() {
229        let primary = Scripted::err(|| JudgmentError::Unauthorized);
230        let fallback = Scripted::ok(0.6);
231        let gate = FallbackJudgment::new(primary, fallback);
232        let response = gate.judge(request()).await.expect("fallback answered");
233        assert_eq!(response.source, JudgmentSource::Fallback);
234        assert_eq!(gate.fallback().calls(), 1);
235    }
236
237    /// An exhausted gateway balance is an availability failure from the
238    /// gate's point of view, so it falls back like an outage.
239    #[tokio::test]
240    async fn primary_exhausted_falls_back() {
241        let primary = Scripted::err(|| JudgmentError::Exhausted);
242        let fallback = Scripted::ok(0.6);
243        let gate = FallbackJudgment::new(primary, fallback);
244        let response = gate.judge(request()).await.expect("fallback answered");
245        assert_eq!(response.source, JudgmentSource::Fallback);
246        assert_eq!(gate.fallback().calls(), 1);
247    }
248
249    /// A bad-request shape is never fixed by a different backend.
250    #[tokio::test]
251    async fn primary_invalid_returns_as_is_without_a_fallback_call() {
252        let primary = Scripted::err(|| JudgmentError::Invalid("bad state shape".to_owned()));
253        let fallback = Scripted::ok(0.6);
254        let gate = FallbackJudgment::new(primary, fallback);
255        let err = gate
256            .judge(request())
257            .await
258            .expect_err("no fallback for Invalid");
259        assert!(matches!(err, JudgmentError::Invalid(_)), "{err:?}");
260        assert_eq!(gate.fallback().calls(), 0);
261    }
262
263    #[tokio::test]
264    async fn primary_malformed_returns_as_is_without_a_fallback_call() {
265        let primary = Scripted::err(|| JudgmentError::Malformed("not json".to_owned()));
266        let fallback = Scripted::ok(0.6);
267        let gate = FallbackJudgment::new(primary, fallback);
268        let err = gate
269            .judge(request())
270            .await
271            .expect_err("no fallback for Malformed");
272        assert!(matches!(err, JudgmentError::Malformed(_)), "{err:?}");
273        assert_eq!(gate.fallback().calls(), 0);
274    }
275}