use crate::oauth::OAuthState;
use crate::oauth::models::{
AccessTokenClaims, OAuthError, RefreshToken, TokenRequest, TokenResponse,
};
use crate::oauth::pkce::verify_pkce;
use axum::{Form, Json, extract::State, http::StatusCode, response::IntoResponse};
use chrono::{Duration, Utc};
use jsonwebtoken::{EncodingKey, Header, encode};
use rand::Rng;
pub async fn token_endpoint(
State(state): State<OAuthState>,
Form(request): Form<TokenRequest>,
) -> Result<impl IntoResponse, (StatusCode, Json<OAuthError>)> {
let is_valid = state
.storage
.verify_client_secret(&request.client_id, &request.client_secret)
.await
.map_err(|_| {
(
StatusCode::UNAUTHORIZED,
Json(OAuthError::invalid_client("Invalid client credentials")),
)
})?;
if !is_valid {
return Err((
StatusCode::UNAUTHORIZED,
Json(OAuthError::invalid_client("Invalid client credentials")),
));
}
match request.grant_type.as_str() {
"authorization_code" => handle_authorization_code_grant(state, request).await,
"refresh_token" => handle_refresh_token_grant(state, request).await,
_ => Err((
StatusCode::BAD_REQUEST,
Json(OAuthError::unsupported_grant_type(format!(
"grant_type '{}' not supported",
request.grant_type
))),
)),
}
}
async fn handle_authorization_code_grant(
state: OAuthState,
request: TokenRequest,
) -> Result<(StatusCode, Json<TokenResponse>), (StatusCode, Json<OAuthError>)> {
let code = request.code.ok_or_else(|| {
(
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_request("code is required")),
)
})?;
let redirect_uri = request.redirect_uri.ok_or_else(|| {
(
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_request("redirect_uri is required")),
)
})?;
let code_verifier = request.code_verifier.ok_or_else(|| {
(
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_request("code_verifier is required")),
)
})?;
let auth_code = state
.storage
.get_authorization_code(&code)
.await
.map_err(|_| {
(
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_grant(
"Invalid or expired authorization code",
)),
)
})?;
if auth_code.client_id != request.client_id {
return Err((
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_grant("client_id mismatch")),
));
}
if auth_code.redirect_uri != redirect_uri {
return Err((
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_grant("redirect_uri mismatch")),
));
}
if !verify_pkce(&code_verifier, &auth_code.code_challenge) {
return Err((
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_grant("PKCE verification failed")),
));
}
state
.storage
.delete_authorization_code(&code)
.await
.map_err(|e| {
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(OAuthError::invalid_request(format!(
"Failed to delete authorization code: {}",
e
))),
)
})?;
let access_token = generate_access_token(
&request.client_id,
&auth_code.scopes,
auth_code.resource.as_deref(),
)?;
let refresh_token_value = generate_refresh_token();
let refresh_token = RefreshToken {
token: refresh_token_value.clone(),
client_id: request.client_id.clone(),
resource: auth_code.resource.clone(),
scopes: auth_code.scopes.clone(),
expires_at: Utc::now() + Duration::days(30), created_at: Utc::now(),
};
state
.storage
.save_refresh_token(&refresh_token)
.await
.map_err(|e| {
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(OAuthError::invalid_request(format!(
"Failed to save refresh token: {}",
e
))),
)
})?;
let response = TokenResponse {
access_token,
token_type: "Bearer".to_string(),
expires_in: 3600, refresh_token: Some(refresh_token_value),
scope: Some(auth_code.scopes.join(" ")),
};
Ok((StatusCode::OK, Json(response)))
}
async fn handle_refresh_token_grant(
state: OAuthState,
request: TokenRequest,
) -> Result<(StatusCode, Json<TokenResponse>), (StatusCode, Json<OAuthError>)> {
let refresh_token_value = request.refresh_token.ok_or_else(|| {
(
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_request("refresh_token is required")),
)
})?;
let old_refresh_token = state
.storage
.get_refresh_token(&refresh_token_value)
.await
.map_err(|_| {
(
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_grant(
"Invalid or expired refresh token",
)),
)
})?;
if old_refresh_token.client_id != request.client_id {
return Err((
StatusCode::BAD_REQUEST,
Json(OAuthError::invalid_grant("client_id mismatch")),
));
}
let access_token = generate_access_token(
&request.client_id,
&old_refresh_token.scopes,
old_refresh_token.resource.as_deref(),
)?;
let new_refresh_token_value = generate_refresh_token();
let new_refresh_token = RefreshToken {
token: new_refresh_token_value.clone(),
client_id: request.client_id.clone(),
resource: old_refresh_token.resource.clone(),
scopes: old_refresh_token.scopes.clone(),
expires_at: Utc::now() + Duration::days(30), created_at: Utc::now(),
};
state
.storage
.save_refresh_token(&new_refresh_token)
.await
.map_err(|e| {
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(OAuthError::invalid_request(format!(
"Failed to save refresh token: {}",
e
))),
)
})?;
state
.storage
.delete_refresh_token(&refresh_token_value)
.await
.map_err(|e| {
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(OAuthError::invalid_request(format!(
"Failed to delete old refresh token: {}",
e
))),
)
})?;
let response = TokenResponse {
access_token,
token_type: "Bearer".to_string(),
expires_in: 3600, refresh_token: Some(new_refresh_token_value),
scope: Some(old_refresh_token.scopes.join(" ")),
};
Ok((StatusCode::OK, Json(response)))
}
fn generate_access_token(
client_id: &str,
scopes: &[String],
resource: Option<&str>,
) -> Result<String, (StatusCode, Json<OAuthError>)> {
let secret = std::env::var("JWT_SECRET")
.unwrap_or_else(|_| "REPLACE_THIS_WITH_SECURE_SECRET_FROM_ENV_OR_VAULT".to_string());
let base_url =
std::env::var("BASE_URL").unwrap_or_else(|_| "http://localhost:3000".to_string());
let now = Utc::now().timestamp();
let claims = AccessTokenClaims {
sub: client_id.to_string(),
aud: resource.map(|r| r.to_string()),
exp: now + 3600, iat: now,
iss: base_url,
scope: scopes.join(" "),
client_id: client_id.to_string(),
};
encode(
&Header::default(),
&claims,
&EncodingKey::from_secret(secret.as_bytes()),
)
.map_err(|e| {
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(OAuthError::invalid_request(format!(
"Failed to generate token: {}",
e
))),
)
})
}
fn generate_refresh_token() -> String {
const CHARSET: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_";
let mut rng = rand::thread_rng();
(0..64)
.map(|_| {
let idx = rng.gen_range(0..CHARSET.len());
CHARSET[idx] as char
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_generate_refresh_token_length() {
let token = generate_refresh_token();
assert_eq!(token.len(), 64);
}
#[test]
fn test_generate_refresh_token_charset() {
let token = generate_refresh_token();
for c in token.chars() {
assert!(c.is_ascii_alphanumeric() || c == '-' || c == '_');
}
}
}