use async_trait::async_trait;
use crate::{JudgmentError, JudgmentProvider, JudgmentRequest, JudgmentResponse, JudgmentSource};
pub struct FallbackJudgment<P, F> {
primary: P,
fallback: F,
}
impl<P, F> FallbackJudgment<P, F> {
#[must_use]
pub const fn new(primary: P, fallback: F) -> Self {
Self { primary, fallback }
}
#[must_use]
pub const fn primary(&self) -> &P {
&self.primary
}
#[must_use]
pub const fn fallback(&self) -> &F {
&self.fallback
}
}
#[async_trait]
impl<P, F> JudgmentProvider for FallbackJudgment<P, F>
where
P: JudgmentProvider<Error = JudgmentError>,
F: JudgmentProvider<Error = JudgmentError>,
{
type Error = JudgmentError;
async fn judge(&self, request: JudgmentRequest) -> Result<JudgmentResponse, JudgmentError> {
match self.primary.judge(request.clone()).await {
Ok(response) => Ok(JudgmentResponse {
source: JudgmentSource::Primary,
..response
}),
Err(
err @ (JudgmentError::RateLimited { .. }
| JudgmentError::Unavailable { .. }
| JudgmentError::Transport(_)
| JudgmentError::Unauthorized
| JudgmentError::Exhausted),
) => {
tracing::warn!(
kind = primary_error_kind(&err),
"judgment primary backend failed; answering with the fallback backend instead"
);
self.fallback
.judge(request)
.await
.map(|response| JudgmentResponse {
source: JudgmentSource::Fallback,
..response
})
}
Err(err @ (JudgmentError::Invalid(_) | JudgmentError::Malformed(_))) => Err(err),
}
}
}
const fn primary_error_kind(error: &JudgmentError) -> &'static str {
match error {
JudgmentError::Invalid(_) => "invalid",
JudgmentError::Unauthorized => "unauthorized",
JudgmentError::RateLimited { .. } => "rate_limited",
JudgmentError::Unavailable { .. } => "unavailable",
JudgmentError::Exhausted => "exhausted",
JudgmentError::Transport(_) => "transport",
JudgmentError::Malformed(_) => "malformed",
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::pedantic, clippy::nursery, missing_docs)]
use std::sync::atomic::{AtomicUsize, Ordering};
use async_trait::async_trait;
use super::super::{
Answer, JudgmentError, JudgmentProvider, JudgmentRequest, JudgmentResponse, JudgmentSource,
JudgmentUsage,
};
use super::FallbackJudgment;
struct Scripted {
result: Result<f64, JudgmentError>,
calls: AtomicUsize,
}
impl Scripted {
fn ok(noul: f64) -> Self {
Self {
result: Ok(noul),
calls: AtomicUsize::new(0),
}
}
fn err(build: impl Fn() -> JudgmentError) -> Self {
Self {
result: Err(build()),
calls: AtomicUsize::new(0),
}
}
fn calls(&self) -> usize {
self.calls.load(Ordering::SeqCst)
}
}
#[async_trait]
impl JudgmentProvider for Scripted {
type Error = JudgmentError;
async fn judge(&self, _: JudgmentRequest) -> Result<JudgmentResponse, JudgmentError> {
self.calls.fetch_add(1, Ordering::SeqCst);
match &self.result {
Ok(noul) => {
let mut answers = std::collections::BTreeMap::new();
answers.insert("q".to_owned(), Answer::Noul { noul: *noul });
Ok(JudgmentResponse {
model: "scripted".to_owned(),
answers,
usage: JudgmentUsage::default(),
source: JudgmentSource::Primary,
})
}
Err(_) => Err(clone_err(&self.result)),
}
}
}
fn clone_err(result: &Result<f64, JudgmentError>) -> JudgmentError {
match result {
Ok(_) => unreachable!("only called on the Err arm"),
Err(JudgmentError::Invalid(s)) => JudgmentError::Invalid(s.clone()),
Err(JudgmentError::Unauthorized) => JudgmentError::Unauthorized,
Err(JudgmentError::Exhausted) => JudgmentError::Exhausted,
Err(JudgmentError::RateLimited { retry_after }) => JudgmentError::RateLimited {
retry_after: *retry_after,
},
Err(JudgmentError::Unavailable { status }) => {
JudgmentError::Unavailable { status: *status }
}
Err(JudgmentError::Transport(_)) => {
JudgmentError::Transport(Box::new(std::io::Error::other("scripted transport")))
}
Err(JudgmentError::Malformed(s)) => JudgmentError::Malformed(s.clone()),
}
}
fn request() -> JudgmentRequest {
JudgmentRequest {
state: serde_json::json!({}),
questions: std::collections::BTreeMap::new(),
}
}
#[tokio::test]
async fn primary_success_never_calls_the_fallback() {
let primary = Scripted::ok(0.8);
let fallback = Scripted::ok(0.1);
let gate = FallbackJudgment::new(primary, fallback);
let response = gate.judge(request()).await.expect("primary answered");
assert_eq!(response.noul("q"), Some(0.8));
assert_eq!(response.source, JudgmentSource::Primary);
assert_eq!(gate.primary().calls(), 1);
assert_eq!(gate.fallback().calls(), 0);
}
#[tokio::test]
async fn primary_rate_limited_falls_back() {
let primary = Scripted::err(|| JudgmentError::RateLimited { retry_after: None });
let fallback = Scripted::ok(0.6);
let gate = FallbackJudgment::new(primary, fallback);
let response = gate.judge(request()).await.expect("fallback answered");
assert_eq!(response.noul("q"), Some(0.6));
assert_eq!(response.source, JudgmentSource::Fallback);
assert_eq!(gate.primary().calls(), 1);
assert_eq!(gate.fallback().calls(), 1);
}
#[tokio::test]
async fn primary_unavailable_falls_back() {
let primary = Scripted::err(|| JudgmentError::Unavailable { status: 503 });
let fallback = Scripted::ok(0.6);
let gate = FallbackJudgment::new(primary, fallback);
let response = gate.judge(request()).await.expect("fallback answered");
assert_eq!(response.source, JudgmentSource::Fallback);
assert_eq!(gate.fallback().calls(), 1);
}
#[tokio::test]
async fn primary_transport_failure_falls_back() {
let primary = Scripted::err(|| {
JudgmentError::Transport(Box::new(std::io::Error::other("connect refused")))
});
let fallback = Scripted::ok(0.6);
let gate = FallbackJudgment::new(primary, fallback);
let response = gate.judge(request()).await.expect("fallback answered");
assert_eq!(response.source, JudgmentSource::Fallback);
assert_eq!(gate.fallback().calls(), 1);
}
#[tokio::test]
async fn primary_unauthorized_falls_back() {
let primary = Scripted::err(|| JudgmentError::Unauthorized);
let fallback = Scripted::ok(0.6);
let gate = FallbackJudgment::new(primary, fallback);
let response = gate.judge(request()).await.expect("fallback answered");
assert_eq!(response.source, JudgmentSource::Fallback);
assert_eq!(gate.fallback().calls(), 1);
}
#[tokio::test]
async fn primary_exhausted_falls_back() {
let primary = Scripted::err(|| JudgmentError::Exhausted);
let fallback = Scripted::ok(0.6);
let gate = FallbackJudgment::new(primary, fallback);
let response = gate.judge(request()).await.expect("fallback answered");
assert_eq!(response.source, JudgmentSource::Fallback);
assert_eq!(gate.fallback().calls(), 1);
}
#[tokio::test]
async fn primary_invalid_returns_as_is_without_a_fallback_call() {
let primary = Scripted::err(|| JudgmentError::Invalid("bad state shape".to_owned()));
let fallback = Scripted::ok(0.6);
let gate = FallbackJudgment::new(primary, fallback);
let err = gate
.judge(request())
.await
.expect_err("no fallback for Invalid");
assert!(matches!(err, JudgmentError::Invalid(_)), "{err:?}");
assert_eq!(gate.fallback().calls(), 0);
}
#[tokio::test]
async fn primary_malformed_returns_as_is_without_a_fallback_call() {
let primary = Scripted::err(|| JudgmentError::Malformed("not json".to_owned()));
let fallback = Scripted::ok(0.6);
let gate = FallbackJudgment::new(primary, fallback);
let err = gate
.judge(request())
.await
.expect_err("no fallback for Malformed");
assert!(matches!(err, JudgmentError::Malformed(_)), "{err:?}");
assert_eq!(gate.fallback().calls(), 0);
}
}