use crate::oauth::OAuthState;
use crate::oauth::models::{AuthorizationCode, AuthorizeRequest, OAuthError};
use crate::oauth::pkce::validate_code_challenge;
use axum::{
Form,
extract::{Query, State},
http::StatusCode,
response::{Html, IntoResponse, Redirect},
};
use chrono::{Duration, Utc};
use serde::Deserialize;
pub async fn authorize_get(
State(state): State<OAuthState>,
Query(params): Query<AuthorizeRequest>,
) -> Result<impl IntoResponse, (StatusCode, Json<OAuthError>)> {
validate_authorize_request(¶ms)?;
let client = state
.storage
.get_client(¶ms.client_id)
.await
.map_err(|_| {
(
StatusCode::UNAUTHORIZED,
Json(OAuthError::invalid_client("Client not found")),
)
})?;
if !client.redirect_uris.contains(¶ms.redirect_uri) {
return Err((
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_request(
"redirect_uri does not match registered URIs",
)),
));
}
let html = render_consent_form(¶ms);
Ok(Html(html))
}
#[derive(Debug, Deserialize)]
pub struct AuthorizeForm {
pub client_id: String,
pub redirect_uri: String,
pub state: Option<String>,
pub code_challenge: String,
pub resource: Option<String>,
pub scope: Option<String>,
pub approved: String,
}
pub async fn authorize_post(
State(state): State<OAuthState>,
Form(form): Form<AuthorizeForm>,
) -> Result<impl IntoResponse, (StatusCode, Json<OAuthError>)> {
if form.approved != "true" {
let error_redirect = format!(
"{}?error=access_denied&error_description=User denied authorization{}",
form.redirect_uri,
form.state
.as_ref()
.map(|s| format!("&state={}", s))
.unwrap_or_default()
);
return Ok(Redirect::to(&error_redirect).into_response());
}
let client = state
.storage
.get_client(&form.client_id)
.await
.map_err(|_| {
(
StatusCode::UNAUTHORIZED,
Json(OAuthError::invalid_client("Client not found")),
)
})?;
if !client.redirect_uris.contains(&form.redirect_uri) {
return Err((
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_request(
"redirect_uri does not match registered URIs",
)),
));
}
let code = generate_authorization_code();
let scopes = form
.scope
.as_ref()
.map(|s| {
s.split_whitespace()
.map(|scope| scope.to_string())
.collect()
})
.unwrap_or_else(Vec::new);
let authorization_code = AuthorizationCode {
code: code.clone(),
client_id: form.client_id.clone(),
redirect_uri: form.redirect_uri.clone(),
code_challenge: form.code_challenge,
resource: form.resource,
scopes,
expires_at: Utc::now() + Duration::minutes(10), created_at: Utc::now(),
};
state
.storage
.save_authorization_code(&authorization_code)
.await
.map_err(|e| {
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(OAuthError::invalid_request(format!(
"Failed to save authorization code: {}",
e
))),
)
})?;
let success_redirect = format!(
"{}?code={}{}",
form.redirect_uri,
code,
form.state
.as_ref()
.map(|s| format!("&state={}", s))
.unwrap_or_default()
);
Ok(Redirect::to(&success_redirect).into_response())
}
fn validate_authorize_request(
params: &AuthorizeRequest,
) -> Result<(), (StatusCode, Json<OAuthError>)> {
if params.response_type != "code" {
return Err((
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_request("response_type must be 'code'")),
));
}
if params.code_challenge_method != "S256" {
return Err((
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_request(
"code_challenge_method must be 'S256'",
)),
));
}
if !validate_code_challenge(¶ms.code_challenge) {
return Err((
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_request("Invalid code_challenge format")),
));
}
if !params.redirect_uri.starts_with("https://")
&& !params.redirect_uri.starts_with("http://localhost")
&& !params.redirect_uri.starts_with("http://127.0.0.1")
{
return Err((
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_request(
"redirect_uri must be HTTPS or http://localhost",
)),
));
}
Ok(())
}
fn render_consent_form(params: &AuthorizeRequest) -> String {
let scopes = params
.scope
.as_ref()
.map(|s| s.split_whitespace().collect::<Vec<_>>())
.unwrap_or_default();
format!(
r#"<!DOCTYPE html>
<html>
<head>
<title>Authorization Request</title>
<style>
body {{ font-family: Arial, sans-serif; max-width: 500px; margin: 50px auto; padding: 20px; }}
.consent-box {{ border: 1px solid #ccc; padding: 20px; border-radius: 5px; }}
.scopes {{ margin: 20px 0; }}
.scope-item {{ padding: 5px 0; }}
.buttons {{ margin-top: 20px; }}
button {{ padding: 10px 20px; margin-right: 10px; cursor: pointer; }}
.approve {{ background-color: #4CAF50; color: white; border: none; }}
.deny {{ background-color: #f44336; color: white; border: none; }}
</style>
</head>
<body>
<div class="consent-box">
<h2>Authorization Request</h2>
<p><strong>Client:</strong> {}</p>
<p><strong>Redirect URI:</strong> {}</p>
<div class="scopes">
<p><strong>Requested Permissions:</strong></p>
{}
</div>
<form method="POST" action="/oauth/authorize">
<input type="hidden" name="client_id" value="{}">
<input type="hidden" name="redirect_uri" value="{}">
<input type="hidden" name="code_challenge" value="{}">
{}
{}
{}
<div class="buttons">
<button type="submit" name="approved" value="true" class="approve">Approve</button>
<button type="submit" name="approved" value="false" class="deny">Deny</button>
</div>
</form>
</div>
</body>
</html>"#,
params.client_id,
params.redirect_uri,
if scopes.is_empty() {
"<p>No specific permissions requested</p>".to_string()
} else {
scopes
.iter()
.map(|s| format!("<div class='scope-item'>• {}</div>", s))
.collect::<Vec<_>>()
.join("\n")
},
params.client_id,
params.redirect_uri,
params.code_challenge,
params
.state
.as_ref()
.map(|s| format!(r#"<input type="hidden" name="state" value="{}">"#, s))
.unwrap_or_default(),
params
.resource
.as_ref()
.map(|r| format!(r#"<input type="hidden" name="resource" value="{}">"#, r))
.unwrap_or_default(),
params
.scope
.as_ref()
.map(|s| format!(r#"<input type="hidden" name="scope" value="{}">"#, s))
.unwrap_or_default(),
)
}
fn generate_authorization_code() -> String {
use rand::Rng;
const CHARSET: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789";
let mut rng = rand::thread_rng();
(0..32)
.map(|_| {
let idx = rng.gen_range(0..CHARSET.len());
CHARSET[idx] as char
})
.collect()
}
use axum::Json;