#![allow(clippy::unused_async)]
use axum::extract::{Path, State};
use axum::http::{HeaderMap, StatusCode};
use axum::response::IntoResponse;
use crate::login::{LoginError, LoginManager};
use crate::proxy::{AppState, error_response, is_admin_authorised};
use crate::subscription::SubscriptionProvider;
pub async fn begin_login(
State(state): State<AppState>,
headers: HeaderMap,
request: Option<axum::Json<BeginLoginRequest>>,
) -> impl IntoResponse {
let Some(manager) = guard(&state, &headers) else {
return unauthorised();
};
let provider = requested_provider(request.as_ref().and_then(|body| body.provider.as_deref()));
let Some(provider) = provider else {
return error_response(
StatusCode::BAD_REQUEST,
"invalid_request_error",
"`provider` must be `claude` or `codex`",
);
};
match manager.begin_for(provider).await {
Ok(view) => (StatusCode::OK, axum::Json(view)).into_response(),
Err(e) => login_error_response(&e),
}
}
#[derive(Default, serde::Deserialize)]
pub struct BeginLoginRequest {
pub provider: Option<String>,
}
fn requested_provider(value: Option<&str>) -> Option<SubscriptionProvider> {
match value
.unwrap_or("claude")
.trim()
.to_ascii_lowercase()
.as_str()
{
"claude" => Some(SubscriptionProvider::Claude),
"codex" => Some(SubscriptionProvider::Codex),
_ => None,
}
}
pub async fn login_status(
State(state): State<AppState>,
headers: HeaderMap,
Path(id): Path<String>,
) -> impl IntoResponse {
let Some(manager) = guard(&state, &headers) else {
return unauthorised();
};
manager.status(&id).map_or_else(
|| login_error_response(&LoginError::NotFound),
|view| (StatusCode::OK, axum::Json(view)).into_response(),
)
}
pub async fn submit_code(
State(state): State<AppState>,
headers: HeaderMap,
Path(id): Path<String>,
axum::Json(req): axum::Json<SubmitCodeRequest>,
) -> impl IntoResponse {
let Some(manager) = guard(&state, &headers) else {
return unauthorised();
};
if req.code.trim().is_empty() {
return error_response(
StatusCode::BAD_REQUEST,
"invalid_request_error",
"`code` must not be empty",
);
}
match manager.submit_code(&id, &req.code).await {
Ok(view) => {
if view.status == crate::login::LoginStatus::Authorized {
let _ = state.oauth_provider.refresh_token();
}
(StatusCode::OK, axum::Json(view)).into_response()
}
Err(e) => login_error_response(&e),
}
}
pub async fn cancel_login(
State(state): State<AppState>,
headers: HeaderMap,
Path(id): Path<String>,
) -> impl IntoResponse {
let Some(manager) = guard(&state, &headers) else {
return unauthorised();
};
if manager.cancel(&id) {
(
StatusCode::OK,
axum::Json(serde_json::json!({"cancelled": id})),
)
.into_response()
} else {
login_error_response(&LoginError::NotFound)
}
}
#[derive(serde::Deserialize)]
pub struct SubmitCodeRequest {
pub code: String,
}
fn guard<'a>(state: &'a AppState, headers: &HeaderMap) -> Option<&'a LoginManager> {
is_admin_authorised(state, headers).then_some(&state.login_manager)
}
fn unauthorised() -> axum::response::Response {
error_response(
StatusCode::UNAUTHORIZED,
"authentication_error",
"admin Bearer key required",
)
}
fn login_error_response(error: &LoginError) -> axum::response::Response {
let (status, kind) = match error {
LoginError::Disabled | LoginError::NotFound => (StatusCode::NOT_FOUND, "not_found_error"),
LoginError::TooManySessions(_) => (StatusCode::TOO_MANY_REQUESTS, "rate_limit_error"),
LoginError::NotPending(_) => (StatusCode::CONFLICT, "invalid_request_error"),
LoginError::Spawn(_) | LoginError::NoUrl(_) | LoginError::Storage(_) => {
(StatusCode::BAD_GATEWAY, "api_error")
}
};
error_response(status, kind, &error.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn provider_defaults_to_claude_and_rejects_unexposed_aliases() {
assert_eq!(requested_provider(None), Some(SubscriptionProvider::Claude));
assert_eq!(
requested_provider(Some("codex")),
Some(SubscriptionProvider::Codex)
);
assert_eq!(requested_provider(Some("chatgpt")), None);
assert_eq!(requested_provider(Some("gemini")), None);
}
}