use chrono::Utc;
use sha2::{Digest, Sha256};
use sqlx::Row;
use sqlx::types::Json;
use subtle::ConstantTimeEq;
use zeroize::Zeroize;
use super::dialect::{TokenDb, TokenPool, restored_time, sql, stored_time};
use super::error::ApiTokenError;
use super::migrate;
use super::token::{
ID_BYTES, IssuedApiToken, PlaintextToken, SECRET_BYTES, format_plaintext, parse_plaintext,
};
use super::{Abilities, ApiToken, ApiTokenId, NewApiToken};
type TokenRow = <TokenDb as sqlx::Database>::Row;
const ISSUE_ATTEMPTS: u32 = 8;
#[derive(Clone, Debug)]
#[non_exhaustive]
pub struct ApiTokens {
pool: TokenPool,
}
impl ApiTokens {
#[must_use]
pub fn new(pool: TokenPool) -> Self {
Self { pool }
}
#[must_use]
pub fn pool(&self) -> &TokenPool {
&self.pool
}
pub async fn migrate(&self) -> Result<(), ApiTokenError> {
migrate::apply(&self.pool).await
}
pub async fn issue(&self, request: &NewApiToken) -> Result<IssuedApiToken, ApiTokenError> {
let mut secret = [0u8; SECRET_BYTES];
fill_random(&mut secret)?;
let digest = digest_of(&secret);
let abilities = request.abilities_ref().clone();
let created_at = Utc::now();
let expires_at = request.expires_at();
for _ in 0..ISSUE_ATTEMPTS {
let mut id_bytes = [0u8; ID_BYTES];
fill_random(&mut id_bytes)?;
let written = sqlx::query(sql::INSERT_NEW)
.bind(id_bytes.to_vec())
.bind(digest.to_vec())
.bind(request.tokenable_id())
.bind(request.name())
.bind(Json(abilities.as_slice()))
.bind(stored_time(expires_at))
.bind(stored_time(created_at))
.execute(&self.pool)
.await?
.rows_affected();
if written == 0 {
continue;
}
let id = ApiTokenId::from_bytes(id_bytes);
let plaintext = PlaintextToken::new(format_plaintext(&id_bytes, &secret));
let token = ApiToken::from_row(
id,
request.tokenable_id().to_owned(),
request.name().to_owned(),
abilities,
expires_at,
created_at,
);
return Ok(IssuedApiToken::new(token, plaintext));
}
Err(ApiTokenError::IdCollision {
attempts: ISSUE_ATTEMPTS,
})
}
pub async fn find(&self, id: ApiTokenId) -> Result<Option<ApiToken>, ApiTokenError> {
let Some(row) = sqlx::query(sql::FIND)
.bind(id.as_bytes().to_vec())
.fetch_optional(&self.pool)
.await?
else {
return Ok(None);
};
let tokenable_id: String = row.try_get(0)?;
let name: String = row.try_get(1)?;
let abilities = abilities_at(&row, 2)?;
let expires_at = restored_time(row.try_get(3)?)?;
let created_at = restored_time(row.try_get(4)?)?;
Ok(Some(ApiToken::from_row(
id,
tokenable_id,
name,
abilities,
expires_at,
created_at,
)))
}
pub async fn authenticate(&self, presented: &str) -> Result<Option<ApiToken>, ApiTokenError> {
let Some((id, mut secret)) = parse_plaintext(presented) else {
return Ok(None);
};
let mut presented_digest = digest_of(&secret);
secret.zeroize();
let found = sqlx::query(sql::AUTHENTICATE)
.bind(id.as_bytes().to_vec())
.fetch_optional(&self.pool)
.await?;
let Some(row) = found else {
presented_digest.zeroize();
return Ok(None);
};
let stored: Vec<u8> = row.try_get(0)?;
let matches: bool = presented_digest.ct_eq(stored.as_slice()).into();
presented_digest.zeroize();
if !matches {
return Ok(None);
}
let tokenable_id: String = row.try_get(1)?;
let name: String = row.try_get(2)?;
let abilities = abilities_at(&row, 3)?;
let expires_at = restored_time(row.try_get(4)?)?;
let created_at = restored_time(row.try_get(5)?)?;
Ok(Some(ApiToken::from_row(
id,
tokenable_id,
name,
abilities,
expires_at,
created_at,
)))
}
pub async fn list_for(&self, tokenable_id: &str) -> Result<Vec<ApiToken>, ApiTokenError> {
let rows = sqlx::query(sql::LIST_FOR)
.bind(tokenable_id)
.fetch_all(&self.pool)
.await?;
let mut tokens = Vec::with_capacity(rows.len());
for row in rows {
let raw_id: Vec<u8> = row.try_get(0)?;
let id_bytes: [u8; ID_BYTES] = raw_id.as_slice().try_into().map_err(|_| {
ApiTokenError::Decode(format!(
"id column holds {} bytes, expected {ID_BYTES}",
raw_id.len()
))
})?;
let tokenable_id: String = row.try_get(1)?;
let name: String = row.try_get(2)?;
let abilities = abilities_at(&row, 3)?;
let expires_at = restored_time(row.try_get(4)?)?;
let created_at = restored_time(row.try_get(5)?)?;
tokens.push(ApiToken::from_row(
ApiTokenId::from_bytes(id_bytes),
tokenable_id,
name,
abilities,
expires_at,
created_at,
));
}
Ok(tokens)
}
pub async fn revoke(&self, id: ApiTokenId) -> Result<bool, ApiTokenError> {
let result = sqlx::query(sql::DELETE)
.bind(id.as_bytes().to_vec())
.execute(&self.pool)
.await?;
Ok(result.rows_affected() > 0)
}
pub async fn revoke_all_for(&self, tokenable_id: &str) -> Result<u64, ApiTokenError> {
let result = sqlx::query(sql::DELETE_FOR)
.bind(tokenable_id)
.execute(&self.pool)
.await?;
Ok(result.rows_affected())
}
pub async fn sweep_expired(&self) -> Result<u64, ApiTokenError> {
let result = sqlx::query(sql::DELETE_EXPIRED).execute(&self.pool).await?;
Ok(result.rows_affected())
}
}
pub(crate) fn digest_of(secret: &[u8]) -> [u8; 32] {
Sha256::digest(secret).into()
}
fn fill_random(buffer: &mut [u8]) -> Result<(), ApiTokenError> {
getrandom::fill(buffer).map_err(|_| ApiTokenError::Entropy)
}
fn abilities_at(row: &TokenRow, index: usize) -> Result<Abilities, ApiTokenError> {
let stored: Json<Vec<String>> = row
.try_get(index)
.map_err(|error| ApiTokenError::Decode(error.to_string()))?;
Ok(Abilities::of(stored.0))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_digest_is_not_the_secret() {
let secret = [0xa5u8; SECRET_BYTES];
let digest = digest_of(&secret);
assert_eq!(digest.len(), 32);
assert_ne!(&digest[..], &secret[..]);
}
#[test]
fn the_digest_is_the_documented_sha256() {
let digest = super::super::token::hex_encode(&digest_of(&[0u8; 32]));
assert_eq!(
digest,
"66687aadf862bd776c8fc18b8e9f8e20089714856ee233b3902a591d0d5f2925"
);
}
#[test]
fn two_secrets_that_differ_in_one_bit_give_different_digests() {
let a = [0u8; SECRET_BYTES];
let mut b = [0u8; SECRET_BYTES];
b[SECRET_BYTES - 1] = 1;
assert_ne!(digest_of(&a), digest_of(&b));
assert_eq!(digest_of(&a), digest_of(&[0u8; SECRET_BYTES]));
}
#[test]
fn the_random_source_fills_the_whole_buffer() {
let mut first = [0u8; SECRET_BYTES];
let mut second = [0u8; SECRET_BYTES];
fill_random(&mut first).expect("the OS randomness source is available");
fill_random(&mut second).expect("the OS randomness source is available");
assert_ne!(first, [0u8; SECRET_BYTES]);
assert_ne!(first, second);
}
}