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 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>>,
}
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)) }
}
pub fn set(&mut self, credentials: impl Credentials) {
self.inner = Some(Arc::new(credentials));
}
pub fn clear(&mut self) {
self.inner = None;
}
pub fn is_set(&self) -> bool {
self.inner.is_some()
}
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) {
match sensitive_value(&format!("Bearer {}", self.0.expose())) {
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()));
match serde_json::to_vec(&object) {
Ok(body) => request.set_body(Some(body)),
Err(_) => request.reject("could not re-encode the JSON body with the credential field"),
}
}
}
fn sensitive_value(text: &str) -> Option<HeaderValue> {
let mut value = HeaderValue::try_from(text).ok()?;
value.set_sensitive(true);
Some(value)
}