use std::fmt;
use serde::{Deserialize, Deserializer, Serialize};
use crate::error::{ApiError, ValidationDetails};
use crate::ids::UserId;
use crate::text;
use crate::time::UnixMillis;
pub const REFRESH_REUSE_GRACE_SECS: u64 = 30;
pub const ACCESS_TOKEN_REFRESH_MARGIN_SECS: u64 = 60;
pub const DEFAULT_ACCESS_TOKEN_TTL_SECS: u64 = 60 * 60;
pub const DEFAULT_REFRESH_TOKEN_TTL_SECS: u64 = 30 * 24 * 60 * 60;
pub const PASSWORD_MIN_CHARS: usize = 10;
pub const PASSWORD_MAX_BYTES: usize = 128;
pub const EMAIL_MAX_BYTES: usize = 254;
pub const DISPLAY_NAME_MAX_CHARS: usize = 32;
pub const STEAM_TICKET_MAX_HEX: usize = 8192;
const _: () = assert!(STEAM_TICKET_MAX_HEX >= 2 * 2560);
pub mod provider {
pub const STEAM: &str = "steam";
}
macro_rules! secret_type {
($(#[$meta:meta])* $name:ident) => {
$(#[$meta])*
#[derive(Clone, Serialize)]
#[serde(transparent)]
pub struct $name(String);
impl<'de> Deserialize<'de> for $name {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
deserialize_secret(deserializer).map(Self)
}
}
impl $name {
pub fn new(secret: impl Into<String>) -> Self {
Self(secret.into())
}
pub fn expose(&self) -> &str {
&self.0
}
pub fn into_inner(self) -> String {
self.0
}
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
}
impl fmt::Debug for $name {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(concat!(stringify!($name), "(<redacted>)"))
}
}
impl From<String> for $name {
fn from(secret: String) -> Self {
Self(secret)
}
}
impl From<&str> for $name {
fn from(secret: &str) -> Self {
Self(secret.to_string())
}
}
};
}
secret_type!(
Password
);
secret_type!(
AccessToken
);
secret_type!(
RefreshToken
);
secret_type!(
Secret
);
fn deserialize_secret<'de, D: Deserializer<'de>>(deserializer: D) -> Result<String, D::Error> {
struct SecretVisitor;
impl<'de> serde::de::Visitor<'de> for SecretVisitor {
type Value = String;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("a string")
}
fn visit_str<E: serde::de::Error>(self, value: &str) -> Result<String, E> {
Ok(value.to_string())
}
fn visit_string<E: serde::de::Error>(self, value: String) -> Result<String, E> {
Ok(value)
}
fn visit_bool<E: serde::de::Error>(self, _: bool) -> Result<String, E> {
Err(E::custom(NOT_A_STRING))
}
fn visit_i64<E: serde::de::Error>(self, _: i64) -> Result<String, E> {
Err(E::custom(NOT_A_STRING))
}
fn visit_u64<E: serde::de::Error>(self, _: u64) -> Result<String, E> {
Err(E::custom(NOT_A_STRING))
}
fn visit_i128<E: serde::de::Error>(self, _: i128) -> Result<String, E> {
Err(E::custom(NOT_A_STRING))
}
fn visit_u128<E: serde::de::Error>(self, _: u128) -> Result<String, E> {
Err(E::custom(NOT_A_STRING))
}
fn visit_f64<E: serde::de::Error>(self, _: f64) -> Result<String, E> {
Err(E::custom(NOT_A_STRING))
}
fn visit_char<E: serde::de::Error>(self, value: char) -> Result<String, E> {
Ok(value.to_string())
}
fn visit_bytes<E: serde::de::Error>(self, _: &[u8]) -> Result<String, E> {
Err(E::custom(NOT_A_STRING))
}
fn visit_none<E: serde::de::Error>(self) -> Result<String, E> {
Err(E::custom(NOT_A_STRING))
}
fn visit_unit<E: serde::de::Error>(self) -> Result<String, E> {
Err(E::custom(NOT_A_STRING))
}
fn visit_seq<A: serde::de::SeqAccess<'de>>(self, _: A) -> Result<String, A::Error> {
Err(serde::de::Error::custom(NOT_A_STRING))
}
fn visit_map<A: serde::de::MapAccess<'de>>(self, _: A) -> Result<String, A::Error> {
Err(serde::de::Error::custom(NOT_A_STRING))
}
}
deserializer.deserialize_any(SecretVisitor)
}
const NOT_A_STRING: &str = "a secret must be a JSON string";
pub const EMAIL_LOCAL_MAX_BYTES: usize = 64;
pub fn is_valid_email(email: &str) -> bool {
let email = email.trim();
let Some((local, domain)) = email.split_once('@') else { return false };
if local.is_empty() || local.len() > EMAIL_LOCAL_MAX_BYTES || domain.is_empty() || email.len() > EMAIL_MAX_BYTES {
return false;
}
let atext = |c: char| c.is_alphanumeric() || (c.is_ascii() && "!#$%&'*+/=?^_`{|}~-".contains(c));
let local_ok = local.split('.').all(|piece| !piece.is_empty() && piece.chars().all(atext));
let label_ok = |label: &str| {
!label.is_empty() && label.len() <= 63 && !label.starts_with('-') && !label.ends_with('-') && label.chars().all(|c| c.is_alphanumeric() || c == '-')
};
local_ok && domain.split('.').all(label_ok)
}
fn check_email(email: &str, details: &mut ValidationDetails) {
if !is_valid_email(email) {
details.add("email", "is not an email address");
}
if let Some(problem) = text::name_problem(email) {
details.add("email", problem);
}
if email.len() > EMAIL_MAX_BYTES {
details.add("email", format!("is longer than {EMAIL_MAX_BYTES} bytes"));
}
}
fn check_password(field: &str, password: &Password, details: &mut ValidationDetails) {
let password = password.expose();
if password.chars().count() < PASSWORD_MIN_CHARS {
details.add(field, format!("is shorter than {PASSWORD_MIN_CHARS} characters"));
}
if password.len() > PASSWORD_MAX_BYTES {
details.add(field, format!("is longer than {PASSWORD_MAX_BYTES} bytes"));
}
if password.chars().any(char::is_control) {
details.add(field, "contains control characters");
}
if !password.is_empty() && password.chars().all(char::is_whitespace) {
details.add(field, "is only whitespace");
}
}
fn check_display_name(name: &str, details: &mut ValidationDetails) {
if name.trim().is_empty() {
details.add("display_name", "is empty");
} else if name.trim() != name {
details.add("display_name", "starts or ends with whitespace");
}
if name.chars().count() > DISPLAY_NAME_MAX_CHARS {
details.add("display_name", format!("is longer than {DISPLAY_NAME_MAX_CHARS} characters"));
}
if let Some(problem) = text::name_problem(name) {
details.add("display_name", problem);
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[non_exhaustive]
pub struct RegisterRequest {
pub email: String,
pub password: Password,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub display_name: Option<String>,
}
impl RegisterRequest {
pub fn new(email: impl Into<String>, password: impl Into<Password>) -> Self {
Self { email: email.into(), password: password.into(), display_name: None }
}
pub fn with_display_name(mut self, name: impl Into<String>) -> Self {
self.display_name = Some(name.into());
self
}
pub fn validate(&self) -> Result<(), ApiError> {
let mut details = ValidationDetails::new();
check_email(&self.email, &mut details);
check_password("password", &self.password, &mut details);
if let Some(name) = &self.display_name {
check_display_name(name, &mut details);
}
details.into_result()
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[non_exhaustive]
pub struct LoginRequest {
pub email: String,
pub password: Password,
}
impl LoginRequest {
pub fn new(email: impl Into<String>, password: impl Into<Password>) -> Self {
Self { email: email.into(), password: password.into() }
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[non_exhaustive]
pub struct SteamLoginRequest {
pub ticket_hex: Secret,
pub identity: String,
}
impl SteamLoginRequest {
pub fn new(ticket_hex: impl Into<Secret>, identity: impl Into<String>) -> Self {
Self { ticket_hex: ticket_hex.into(), identity: identity.into() }
}
pub fn validate(&self) -> Result<(), ApiError> {
let mut details = ValidationDetails::new();
let ticket = self.ticket_hex.expose();
if ticket.is_empty() || !ticket.len().is_multiple_of(2) || !ticket.bytes().all(|b| b.is_ascii_hexdigit()) {
details.add("ticket_hex", "is not hex-encoded bytes");
}
if ticket.len() > STEAM_TICKET_MAX_HEX {
details.add("ticket_hex", format!("is longer than {STEAM_TICKET_MAX_HEX} characters"));
}
if self.identity.trim().is_empty() {
details.add("identity", "is empty");
}
details.into_result()
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[non_exhaustive]
pub struct RefreshRequest {
pub refresh_token: RefreshToken,
}
impl RefreshRequest {
pub fn new(refresh_token: impl Into<RefreshToken>) -> Self {
Self { refresh_token: refresh_token.into() }
}
}
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
#[non_exhaustive]
pub struct LogoutRequest {
#[serde(default)]
pub everywhere: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub refresh_token: Option<RefreshToken>,
}
impl LogoutRequest {
pub fn this_session() -> Self {
Self { everywhere: false, refresh_token: None }
}
pub fn everywhere() -> Self {
Self { everywhere: true, refresh_token: None }
}
pub fn with_refresh_token(mut self, refresh_token: impl Into<RefreshToken>) -> Self {
self.refresh_token = Some(refresh_token.into());
self
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[non_exhaustive]
pub struct TokenPair {
pub token_type: String,
pub access_token: AccessToken,
pub access_expires_at: UnixMillis,
pub refresh_token: RefreshToken,
pub refresh_expires_at: UnixMillis,
}
impl TokenPair {
pub fn new(access_token: AccessToken, access_expires_at: UnixMillis, refresh_token: RefreshToken, refresh_expires_at: UnixMillis) -> Self {
Self { token_type: "Bearer".into(), access_token, access_expires_at, refresh_token, refresh_expires_at }
}
pub fn authorization_header(&self) -> String {
format!("Bearer {}", self.access_token.expose())
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct LinkedIdentity {
pub provider: String,
pub subject: String,
}
impl LinkedIdentity {
pub fn new(provider: impl Into<String>, subject: impl Into<String>) -> Self {
Self { provider: provider.into(), subject: subject.into() }
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct Account {
pub id: UserId,
#[serde(default)]
pub email: Option<String>,
#[serde(default)]
pub email_verified: bool,
#[serde(default)]
pub display_name: Option<String>,
#[serde(default)]
pub roles: Vec<String>,
#[serde(default)]
pub identities: Vec<LinkedIdentity>,
pub created_at: UnixMillis,
}
impl Account {
pub fn new(id: UserId, created_at: UnixMillis) -> Self {
Self { id, email: None, email_verified: false, display_name: None, roles: Vec::new(), identities: Vec::new(), created_at }
}
pub fn with_email(mut self, email: impl Into<String>, verified: bool) -> Self {
self.email = Some(email.into());
self.email_verified = verified;
self
}
pub fn with_display_name(mut self, name: impl Into<String>) -> Self {
self.display_name = Some(name.into());
self
}
pub fn with_roles(mut self, roles: Vec<String>) -> Self {
self.roles = roles;
self
}
pub fn with_identities(mut self, identities: Vec<LinkedIdentity>) -> Self {
self.identities = identities;
self
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[non_exhaustive]
pub struct AuthSession {
pub account: Account,
pub tokens: TokenPair,
}
impl AuthSession {
pub fn new(account: Account, tokens: TokenPair) -> Self {
Self { account, tokens }
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct UpdateAccountRequest {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub display_name: Option<String>,
}
impl UpdateAccountRequest {
pub fn new() -> Self {
Self::default()
}
pub fn with_display_name(mut self, name: impl Into<String>) -> Self {
self.display_name = Some(name.into());
self
}
pub fn validate(&self) -> Result<(), ApiError> {
let mut details = ValidationDetails::new();
if let Some(name) = &self.display_name {
check_display_name(name, &mut details);
}
details.into_result()
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[non_exhaustive]
pub struct ChangePasswordRequest {
pub current_password: Password,
pub new_password: Password,
}
impl ChangePasswordRequest {
pub fn new(current_password: impl Into<Password>, new_password: impl Into<Password>) -> Self {
Self { current_password: current_password.into(), new_password: new_password.into() }
}
pub fn validate(&self) -> Result<(), ApiError> {
let mut details = ValidationDetails::new();
check_password("new_password", &self.new_password, &mut details);
details.into_result()
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[non_exhaustive]
pub struct VerifyEmailRequest {
pub token: Secret,
}
impl VerifyEmailRequest {
pub fn new(token: impl Into<Secret>) -> Self {
Self { token: token.into() }
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct ForgotPasswordRequest {
pub email: String,
}
impl ForgotPasswordRequest {
pub fn new(email: impl Into<String>) -> Self {
Self { email: email.into() }
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[non_exhaustive]
pub struct ResetPasswordRequest {
pub token: Secret,
pub new_password: Password,
}
impl ResetPasswordRequest {
pub fn new(token: impl Into<Secret>, new_password: impl Into<Password>) -> Self {
Self { token: token.into(), new_password: new_password.into() }
}
pub fn validate(&self) -> Result<(), ApiError> {
let mut details = ValidationDetails::new();
check_password("new_password", &self.new_password, &mut details);
details.into_result()
}
}
mod calls {
use super::*;
use crate::envelope::Ack;
use crate::http_call::{payload_call, HttpCall, NoPayload, PathParams, PayloadKind, NO_PAYLOAD};
use crate::routes::{self, HttpMethod, Route};
payload_call!(RegisterRequest, Post, routes::auth::REGISTER, false, Json, AuthSession);
payload_call!(LoginRequest, Post, routes::auth::LOGIN, false, Json, AuthSession);
payload_call!(SteamLoginRequest, Post, routes::auth::STEAM, false, Json, AuthSession);
payload_call!(RefreshRequest, Post, routes::auth::REFRESH, false, Json, TokenPair);
payload_call!(LogoutRequest, Post, routes::auth::LOGOUT, false, Json, Ack);
payload_call!(VerifyEmailRequest, Post, routes::auth::VERIFY_EMAIL, false, Json, Ack);
payload_call!(ForgotPasswordRequest, Post, routes::auth::FORGOT_PASSWORD, false, Json, Ack);
payload_call!(ResetPasswordRequest, Post, routes::auth::RESET_PASSWORD, false, Json, Ack);
payload_call!(UpdateAccountRequest, Patch, routes::account::ME, true, Json, Account);
payload_call!(ChangePasswordRequest, Post, routes::account::PASSWORD, true, Json, Ack);
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
#[non_exhaustive]
pub struct ResendVerification {}
impl ResendVerification {
pub const fn new() -> Self {
Self {}
}
}
impl HttpCall for ResendVerification {
type Payload = NoPayload;
type Response = Ack;
const ROUTE: Route = Route::new(HttpMethod::Post, routes::auth::RESEND_VERIFICATION, true);
const PAYLOAD: PayloadKind = PayloadKind::Empty;
fn payload(&self) -> &NoPayload {
&NO_PAYLOAD
}
fn from_parts(_params: &PathParams, _payload: NoPayload) -> Result<Self, ApiError> {
Ok(Self::new())
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
#[non_exhaustive]
pub struct GetAccount {}
impl GetAccount {
pub const fn new() -> Self {
Self {}
}
}
impl HttpCall for GetAccount {
type Payload = NoPayload;
type Response = Account;
const ROUTE: Route = Route::new(HttpMethod::Get, routes::account::ME, true);
const PAYLOAD: PayloadKind = PayloadKind::Empty;
fn payload(&self) -> &NoPayload {
&NO_PAYLOAD
}
fn from_parts(_params: &PathParams, _payload: NoPayload) -> Result<Self, ApiError> {
Ok(Self::new())
}
}
pub fn is_valid_provider(provider: &str) -> bool {
provider.len() <= 32
&& provider.bytes().next().is_some_and(|b| b.is_ascii_lowercase())
&& provider.bytes().all(|b| b.is_ascii_lowercase() || b.is_ascii_digit() || b == b'_')
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct UnlinkIdentity {
pub provider: String,
}
impl UnlinkIdentity {
pub fn new(provider: impl Into<String>) -> Self {
Self { provider: provider.into() }
}
}
impl HttpCall for UnlinkIdentity {
type Payload = NoPayload;
type Response = Ack;
const ROUTE: Route = Route::new(HttpMethod::Delete, routes::account::IDENTITY, true);
const PAYLOAD: PayloadKind = PayloadKind::Empty;
fn payload(&self) -> &NoPayload {
&NO_PAYLOAD
}
fn path_params(&self) -> PathParams {
PathParams::new().with("provider", &self.provider)
}
fn from_parts(params: &PathParams, _payload: NoPayload) -> Result<Self, ApiError> {
Ok(Self::new(params.checked("provider", is_valid_provider, "is not a provider name")?))
}
}
}
pub use calls::{is_valid_provider, GetAccount, ResendVerification, UnlinkIdentity};
#[cfg(test)]
mod tests {
use super::*;
use crate::error::codes;
#[test]
fn register_validation() {
assert!(RegisterRequest::new("ada@example.com", "correct horse battery").validate().is_ok());
let error = RegisterRequest::new("not-an-email", "short").with_display_name(" ").validate().err();
let details = error.as_ref().and_then(|e| e.details_as::<ValidationDetails>()).unwrap_or_default();
assert!(error.is_some_and(|e| e.is(codes::VALIDATION_FAILED)));
assert_eq!(details.fields.keys().map(String::as_str).collect::<Vec<_>>(), ["display_name", "email", "password"]);
assert!(RegisterRequest::new("a b@example.com", "correct horse battery").validate().is_err());
assert!(RegisterRequest::new("a@b@example.com", "correct horse battery").validate().is_err());
assert!(RegisterRequest::new("x<victim@example.com>", "correct horse battery").validate().is_err());
let long = "x".repeat(PASSWORD_MAX_BYTES + 1);
assert!(RegisterRequest::new("ada@example.com", long).validate().is_err());
assert!(RegisterRequest::new("ada@example.com", "correct horse battery").with_display_name("x".repeat(33)).validate().is_err());
}
#[test]
fn plain_addresses_only() {
for ok in ["ada@example.com", " ada@example.com ", "a.b+tag@sub.example.co", "o'neil@example.com", "zoë@exämple.de", "a@b", "x_y-z@1.example"] {
assert!(is_valid_email(ok), "{ok}");
}
for bad in [
"x<victim@example.com>",
"<victim@example.com>",
"Ada <ada@example.com>",
"\"a\"@example.com",
"ada(comment)@example.com",
"ada@[192.0.2.1]",
"a,b@example.com",
"a;b@example.com",
"a:b@example.com",
"a\\b@example.com",
".ada@example.com",
"ada.@example.com",
"a..b@example.com",
"ada@example..com",
"ada@-example.com",
"ada@example-.com",
"ada@",
"@example.com",
"ada example@example.com",
"ada@@example.com",
] {
assert!(!is_valid_email(bad), "{bad}");
}
assert!(!is_valid_email(&format!("{}@example.com", "a".repeat(65))));
assert!(!is_valid_email(&format!("a@{}.com", "b".repeat(64))));
}
#[test]
fn steam_validation() {
assert!(SteamLoginRequest::new("0a1B", "my-game").validate().is_ok());
assert!(SteamLoginRequest::new("0a1", "my-game").validate().is_err());
assert!(SteamLoginRequest::new("zz", "my-game").validate().is_err());
assert!(SteamLoginRequest::new("", "my-game").validate().is_err());
assert!(SteamLoginRequest::new("00", " ").validate().is_err());
assert!(SteamLoginRequest::new("0".repeat(STEAM_TICKET_MAX_HEX + 2), "g").validate().is_err());
}
#[test]
fn other_validation() {
assert!(ChangePasswordRequest::new("old", "a new long password").validate().is_ok());
assert!(ChangePasswordRequest::new("old", "short").validate().is_err());
assert!(ResetPasswordRequest::new("tok", "short").validate().is_err());
assert!(UpdateAccountRequest::new().validate().is_ok());
assert!(UpdateAccountRequest::new().with_display_name("bad\nname").validate().is_err());
}
#[test]
fn secrets_are_redacted() {
let pair = TokenPair::new(AccessToken::new("acc-SECRET"), UnixMillis(1), RefreshToken::new("ref-SECRET"), UnixMillis(2));
let debug = format!("{pair:?} {:?}", LoginRequest::new("a@example.com", "pw-SECRET"));
assert!(!debug.contains("SECRET"), "{debug}");
assert!(debug.contains("AccessToken(<redacted>)") && debug.contains("Password(<redacted>)"));
assert_eq!(pair.authorization_header(), "Bearer acc-SECRET");
assert_eq!(pair.access_token.clone().into_inner(), "acc-SECRET");
assert!(!Secret::from(String::from("x")).is_empty());
}
}