use crate::ai::gemini::{GeminiClient, GenerateContentRequest};
use crate::ai::quota::QuotaManager;
use crate::ai::{AiErrorClass, classify_ai_error};
use axum::{
extract::{Json, State},
http::StatusCode,
response::IntoResponse,
};
use std::sync::Arc;
use tracing::error;
pub struct ProxyState {
pub client: Arc<GeminiClient>,
pub quota_manager: Arc<QuotaManager>,
}
pub async fn handle_generate(
State(state): State<Arc<ProxyState>>,
Json(request): Json<GenerateContentRequest>,
) -> impl IntoResponse {
let mut local_transient_errors = 0;
loop {
let _slept = state.quota_manager.wait_for_access().await;
match state.client.generate_content_single(&request).await {
Ok(response) => {
state.quota_manager.report_success().await;
return (StatusCode::OK, Json(response)).into_response();
}
Err(e) => {
match classify_ai_error(&e) {
AiErrorClass::RateLimit { retry_after } => {
state.quota_manager.report_quota_error(retry_after).await;
continue;
}
AiErrorClass::Transient { retry_after } => {
local_transient_errors += 1;
let backoff_secs =
(1.0 * (2.0_f64.powi(local_transient_errors - 1))).min(60.0);
let backoff =
std::time::Duration::from_secs_f64(backoff_secs).max(retry_after);
tracing::warn!(
"AI provider transient error (streak: {}). Locally backing off for {:.2}s",
local_transient_errors,
backoff.as_secs_f64()
);
tokio::time::sleep(backoff).await;
continue;
}
AiErrorClass::Fatal => {}
}
error!("Gemini Proxy Error: {}", e);
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": format!("{:#}", e)})),
)
.into_response();
}
}
}
}