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 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 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 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 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 #[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 #[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 #[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}