1use async_trait::async_trait;
5
6use crate::{JudgmentError, JudgmentProvider, JudgmentRequest, JudgmentResponse, JudgmentSource};
7
8pub struct FallbackJudgment<P, F> {
15 primary: P,
16 fallback: F,
17}
18
19impl<P, F> FallbackJudgment<P, F> {
20 #[must_use]
22 pub const fn new(primary: P, fallback: F) -> Self {
23 Self { primary, fallback }
24 }
25
26 #[must_use]
28 pub const fn primary(&self) -> &P {
29 &self.primary
30 }
31
32 #[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 Ok(response) => Ok(JudgmentResponse {
54 source: JudgmentSource::Primary,
55 ..response
56 }),
57 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 Err(err @ (JudgmentError::Invalid(_) | JudgmentError::Malformed(_))) => Err(err),
90 }
91 }
92}
93
94const 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 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 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 #[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 #[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 #[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}