use base64ct::{Base64, Encoding};
use crate::error::{Result, fmt};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum AuthMode {
None,
Basic { username: String, password: String },
Bearer { token: String },
Verbatim { value: String },
}
impl AuthMode {
pub fn from_parts(
username: Option<&str>,
password: Option<&str>,
token: Option<&str>,
verbatim: Option<&str>,
) -> Result<Self> {
let basic_partial = username.is_some() ^ password.is_some();
if basic_partial {
return Err(fmt!(
ConfigError,
"Basic auth requires both \"username\" and \"password\""
));
}
let basic_set = username.is_some() && password.is_some();
let token_set = token.is_some();
let verbatim_set = verbatim.is_some();
let count = (basic_set as u8) + (token_set as u8) + (verbatim_set as u8);
if count > 1 {
return Err(fmt!(
ConfigError,
"Auth modes are mutually exclusive; pick at most one of (username/password), token, or auth"
));
}
if basic_set {
let user = username.unwrap();
let pass = password.unwrap();
if user.contains(':') {
return Err(fmt!(AuthError, "Basic auth username must not contain ':'"));
}
reject_control_bytes(user, "Basic auth username")?;
reject_control_bytes(pass, "Basic auth password")?;
return Ok(AuthMode::Basic {
username: user.to_string(),
password: pass.to_string(),
});
}
if let Some(t) = token {
reject_control_bytes(t, "Bearer token")?;
return Ok(AuthMode::Bearer {
token: t.to_string(),
});
}
if let Some(v) = verbatim {
reject_control_bytes(v, "verbatim auth value")?;
return Ok(AuthMode::Verbatim {
value: v.to_string(),
});
}
Ok(AuthMode::None)
}
pub(crate) fn validate(&self) -> Result<()> {
match self {
AuthMode::None => Ok(()),
AuthMode::Basic { username, password } => {
if username.contains(':') {
return Err(fmt!(AuthError, "Basic auth username must not contain ':'"));
}
reject_control_bytes(username, "Basic auth username")?;
reject_control_bytes(password, "Basic auth password")?;
Ok(())
}
AuthMode::Bearer { token } => reject_control_bytes(token, "Bearer token"),
AuthMode::Verbatim { value } => reject_control_bytes(value, "verbatim auth value"),
}
}
pub fn header_value(&self) -> Option<String> {
match self {
AuthMode::None => None,
AuthMode::Basic { username, password } => {
let pair = format!("{}:{}", username, password);
let encoded = Base64::encode_string(pair.as_bytes());
Some(format!("Basic {}", encoded))
}
AuthMode::Bearer { token } => Some(format!("Bearer {}", token)),
AuthMode::Verbatim { value } => Some(value.clone()),
}
}
}
fn reject_control_bytes(s: &str, what: &str) -> Result<()> {
if let Some(b) = s.bytes().find(|&b| b < 0x20 || b == 0x7F) {
return Err(fmt!(
AuthError,
"{} must not contain control byte 0x{:02X}",
what,
b
));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::ErrorCode;
#[test]
fn none_when_nothing_set() {
let m = AuthMode::from_parts(None, None, None, None).unwrap();
assert_eq!(m, AuthMode::None);
assert_eq!(m.header_value(), None);
}
#[test]
fn basic_header_format() {
let m = AuthMode::from_parts(Some("admin"), Some("quest"), None, None).unwrap();
assert_eq!(m.header_value().unwrap(), "Basic YWRtaW46cXVlc3Q=");
}
#[test]
fn bearer_header_format() {
let m = AuthMode::from_parts(None, None, Some("eyJhbGciOi"), None).unwrap();
assert_eq!(m.header_value().unwrap(), "Bearer eyJhbGciOi");
}
#[test]
fn verbatim_header_format() {
let m = AuthMode::from_parts(None, None, None, Some("Custom xyz")).unwrap();
assert_eq!(m.header_value().unwrap(), "Custom xyz");
}
#[test]
fn basic_partial_rejected() {
let err = AuthMode::from_parts(Some("u"), None, None, None).unwrap_err();
assert_eq!(err.code(), ErrorCode::ConfigError);
let err = AuthMode::from_parts(None, Some("p"), None, None).unwrap_err();
assert_eq!(err.code(), ErrorCode::ConfigError);
}
#[test]
fn mutually_exclusive() {
let err = AuthMode::from_parts(Some("u"), Some("p"), Some("t"), None).unwrap_err();
assert_eq!(err.code(), ErrorCode::ConfigError);
let err = AuthMode::from_parts(None, None, Some("t"), Some("v")).unwrap_err();
assert_eq!(err.code(), ErrorCode::ConfigError);
let err = AuthMode::from_parts(Some("u"), Some("p"), None, Some("v")).unwrap_err();
assert_eq!(err.code(), ErrorCode::ConfigError);
}
#[test]
fn token_with_newline_rejected() {
let err = AuthMode::from_parts(None, None, Some("a\nb"), None).unwrap_err();
assert_eq!(err.code(), ErrorCode::AuthError);
}
#[test]
fn verbatim_with_cr_rejected() {
let err = AuthMode::from_parts(None, None, None, Some("a\rb")).unwrap_err();
assert_eq!(err.code(), ErrorCode::AuthError);
}
#[test]
fn basic_username_with_colon_rejected() {
let err =
AuthMode::from_parts(Some("admin:override"), Some("realpass"), None, None).unwrap_err();
assert_eq!(err.code(), ErrorCode::AuthError);
}
#[test]
fn basic_username_with_control_byte_rejected() {
for bad in ["a\nb", "a\rb", "a\0b", "a\tb", "a\x7Fb", "a\x01b"] {
let err = AuthMode::from_parts(Some(bad), Some("p"), None, None).unwrap_err();
assert_eq!(
err.code(),
ErrorCode::AuthError,
"expected reject for {bad:?}"
);
}
}
#[test]
fn basic_password_with_control_byte_rejected() {
for bad in ["a\nb", "a\rb", "a\0b", "a\x7Fb", "a\x01b"] {
let err = AuthMode::from_parts(Some("u"), Some(bad), None, None).unwrap_err();
assert_eq!(
err.code(),
ErrorCode::AuthError,
"expected reject for {bad:?}"
);
}
}
#[test]
fn token_with_control_byte_rejected() {
for bad in ["a\0b", "a\x01b", "a\x7Fb"] {
let err = AuthMode::from_parts(None, None, Some(bad), None).unwrap_err();
assert_eq!(
err.code(),
ErrorCode::AuthError,
"expected reject for {bad:?}"
);
}
}
#[test]
fn verbatim_with_control_byte_rejected() {
for bad in ["a\0b", "a\x01b", "a\x7Fb"] {
let err = AuthMode::from_parts(None, None, None, Some(bad)).unwrap_err();
assert_eq!(
err.code(),
ErrorCode::AuthError,
"expected reject for {bad:?}"
);
}
}
#[test]
fn basic_high_bytes_accepted() {
let m = AuthMode::from_parts(Some("üser"), Some("päss"), None, None).unwrap();
assert!(matches!(m, AuthMode::Basic { .. }));
}
}