use base64::Engine;
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use rand::seq::SliceRandom;
use sha2::{Digest, Sha256};
use std::sync::Mutex;
use better_auth_core::entity::{AuthApiKey as _, AuthUser};
use better_auth_core::store::ConsumeApiKeyResult;
use better_auth_core::{AuthContext, AuthError, AuthResult, BeforeRequestAction};
use better_auth_core::{AuthRequest, AuthResponse};
pub(super) mod handlers;
pub(super) mod types;
#[cfg(test)]
mod tests;
use handlers::*;
use types::*;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ApiKeyErrorCode {
InvalidApiKey,
KeyDisabled,
KeyExpired,
UsageExceeded,
KeyNotFound,
RateLimited,
UnauthorizedSession,
InvalidPrefixLength,
InvalidNameLength,
MetadataDisabled,
NoValuesToUpdate,
KeyDisabledExpiration,
ExpiresInTooSmall,
ExpiresInTooLarge,
InvalidRemaining,
RefillAmountAndIntervalRequired,
NameRequired,
InvalidUserIdFromApiKey,
ServerOnlyProperty,
FailedToUpdateApiKey,
InvalidMetadataType,
}
impl ApiKeyErrorCode {
pub fn as_str(self) -> &'static str {
match self {
Self::InvalidApiKey => "INVALID_API_KEY",
Self::KeyDisabled => "KEY_DISABLED",
Self::KeyExpired => "KEY_EXPIRED",
Self::UsageExceeded => "USAGE_EXCEEDED",
Self::KeyNotFound => "KEY_NOT_FOUND",
Self::RateLimited => "RATE_LIMITED",
Self::UnauthorizedSession => "UNAUTHORIZED_SESSION",
Self::InvalidPrefixLength => "INVALID_PREFIX_LENGTH",
Self::InvalidNameLength => "INVALID_NAME_LENGTH",
Self::MetadataDisabled => "METADATA_DISABLED",
Self::NoValuesToUpdate => "NO_VALUES_TO_UPDATE",
Self::KeyDisabledExpiration => "KEY_DISABLED_EXPIRATION",
Self::ExpiresInTooSmall => "EXPIRES_IN_IS_TOO_SMALL",
Self::ExpiresInTooLarge => "EXPIRES_IN_IS_TOO_LARGE",
Self::InvalidRemaining => "INVALID_REMAINING",
Self::RefillAmountAndIntervalRequired => "REFILL_AMOUNT_AND_INTERVAL_REQUIRED",
Self::NameRequired => "NAME_REQUIRED",
Self::InvalidUserIdFromApiKey => "INVALID_USER_ID_FROM_API_KEY",
Self::ServerOnlyProperty => "SERVER_ONLY_PROPERTY",
Self::FailedToUpdateApiKey => "FAILED_TO_UPDATE_API_KEY",
Self::InvalidMetadataType => "INVALID_METADATA_TYPE",
}
}
pub fn message(self) -> &'static str {
match self {
Self::InvalidApiKey => "Invalid API key.",
Self::KeyDisabled => "API Key is disabled",
Self::KeyExpired => "API Key has expired",
Self::UsageExceeded => "API Key has reached its usage limit",
Self::KeyNotFound => "API Key not found",
Self::RateLimited => "Rate limit exceeded.",
Self::UnauthorizedSession => "Unauthorized or invalid session",
Self::InvalidPrefixLength => "The prefix length is either too large or too small.",
Self::InvalidNameLength => "The name length is either too large or too small.",
Self::MetadataDisabled => "Metadata is disabled.",
Self::NoValuesToUpdate => "No values to update.",
Self::KeyDisabledExpiration => "Custom key expiration values are disabled.",
Self::ExpiresInTooSmall => {
"The expiresIn is smaller than the predefined minimum value."
}
Self::ExpiresInTooLarge => "The expiresIn is larger than the predefined maximum value.",
Self::InvalidRemaining => "The remaining count is either too large or too small.",
Self::RefillAmountAndIntervalRequired => {
"refillAmount and refillInterval must both be provided together"
}
Self::NameRequired => "API Key name is required.",
Self::InvalidUserIdFromApiKey => "The user id from the API key is invalid.",
Self::ServerOnlyProperty => {
"The property you're trying to set can only be set from the server auth instance only."
}
Self::FailedToUpdateApiKey => "Failed to update API key",
Self::InvalidMetadataType => "metadata must be an object or undefined",
}
}
}
fn api_key_error(code: ApiKeyErrorCode) -> AuthError {
AuthError::bad_request(code.message())
}
pub(super) struct ApiKeyValidationError {
#[cfg_attr(
not(test),
expect(dead_code, reason = "read by the test module's verify_key helper")
)]
pub(super) code: ApiKeyErrorCode,
pub(super) message: String,
}
impl ApiKeyValidationError {
fn new(code: ApiKeyErrorCode) -> Self {
Self {
message: code.message().to_string(),
code,
}
}
}
pub struct ApiKeyPlugin {
pub(super) config: ApiKeyConfig,
last_expired_check: Mutex<Option<std::time::Instant>>,
}
#[derive(Debug, Clone)]
pub struct ApiKeyConfig {
pub key_length: usize,
pub prefix: Option<String>,
pub default_remaining: Option<i64>,
pub api_key_header: String,
pub disable_key_hashing: bool,
pub starting_characters_length: usize,
pub store_starting_characters: bool,
pub max_prefix_length: usize,
pub min_prefix_length: usize,
pub max_name_length: usize,
pub min_name_length: usize,
pub require_name: bool,
pub enable_metadata: bool,
pub key_expiration: KeyExpirationConfig,
pub rate_limit: RateLimitDefaults,
pub enable_session_for_api_keys: bool,
}
#[derive(Debug, Clone)]
pub struct KeyExpirationConfig {
pub default_expires_in: Option<i64>,
pub disable_custom_expires_time: bool,
pub max_expires_in: i64,
pub min_expires_in: i64,
}
impl Default for KeyExpirationConfig {
fn default() -> Self {
Self {
default_expires_in: None,
disable_custom_expires_time: false,
max_expires_in: 365,
min_expires_in: 1,
}
}
}
#[derive(Debug, Clone)]
pub struct RateLimitDefaults {
pub enabled: bool,
pub time_window: i64,
pub max_requests: i64,
}
impl Default for RateLimitDefaults {
fn default() -> Self {
Self {
enabled: true,
time_window: 86_400_000, max_requests: 10,
}
}
}
impl Default for ApiKeyConfig {
fn default() -> Self {
Self {
key_length: 64,
prefix: None,
default_remaining: None,
api_key_header: "x-api-key".to_string(),
disable_key_hashing: false,
starting_characters_length: 6,
store_starting_characters: true,
max_prefix_length: 32,
min_prefix_length: 1,
max_name_length: 32,
min_name_length: 1,
require_name: false,
enable_metadata: false,
key_expiration: KeyExpirationConfig::default(),
rate_limit: RateLimitDefaults::default(),
enable_session_for_api_keys: false,
}
}
}
#[bon::bon]
impl ApiKeyPlugin {
#[builder]
pub fn new(
#[builder(default = 64)] key_length: usize,
prefix: Option<String>,
default_remaining: Option<i64>,
#[builder(default = "x-api-key".to_string())] api_key_header: String,
#[builder(default = false)] disable_key_hashing: bool,
#[builder(default = 6)] starting_characters_length: usize,
#[builder(default = true)] store_starting_characters: bool,
#[builder(default = 32)] max_prefix_length: usize,
#[builder(default = 1)] min_prefix_length: usize,
#[builder(default = 32)] max_name_length: usize,
#[builder(default = 1)] min_name_length: usize,
#[builder(default = false)] require_name: bool,
#[builder(default = false)] enable_metadata: bool,
#[builder(default)] key_expiration: KeyExpirationConfig,
#[builder(default)] rate_limit: RateLimitDefaults,
#[builder(default = false)] enable_session_for_api_keys: bool,
) -> Self {
Self {
config: ApiKeyConfig {
key_length,
prefix,
default_remaining,
api_key_header,
disable_key_hashing,
starting_characters_length,
store_starting_characters,
max_prefix_length,
min_prefix_length,
max_name_length,
min_name_length,
require_name,
enable_metadata,
key_expiration,
rate_limit,
enable_session_for_api_keys,
},
last_expired_check: Mutex::new(None),
}
}
pub fn with_config(config: ApiKeyConfig) -> Self {
Self {
config,
last_expired_check: Mutex::new(None),
}
}
pub(super) fn generate_key(&self, custom_prefix: Option<&str>) -> (String, String, String) {
const ALPHABET: &[u8] = b"abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ";
let mut rng = rand::thread_rng();
let raw: String = (0..self.config.key_length)
.map(|_| {
ALPHABET
.choose(&mut rng)
.copied()
.map(char::from)
.unwrap_or('a')
})
.collect();
let prefix = custom_prefix
.or(self.config.prefix.as_deref())
.unwrap_or("");
let full_key = format!("{}{}", prefix, raw);
let start_len = self.config.starting_characters_length;
let start: String = full_key.chars().take(start_len).collect();
let hash = if self.config.disable_key_hashing {
full_key.clone()
} else {
Self::hash_key(&full_key)
};
(full_key, hash, start)
}
fn hash_key(key: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(key.as_bytes());
let digest = hasher.finalize();
URL_SAFE_NO_PAD.encode(digest)
}
pub(super) async fn maybe_delete_expired(
&self,
ctx: &AuthContext<impl better_auth_core::AuthSchema>,
) {
let should_run = {
let mut last = self
.last_expired_check
.lock()
.unwrap_or_else(|e| e.into_inner());
let now = std::time::Instant::now();
match *last {
Some(prev) if now.duration_since(prev).as_secs() < 10 => false,
_ => {
*last = Some(now);
true
}
}
};
if should_run {
let _ = ctx.database.delete_expired_api_keys().await;
}
}
pub(super) fn validate_prefix(&self, prefix: Option<&str>) -> AuthResult<()> {
if let Some(p) = prefix {
let len = p.len();
if len < self.config.min_prefix_length || len > self.config.max_prefix_length {
return Err(api_key_error(ApiKeyErrorCode::InvalidPrefixLength));
}
}
Ok(())
}
pub(super) fn validate_name(&self, name: Option<&str>, is_create: bool) -> AuthResult<()> {
if is_create && self.config.require_name && name.is_none() {
return Err(api_key_error(ApiKeyErrorCode::NameRequired));
}
if let Some(n) = name {
let len = n.len();
if len < self.config.min_name_length || len > self.config.max_name_length {
return Err(api_key_error(ApiKeyErrorCode::InvalidNameLength));
}
}
Ok(())
}
pub(super) fn validate_expires_in(&self, expires_in: Option<i64>) -> AuthResult<Option<i64>> {
let cfg = &self.config.key_expiration;
if let Some(secs) = expires_in {
if cfg.disable_custom_expires_time {
return Err(api_key_error(ApiKeyErrorCode::KeyDisabledExpiration));
}
let days = secs as f64 / 86_400.0;
if days < cfg.min_expires_in as f64 {
return Err(api_key_error(ApiKeyErrorCode::ExpiresInTooSmall));
}
if days > cfg.max_expires_in as f64 {
return Err(api_key_error(ApiKeyErrorCode::ExpiresInTooLarge));
}
Ok(Some(secs))
} else {
Ok(cfg.default_expires_in)
}
}
pub(super) fn validate_metadata(&self, metadata: &Option<serde_json::Value>) -> AuthResult<()> {
if metadata.is_some() && !self.config.enable_metadata {
return Err(api_key_error(ApiKeyErrorCode::MetadataDisabled));
}
if let Some(v) = metadata
&& !v.is_object()
&& !v.is_null()
{
return Err(api_key_error(ApiKeyErrorCode::InvalidMetadataType));
}
Ok(())
}
pub(super) fn validate_refill(
refill_interval: Option<i64>,
refill_amount: Option<i64>,
) -> AuthResult<()> {
match (refill_interval, refill_amount) {
(Some(_), None) | (None, Some(_)) => Err(api_key_error(
ApiKeyErrorCode::RefillAmountAndIntervalRequired,
)),
_ => Ok(()),
}
}
async fn handle_create(
&self,
req: &AuthRequest,
ctx: &AuthContext<impl better_auth_core::AuthSchema>,
) -> AuthResult<AuthResponse> {
let (user, _session) = ctx.require_session(req).await?;
let body: CreateKeyRequest = match better_auth_core::validate_request_body(req) {
Ok(v) => v,
Err(resp) => return Ok(resp),
};
let response = create_key_core(&body, user.id(), self, ctx).await?;
Ok(AuthResponse::json(200, &response)?)
}
async fn handle_get(
&self,
req: &AuthRequest,
ctx: &AuthContext<impl better_auth_core::AuthSchema>,
) -> AuthResult<AuthResponse> {
let (user, _session) = ctx.require_session(req).await?;
let id = req
.query
.get("id")
.ok_or_else(|| AuthError::bad_request("Query parameter 'id' is required"))?;
let response = get_key_core(id, user.id(), self, ctx).await?;
Ok(AuthResponse::json(200, &response)?)
}
async fn handle_list(
&self,
req: &AuthRequest,
ctx: &AuthContext<impl better_auth_core::AuthSchema>,
) -> AuthResult<AuthResponse> {
let (user, _session) = ctx.require_session(req).await?;
let response = list_keys_core(user.id(), self, ctx).await?;
Ok(AuthResponse::json(200, &response)?)
}
async fn handle_update(
&self,
req: &AuthRequest,
ctx: &AuthContext<impl better_auth_core::AuthSchema>,
) -> AuthResult<AuthResponse> {
let (user, _session) = ctx.require_session(req).await?;
let body: UpdateKeyRequest = match better_auth_core::validate_request_body(req) {
Ok(v) => v,
Err(resp) => return Ok(resp),
};
let response = update_key_core(&body, user.id(), self, ctx).await?;
Ok(AuthResponse::json(200, &response)?)
}
async fn handle_delete(
&self,
req: &AuthRequest,
ctx: &AuthContext<impl better_auth_core::AuthSchema>,
) -> AuthResult<AuthResponse> {
let (user, _session) = ctx.require_session(req).await?;
let body: DeleteKeyRequest = match better_auth_core::validate_request_body(req) {
Ok(v) => v,
Err(resp) => return Ok(resp),
};
let response = delete_key_core(&body, user.id(), self, ctx).await?;
Ok(AuthResponse::json(200, &response)?)
}
pub(super) async fn validate_api_key(
&self,
ctx: &AuthContext<impl better_auth_core::AuthSchema>,
raw_key: &str,
required_permissions: Option<&serde_json::Value>,
) -> Result<ApiKeyView, ApiKeyValidationError> {
let hashed = if self.config.disable_key_hashing {
raw_key.to_string()
} else {
Self::hash_key(raw_key)
};
let api_key = ctx
.database
.get_api_key_by_hash(&hashed)
.await
.map_err(|_| ApiKeyValidationError::new(ApiKeyErrorCode::InvalidApiKey))?
.ok_or_else(|| ApiKeyValidationError::new(ApiKeyErrorCode::InvalidApiKey))?;
if !api_key.enabled() {
return Err(ApiKeyValidationError::new(ApiKeyErrorCode::KeyDisabled));
}
if let Some(expires_at_str) = api_key.expires_at()
&& let Ok(expires_at) = chrono::DateTime::parse_from_rfc3339(expires_at_str)
&& chrono::Utc::now() > expires_at
{
let _ = ctx.database.delete_api_key(&api_key.id()).await;
return Err(ApiKeyValidationError::new(ApiKeyErrorCode::KeyExpired));
}
if let Some(required) = required_permissions {
let key_perms_str = api_key.permissions().unwrap_or("");
if key_perms_str.is_empty() {
return Err(ApiKeyValidationError::new(ApiKeyErrorCode::KeyNotFound));
}
if !check_permissions(key_perms_str, required) {
return Err(ApiKeyValidationError::new(ApiKeyErrorCode::KeyNotFound));
}
}
let updated = match ctx
.database
.consume_api_key_usage(&api_key.id(), self.config.rate_limit.enabled)
.await
.map_err(|_| ApiKeyValidationError::new(ApiKeyErrorCode::FailedToUpdateApiKey))?
{
ConsumeApiKeyResult::Allowed(key) => *key,
ConsumeApiKeyResult::RateLimited => {
return Err(ApiKeyValidationError::new(ApiKeyErrorCode::RateLimited));
}
ConsumeApiKeyResult::UsageExhausted => {
return Err(ApiKeyValidationError::new(ApiKeyErrorCode::UsageExceeded));
}
};
self.maybe_delete_expired(ctx).await;
Ok(ApiKeyView::from(&updated))
}
}
better_auth_core::impl_auth_plugin! {
ApiKeyPlugin, "api-key";
routes {
post "/api-key/create" => handle_create, "api_key_create";
get "/api-key/get" => handle_get, "api_key_get";
post "/api-key/update" => handle_update, "api_key_update";
post "/api-key/delete" => handle_delete, "api_key_delete";
get "/api-key/list" => handle_list, "api_key_list";
}
extra {
async fn before_request(
&self,
req: &AuthRequest,
ctx: &AuthContext<S>,
) -> AuthResult<Option<BeforeRequestAction>> {
if !self.config.enable_session_for_api_keys {
return Ok(None);
}
let raw_key = match req.headers.get(&self.config.api_key_header) {
Some(k) if !k.is_empty() => k.clone(),
_ => return Ok(None),
};
let view = self
.validate_api_key(ctx, &raw_key, None)
.await
.map_err(|e| AuthError::bad_request(e.message))?;
let user = ctx
.database
.get_user_by_id(&view.user_id)
.await?
.ok_or_else(|| api_key_error(ApiKeyErrorCode::InvalidUserIdFromApiKey))?;
if req.path() == "/get-session" {
let session_json = serde_json::json!({
"user": {
"id": user.id(),
"email": user.email(),
"name": user.name(),
},
"session": {
"id": view.id,
"token": raw_key,
"userId": view.user_id,
}
});
return Ok(Some(BeforeRequestAction::Respond(AuthResponse::json(
200,
&session_json,
)?)));
}
Ok(Some(BeforeRequestAction::InjectSession {
user_id: view.user_id,
session_token: raw_key,
}))
}
}
}