use std::fmt;
use std::sync::Arc;
use bevy_ecs::resource::Resource;
use http::header::{HeaderName, HeaderValue, AUTHORIZATION};
use crate::request::OutgoingRequest;
#[derive(Clone, Default)]
pub struct Secret(String);
impl Drop for Secret {
fn drop(&mut self) {
self.wipe();
}
}
impl Secret {
pub(crate) fn wipe(&mut self) {
zeroize::Zeroize::zeroize(&mut self.0);
}
#[cfg(test)]
pub(crate) fn allocation(&mut self) -> (usize, usize, bool) {
let vec = unsafe { self.0.as_mut_vec() };
let zeroed = vec.spare_capacity_mut().iter().all(|b| unsafe { b.assume_init() } == 0);
(vec.len(), vec.capacity(), zeroed)
}
}
impl Secret {
pub fn new(secret: impl Into<String>) -> Self {
Self(secret.into())
}
pub fn expose(&self) -> &str {
&self.0
}
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
}
impl From<String> for Secret {
fn from(s: String) -> Self {
Self(s)
}
}
impl From<&str> for Secret {
fn from(s: &str) -> Self {
Self(s.to_string())
}
}
impl fmt::Debug for Secret {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("Secret(<redacted>)")
}
}
impl fmt::Display for Secret {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("<redacted>")
}
}
pub trait Credentials: Send + Sync + 'static {
fn apply(&self, request: &mut OutgoingRequest);
fn ws_auth_message(&self) -> Option<String> {
None
}
}
#[derive(Resource, Default, Clone)]
pub struct BackendCredentials {
inner: Option<Arc<dyn Credentials>>,
version: u64,
}
static NEXT_CREDENTIALS_VERSION: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
fn next_version() -> u64 {
NEXT_CREDENTIALS_VERSION.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
}
impl fmt::Debug for BackendCredentials {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BackendCredentials").field("set", &self.inner.is_some()).finish()
}
}
impl BackendCredentials {
pub fn new(credentials: impl Credentials) -> Self {
Self { inner: Some(Arc::new(credentials)), version: next_version() }
}
pub fn set(&mut self, credentials: impl Credentials) {
self.inner = Some(Arc::new(credentials));
self.version = next_version();
}
pub fn clear(&mut self) {
self.inner = None;
self.version = next_version();
}
pub fn is_set(&self) -> bool {
self.inner.is_some()
}
#[cfg_attr(not(feature = "ws"), allow(dead_code))]
pub(crate) fn version(&self) -> u64 {
self.version
}
pub(crate) fn apply(&self, request: &mut OutgoingRequest) {
if let Some(credentials) = &self.inner {
credentials.apply(request);
}
}
#[cfg(feature = "ws")]
pub(crate) fn ws_auth_message(&self) -> Option<String> {
self.inner.as_ref().and_then(|c| c.ws_auth_message())
}
}
#[derive(Clone, Debug)]
pub struct BearerToken(Secret);
impl BearerToken {
pub fn new(token: impl Into<Secret>) -> Self {
Self(token.into())
}
pub fn token(&self) -> &Secret {
&self.0
}
}
impl Credentials for BearerToken {
fn apply(&self, request: &mut OutgoingRequest) {
let token = self.0.expose();
let mut text = zeroize::Zeroizing::new(String::with_capacity(token.len().saturating_add(7)));
text.push_str("Bearer ");
text.push_str(token);
match sensitive_value(&text) {
Some(value) => {
request.headers_mut().insert(AUTHORIZATION, value);
}
None => request.reject("the bearer token contains characters that are not allowed in a header"),
}
}
}
#[derive(Clone, Debug)]
pub struct ApiKeyHeader {
name: String,
key: Secret,
}
impl ApiKeyHeader {
pub fn new(name: impl Into<String>, key: impl Into<Secret>) -> Self {
Self { name: name.into(), key: key.into() }
}
pub fn name(&self) -> &str {
&self.name
}
}
impl Credentials for ApiKeyHeader {
fn apply(&self, request: &mut OutgoingRequest) {
let Ok(name) = HeaderName::try_from(self.name.as_str()) else {
request.reject(format!("`{}` is not a valid header name", self.name));
return;
};
match sensitive_value(self.key.expose()) {
Some(value) => {
request.headers_mut().insert(name, value);
}
None => request.reject(format!("the key for header `{}` contains characters that are not allowed in a header", self.name)),
}
}
}
#[derive(Clone, Debug)]
pub struct ApiKeyQuery {
name: String,
key: Secret,
}
impl ApiKeyQuery {
pub fn new(name: impl Into<String>, key: impl Into<Secret>) -> Self {
Self { name: name.into(), key: key.into() }
}
pub fn name(&self) -> &str {
&self.name
}
}
impl Credentials for ApiKeyQuery {
fn apply(&self, request: &mut OutgoingRequest) {
let query = request.query_mut();
query.retain(|(name, _)| name != &self.name);
query.push((self.name.clone(), self.key.expose().to_string()));
}
}
#[cfg(feature = "json")]
#[cfg_attr(docsrs, doc(cfg(feature = "json")))]
#[derive(Clone, Debug)]
pub struct JsonBodyField {
name: String,
value: Secret,
}
#[cfg(feature = "json")]
impl JsonBodyField {
pub fn new(name: impl Into<String>, value: impl Into<Secret>) -> Self {
Self { name: name.into(), value: value.into() }
}
pub fn name(&self) -> &str {
&self.name
}
}
#[cfg(feature = "json")]
impl Credentials for JsonBodyField {
fn apply(&self, request: &mut OutgoingRequest) {
if request.purpose() == crate::RequestPurpose::WebSocketHandshake {
request.reject("JsonBodyField cannot authenticate a WebSocket handshake (it has no body); use first-message auth (Credentials::ws_auth_message)");
return;
}
if request.purpose() != crate::RequestPurpose::Http {
return;
}
if request.is_multipart() {
request.reject("JsonBodyField cannot authenticate a multipart upload (its body is a form, not JSON); use a header credential (BearerToken, ApiKeyHeader) or add the field to the form yourself");
return;
}
let Some(body) = request.body() else { return };
let Ok(serde_json::Value::Object(mut object)) = serde_json::from_slice::<serde_json::Value>(body) else {
return;
};
object.insert(self.name.clone(), serde_json::Value::String(self.value.expose().to_string()));
let encoded = crate::body::WipedBytes::json(&object);
let mut parsed = serde_json::Value::Object(object);
wipe_json(&mut parsed);
match encoded {
Ok(body) => request.set_wiped_body(body),
Err(_) => request.reject("could not re-encode the JSON body with the credential field"),
}
}
}
#[cfg(feature = "json")]
pub(crate) fn wipe_json(value: &mut serde_json::Value) {
use zeroize::Zeroize;
match value {
serde_json::Value::String(text) => text.zeroize(),
serde_json::Value::Array(items) => items.iter_mut().for_each(wipe_json),
serde_json::Value::Object(map) => {
for (mut key, mut item) in std::mem::take(map) {
key.zeroize();
wipe_json(&mut item);
}
}
serde_json::Value::Null | serde_json::Value::Bool(_) | serde_json::Value::Number(_) => {}
}
}
fn sensitive_value(text: &str) -> Option<HeaderValue> {
let mut value = HeaderValue::try_from(text).ok()?;
value.set_sensitive(true);
Some(value)
}