1#[cfg(feature = "server")]
4use axum::http::StatusCode;
5#[cfg(feature = "server")]
6use axum::response::{IntoResponse, Response};
7#[cfg(feature = "server")]
8use axum::Json;
9#[cfg(feature = "server")]
10use serde_json::json;
11use thiserror::Error;
12
13#[derive(Debug, Error)]
14pub enum GatewayError {
15 #[error("unknown model alias '{0}'")]
16 UnknownModel(String),
17 #[error("native feature '{feature}' is not available on route '{route}'")]
18 NativeFeatureUnsupported { feature: String, route: String },
19 #[error("invalid request: {0}")]
20 BadRequest(String),
21 #[error("all legs of route '{route}' failed")]
22 AllLegsFailed {
23 route: String,
24 failures: Vec<LegFailure>,
25 },
26 #[error("all providers for route '{0}' are unavailable")]
27 AllCircuitsOpen(String),
28 #[error("upstream timed out")]
29 UpstreamTimeout,
30 #[error("upstream error {status}: {body}")]
31 Upstream { status: u16, body: String },
32 #[error("request blocked by content policy '{policy}'")]
33 ContentBlocked {
34 policy: String,
35 scanners: Vec<String>,
36 },
37}
38
39#[derive(Debug, Clone, serde::Serialize)]
40pub struct LegFailure {
41 pub provider: String,
42 pub model: String,
43 pub message: String,
44}
45
46impl GatewayError {
47 #[cfg(feature = "server")]
48 pub fn status(&self) -> StatusCode {
49 match self {
50 GatewayError::UnknownModel(_) => StatusCode::NOT_FOUND,
51 GatewayError::NativeFeatureUnsupported { .. } | GatewayError::BadRequest(_) => {
52 StatusCode::BAD_REQUEST
53 }
54 GatewayError::AllLegsFailed { .. } | GatewayError::Upstream { .. } => {
55 StatusCode::BAD_GATEWAY
56 }
57 GatewayError::AllCircuitsOpen(_) => StatusCode::SERVICE_UNAVAILABLE,
58 GatewayError::UpstreamTimeout => StatusCode::GATEWAY_TIMEOUT,
59 GatewayError::ContentBlocked { .. } => StatusCode::BAD_REQUEST,
60 }
61 }
62
63 #[cfg(feature = "server")]
64 pub(crate) fn code(&self) -> &'static str {
65 match self {
66 GatewayError::UnknownModel(_) => "model_not_found",
67 GatewayError::NativeFeatureUnsupported { .. } => "native_feature_unsupported",
68 GatewayError::BadRequest(_) => "invalid_request_error",
69 GatewayError::AllLegsFailed { .. } => "all_legs_failed",
70 GatewayError::AllCircuitsOpen(_) => "circuit_open",
71 GatewayError::UpstreamTimeout => "upstream_timeout",
72 GatewayError::Upstream { .. } => "upstream_error",
73 GatewayError::ContentBlocked { .. } => "content_blocked",
74 }
75 }
76}
77
78#[cfg(feature = "server")]
79impl IntoResponse for GatewayError {
80 fn into_response(self) -> Response {
81 let mut error = json!({
82 "type": self.code(),
83 "message": self.to_string(),
84 "code": self.code(),
85 });
86 if let GatewayError::AllLegsFailed { failures, .. } = &self {
87 error["failures"] = json!(failures);
88 }
89 if let GatewayError::ContentBlocked { policy, scanners } = &self {
90 error["type"] = json!("content_policy_violation");
91 error["message"] = json!(format!(
92 "Request blocked by content policy '{policy}' (scanners: {})",
93 scanners.join(", ")
94 ));
95 error["scanners"] = json!(scanners);
96 }
97 (self.status(), Json(json!({ "error": error }))).into_response()
98 }
99}
100
101#[cfg(all(test, feature = "server"))]
102mod tests {
103 use super::*;
104 use axum::http::StatusCode;
105
106 #[test]
107 fn content_blocked_is_bad_request_with_named_scanners() {
108 let e = GatewayError::ContentBlocked {
109 policy: "strict".into(),
110 scanners: vec!["secrets".into(), "role_override".into()],
111 };
112 assert_eq!(e.status(), StatusCode::BAD_REQUEST);
113 assert_eq!(e.code(), "content_blocked");
114 let msg = e.to_string();
115 assert!(msg.contains("strict"), "message names the policy: {msg}");
116 }
117}