use axum::{
Router,
extract::{FromRequestParts, Path, State},
http::{HeaderMap, StatusCode, request::Parts},
routing::{delete, get, post},
};
use ocre::{
ApiError, ApiResult, Created, Ctx, Error, Json,
jwt::{self, Claims, Location},
};
use serde::{Deserialize, Serialize};
use crate::models::{
api_key::{self, ApiKey, NewApiKey},
user::{self, NewUser, User},
};
pub const TOKEN_TTL_SECONDS: i64 = 3600;
pub const TOKEN_LOCATIONS: &[Location] = &[Location::Bearer];
pub const RATE_LIMITER: &str = "AUTH_RATE_LIMITER";
pub fn routes() -> Router<Ctx> {
Router::new()
.route("/api/auth/signup", post(signup))
.route("/api/auth/token", post(token))
.route("/api/auth/me", get(me).delete(delete_me))
.route("/api/auth/keys", get(list_keys).post(create_key))
.route("/api/auth/keys/{id}", delete(revoke_key))
}
pub async fn throttle(ctx: &Ctx, headers: &HeaderMap, action: &str) -> ocre::Result<()> {
let ip = ocre::remote_ip(headers).map_or_else(|| "unknown".to_owned(), |ip| ip.to_string());
ocre::security::rate_limit(ctx, RATE_LIMITER, &format!("{action}:{ip}")).await
}
pub struct BearerUser(pub User);
impl FromRequestParts<Ctx> for BearerUser {
type Rejection = ApiError;
async fn from_request_parts(parts: &mut Parts, ctx: &Ctx) -> Result<Self, ApiError> {
let token = jwt::token_from(&parts.headers, &parts.uri, TOKEN_LOCATIONS).ok_or(Error::Unauthorized)?;
let user_id = if token.contains('.') {
jwt::decode(ctx, &token)?.sub.parse::<i64>().map_err(|_| Error::Unauthorized)?
} else {
api_key::authenticate(ctx, &token).await?.ok_or(Error::Unauthorized)?
};
let user = user::find(ctx, user_id).await?.ok_or(Error::Unauthorized)?;
Ok(Self(user))
}
}
#[derive(Clone, Default, Deserialize)]
#[serde(default)]
pub struct Credentials {
pub email: String,
pub password: String,
}
#[derive(Clone, Default, Deserialize)]
#[serde(default)]
pub struct PasswordConfirmation {
pub password: String,
}
#[derive(Serialize)]
struct TokenResponse {
token: String,
token_type: &'static str,
expires_in: i64,
}
#[derive(Serialize)]
struct CreatedKey {
key: String,
api_key: ApiKey,
}
async fn signup(State(ctx): State<Ctx>, headers: HeaderMap, Json(new): Json<NewUser>) -> ApiResult<Created<User>> {
throttle(&ctx, &headers, "signup").await?;
Ok(Created(user::create(&ctx, new).await?))
}
async fn token(
State(ctx): State<Ctx>,
headers: HeaderMap,
Json(credentials): Json<Credentials>,
) -> ApiResult<Json<TokenResponse>> {
throttle(&ctx, &headers, "token").await?;
let user =
user::authenticate(&ctx, &credentials.email, &credentials.password).await?.ok_or(Error::Unauthorized)?;
let token = jwt::encode(&ctx, &Claims::new(user.id.to_string(), TOKEN_TTL_SECONDS))?;
Ok(Json(TokenResponse { token, token_type: "Bearer", expires_in: TOKEN_TTL_SECONDS }))
}
async fn me(BearerUser(user): BearerUser) -> Json<User> {
Json(user)
}
async fn delete_me(
State(ctx): State<Ctx>,
BearerUser(user): BearerUser,
headers: HeaderMap,
Json(confirmation): Json<PasswordConfirmation>,
) -> ApiResult<StatusCode> {
throttle(&ctx, &headers, "account_delete").await?;
if !user::deletion_confirmed(&user, &confirmation.password).await? {
return Err(Error::Forbidden.into());
}
user::delete(&ctx, user.id).await?;
Ok(StatusCode::NO_CONTENT)
}
async fn list_keys(State(ctx): State<Ctx>, BearerUser(user): BearerUser) -> ApiResult<Json<Vec<ApiKey>>> {
Ok(Json(api_key::for_user(&ctx, user.id).await?))
}
async fn create_key(
State(ctx): State<Ctx>,
BearerUser(user): BearerUser,
Json(new): Json<NewApiKey>,
) -> ApiResult<Created<CreatedKey>> {
let (api_key, key) = api_key::create(&ctx, user.id, new).await?;
Ok(Created(CreatedKey { key, api_key }))
}
async fn revoke_key(State(ctx): State<Ctx>, BearerUser(user): BearerUser, Path(id): Path<i64>) -> ApiResult<StatusCode> {
if api_key::revoke(&ctx, user.id, id).await? { Ok(StatusCode::NO_CONTENT) } else { Err(Error::NotFound.into()) }
}