use axum::Json;
use axum::extract::{Query, State};
use axum::http::HeaderMap;
use axum::response::{IntoResponse, Redirect, Response};
use serde::{Deserialize, Serialize};
use crate::routes::oauth::extractors::OAuthRepo;
use crate::routes::oauth::{OAuthHttpError, internal};
use crate::services::request_base_url::RequestBaseUrl;
use systemprompt_identifiers::{ClientId, UserId};
use systemprompt_models::oauth::OAuthServerConfig;
use systemprompt_oauth::OAuthState;
use systemprompt_oauth::repository::{MintAuthCodeParams, OAuthRepository};
use systemprompt_oauth::services::is_browser_request;
use systemprompt_oauth::services::validation::validate_redirect_uri;
#[derive(Debug, Deserialize)]
pub struct WebAuthnCompleteQuery {
pub user_id: UserId,
pub auth_token: Option<String>,
pub response_type: Option<String>,
pub client_id: Option<ClientId>,
pub redirect_uri: Option<String>,
pub scope: Option<String>,
pub state: Option<String>,
pub code_challenge: Option<String>,
pub code_challenge_method: Option<String>,
pub response_mode: Option<String>,
pub resource: Option<String>,
}
async fn verify_completion(
params: &WebAuthnCompleteQuery,
state: &OAuthState,
repo: &OAuthRepository,
) -> Result<(UserId, String), OAuthHttpError> {
let auth_token = params
.auth_token
.as_deref()
.ok_or_else(|| OAuthHttpError::invalid_request("Missing auth_token parameter"))?;
let webauthn_service = state.webauthn()?;
let verified_user_id = webauthn_service
.consume_verified_authentication(auth_token)
.await
.map_err(|_e| OAuthHttpError::access_denied("Invalid or expired authentication token"))?;
if params.user_id != verified_user_id {
return Err(OAuthHttpError::access_denied(
"User identity verification failed",
));
}
let client_id = params
.client_id
.as_ref()
.ok_or_else(|| OAuthHttpError::invalid_request("Missing client_id parameter"))?;
let client = repo
.find_client_by_id(client_id)
.await?
.ok_or_else(|| OAuthHttpError::invalid_request("Unknown client_id"))?;
let redirect_uri = validate_redirect_uri(&client.redirect_uris, params.redirect_uri.as_deref())
.map_err(|_e| OAuthHttpError::invalid_request("redirect_uri not registered for client"))?;
let requested_scopes = OAuthRepository::parse_scopes(params.scope.as_deref().unwrap_or(""));
OAuthRepository::validate_scopes(&requested_scopes)
.and_then(|_valid| {
OAuthRepository::validate_scopes_for_client(&client.scopes, &requested_scopes)
})
.map_err(|e| internal::classify_validation(e, OAuthHttpError::invalid_scope))?;
let has_challenge = params
.code_challenge
.as_deref()
.is_some_and(|challenge| !challenge.is_empty());
if !has_challenge || params.code_challenge_method.as_deref() != Some("S256") {
return Err(OAuthHttpError::invalid_request(
"PKCE code_challenge with method S256 is required",
));
}
Ok((verified_user_id, redirect_uri))
}
pub async fn handle_webauthn_complete(
headers: HeaderMap,
base: RequestBaseUrl,
Query(params): Query<WebAuthnCompleteQuery>,
State(state): State<OAuthState>,
OAuthRepo(repo): OAuthRepo,
) -> Result<Response, OAuthHttpError> {
let (verified_user_id, redirect_uri) = verify_completion(¶ms, &state, &repo).await?;
let user = state.user_provider().find_by_id(&verified_user_id).await?;
if user.is_none() {
return Err(OAuthHttpError::access_denied("User not found"));
}
let client_id = params
.client_id
.as_ref()
.ok_or_else(|| OAuthHttpError::invalid_request("client_id is required"))?;
let authorization_code = repo
.mint_authorization_code(MintAuthCodeParams {
client_id,
user_id: ¶ms.user_id,
redirect_uri: &redirect_uri,
scope: params.scope.as_deref(),
code_challenge: params.code_challenge.as_deref().unwrap_or_default(),
code_challenge_method: params.code_challenge_method.as_deref().unwrap_or_default(),
resource: params.resource.as_deref(),
})
.await?;
let issuer = OAuthServerConfig::from_api_server_url(base.as_str()).issuer;
Ok(create_successful_response(
&headers,
&CompletedAuthorization {
redirect_uri: &redirect_uri,
authorization_code: authorization_code.as_str(),
client_id,
state: params.state.as_deref(),
},
&issuer,
))
}
#[derive(Debug, Serialize)]
pub struct WebAuthnCompleteResponse {
pub authorization_code: String,
pub state: String,
pub redirect_uri: String,
pub client_id: ClientId,
}
struct CompletedAuthorization<'a> {
redirect_uri: &'a str,
authorization_code: &'a str,
client_id: &'a ClientId,
state: Option<&'a str>,
}
fn create_successful_response(
headers: &HeaderMap,
completed: &CompletedAuthorization<'_>,
issuer: &str,
) -> Response {
let CompletedAuthorization {
redirect_uri,
authorization_code,
client_id,
state,
} = *completed;
let state = state.filter(|s| !s.is_empty());
if is_browser_request(headers) {
let mut target = format!(
"{redirect_uri}?code={authorization_code}&client_id={}",
urlencoding::encode(client_id.as_str())
);
if let Some(state_val) = state {
target.push_str(&format!("&state={}", urlencoding::encode(state_val)));
}
target.push_str(&format!("&iss={}", urlencoding::encode(issuer)));
Redirect::to(&target).into_response()
} else {
let response_data = WebAuthnCompleteResponse {
authorization_code: authorization_code.to_owned(),
state: state.unwrap_or("").to_owned(),
redirect_uri: redirect_uri.to_owned(),
client_id: client_id.clone(),
};
Json(response_data).into_response()
}
}