use derec_proto::Protocol;
#[cfg(any(feature = "serde", target_arch = "wasm32"))]
use serde::{Deserialize, Serialize};
pub const MAX_TRANSPORT_URI_LEN: usize = 2048;
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
#[cfg_attr(
any(feature = "serde", target_arch = "wasm32"),
derive(Serialize, Deserialize)
)]
pub struct TransportProtocol {
pub uri: String,
#[cfg_attr(
any(feature = "serde", target_arch = "wasm32"),
serde(with = "protocol_as_i32")
)]
pub protocol: Protocol,
}
#[cfg(any(feature = "serde", target_arch = "wasm32"))]
mod protocol_as_i32 {
use super::Protocol;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
pub fn serialize<S: Serializer>(p: &Protocol, ser: S) -> Result<S::Ok, S::Error> {
i32::from(*p).serialize(ser)
}
pub fn deserialize<'de, D: Deserializer<'de>>(de: D) -> Result<Protocol, D::Error> {
let raw = i32::deserialize(de)?;
Protocol::try_from(raw)
.map_err(|_| serde::de::Error::custom(format!("unknown Protocol discriminant {raw}")))
}
}
impl TransportProtocol {
pub fn new(uri: impl Into<String>, protocol: Protocol) -> Self {
Self {
uri: uri.into(),
protocol,
}
}
pub fn validate(&self) -> Result<(), TransportValidationError> {
if self.uri.is_empty() {
return Err(TransportValidationError::EmptyUri);
}
if self.uri.len() > MAX_TRANSPORT_URI_LEN {
return Err(TransportValidationError::UriTooLong {
got: self.uri.len(),
limit: MAX_TRANSPORT_URI_LEN,
});
}
if self.uri.bytes().any(|b| b < 0x20 || b == 0x7F) {
return Err(TransportValidationError::ControlCharacters);
}
match self.protocol {
Protocol::Https => {
if self.uri.starts_with("https://") {
} else if cfg!(feature = "unsafe-http")
&& self.uri.starts_with("http://")
{
#[cfg(all(feature = "unsafe-http", feature = "logging"))]
tracing::warn!(
uri = %self.uri,
"accepting plaintext http:// transport URI — \
confidentiality and authenticity are NOT provided \
by the transport layer; use https:// in production",
);
} else {
return Err(TransportValidationError::SchemeMismatch {
expected: "https://",
protocol: self.protocol,
});
}
}
}
Ok(())
}
}
impl TryFrom<&str> for TransportProtocol {
type Error = TransportValidationError;
fn try_from(uri: &str) -> Result<Self, Self::Error> {
let tp = Self {
uri: uri.to_owned(),
protocol: Protocol::Https,
};
tp.validate()?;
Ok(tp)
}
}
impl TryFrom<String> for TransportProtocol {
type Error = TransportValidationError;
fn try_from(uri: String) -> Result<Self, Self::Error> {
let tp = Self {
uri,
protocol: Protocol::Https,
};
tp.validate()?;
Ok(tp)
}
}
impl From<TransportProtocol> for derec_proto::TransportProtocol {
fn from(tp: TransportProtocol) -> Self {
Self {
uri: tp.uri,
protocol: tp.protocol.into(),
}
}
}
impl From<&TransportProtocol> for derec_proto::TransportProtocol {
fn from(tp: &TransportProtocol) -> Self {
Self {
uri: tp.uri.clone(),
protocol: tp.protocol.into(),
}
}
}
impl TryFrom<derec_proto::TransportProtocol> for TransportProtocol {
type Error = TransportValidationError;
fn try_from(p: derec_proto::TransportProtocol) -> Result<Self, Self::Error> {
let protocol = Protocol::try_from(p.protocol)
.map_err(|_| TransportValidationError::UnknownProtocol(p.protocol))?;
let tp = Self {
uri: p.uri,
protocol,
};
tp.validate()?;
Ok(tp)
}
}
impl TryFrom<&derec_proto::TransportProtocol> for TransportProtocol {
type Error = TransportValidationError;
fn try_from(p: &derec_proto::TransportProtocol) -> Result<Self, Self::Error> {
let protocol = Protocol::try_from(p.protocol)
.map_err(|_| TransportValidationError::UnknownProtocol(p.protocol))?;
let tp = Self {
uri: p.uri.clone(),
protocol,
};
tp.validate()?;
Ok(tp)
}
}
pub trait TransportProtocolExt {
fn validate(&self) -> Result<(), TransportValidationError>;
}
impl TransportProtocolExt for derec_proto::TransportProtocol {
fn validate(&self) -> Result<(), TransportValidationError> {
TransportProtocol::try_from(self).map(|_| ())
}
}
pub trait IntoOwnTransport {
fn into_own_transport(self) -> Result<TransportProtocol, TransportValidationError>;
}
impl IntoOwnTransport for TransportProtocol {
fn into_own_transport(self) -> Result<TransportProtocol, TransportValidationError> {
self.validate()?;
Ok(self)
}
}
impl IntoOwnTransport for &str {
fn into_own_transport(self) -> Result<TransportProtocol, TransportValidationError> {
TransportProtocol::try_from(self)
}
}
impl IntoOwnTransport for String {
fn into_own_transport(self) -> Result<TransportProtocol, TransportValidationError> {
TransportProtocol::try_from(self)
}
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum TransportValidationError {
#[error("transport uri is empty")]
EmptyUri,
#[error("transport uri length {got} exceeds cap {limit} bytes — refusing to propagate")]
UriTooLong { got: usize, limit: usize },
#[error("transport uri contains control characters (bytes < 0x20 or = 0x7F are not allowed)")]
ControlCharacters,
#[error("unknown TransportProtocol.protocol discriminant: {0}")]
UnknownProtocol(i32),
#[error(
"transport uri must start with `{expected}` for protocol {protocol:?} \
— rejecting plaintext / mismatched scheme"
)]
SchemeMismatch {
expected: &'static str,
protocol: Protocol,
},
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn try_from_str_derives_https_and_validates() {
let tp = TransportProtocol::try_from("https://owner.example.com").unwrap();
assert_eq!(tp.uri, "https://owner.example.com");
assert_eq!(tp.protocol, Protocol::Https);
}
#[test]
fn try_from_string_takes_ownership_and_validates() {
let uri = String::from("https://owner.example.com");
let tp = TransportProtocol::try_from(uri).unwrap();
assert_eq!(tp.protocol, Protocol::Https);
}
#[cfg(feature = "unsafe-http")]
#[test]
fn try_from_str_accepts_plaintext_http_when_unsafe_http_enabled() {
let tp = TransportProtocol::try_from("http://owner.example.com").unwrap();
assert_eq!(tp.protocol, Protocol::Https);
}
#[cfg(not(feature = "unsafe-http"))]
#[test]
fn try_from_str_rejects_plaintext_http_by_default() {
assert!(matches!(
TransportProtocol::try_from("http://owner.example.com"),
Err(TransportValidationError::SchemeMismatch {
expected: "https://",
protocol: Protocol::Https,
})
));
}
#[test]
fn try_from_str_rejects_unsupported_scheme() {
assert!(matches!(
TransportProtocol::try_from("ws://owner.example.com"),
Err(TransportValidationError::SchemeMismatch {
expected: "https://",
protocol: Protocol::Https,
})
));
}
#[test]
fn validate_rejects_unsupported_scheme() {
let tp = TransportProtocol::new("ws://owner.example.com", Protocol::Https);
assert!(matches!(
tp.validate(),
Err(TransportValidationError::SchemeMismatch {
expected: "https://",
protocol: Protocol::Https,
})
));
}
#[test]
fn validate_rejects_control_characters() {
let tp = TransportProtocol::new("https://owner.example.com\n", Protocol::Https);
assert!(matches!(
tp.validate(),
Err(TransportValidationError::ControlCharacters)
));
}
#[test]
fn validate_rejects_oversize_uri() {
let oversize = format!("https://{}", "a".repeat(MAX_TRANSPORT_URI_LEN));
assert!(matches!(
TransportProtocol::try_from(oversize),
Err(TransportValidationError::UriTooLong { .. })
));
}
#[test]
fn try_from_proto_rejects_unknown_enum() {
let proto = derec_proto::TransportProtocol {
uri: "https://x".to_owned(),
protocol: 9999,
};
let res: Result<TransportProtocol, _> = (&proto).try_into();
assert!(matches!(
res,
Err(TransportValidationError::UnknownProtocol(9999))
));
}
#[test]
fn try_from_proto_also_runs_uri_validation() {
let proto = derec_proto::TransportProtocol {
uri: "ws://owner.example.com".to_owned(),
protocol: 0, };
let res: Result<TransportProtocol, _> = (&proto).try_into();
assert!(matches!(
res,
Err(TransportValidationError::SchemeMismatch {
expected: "https://",
..
})
));
}
#[test]
fn roundtrip_to_proto_and_back() {
let original = TransportProtocol::new("https://owner.example.com", Protocol::Https);
let proto: derec_proto::TransportProtocol = original.clone().into();
let back: TransportProtocol = proto.try_into().unwrap();
assert_eq!(original, back);
}
#[test]
fn into_own_transport_accepts_str_string_and_typed() {
let from_str = IntoOwnTransport::into_own_transport("https://owner.example.com").unwrap();
assert_eq!(from_str.uri, "https://owner.example.com");
let from_string =
IntoOwnTransport::into_own_transport(String::from("https://owner.example.com"))
.unwrap();
assert_eq!(from_string.uri, "https://owner.example.com");
let typed = TransportProtocol::new("https://owner.example.com", Protocol::Https);
let from_typed = IntoOwnTransport::into_own_transport(typed.clone()).unwrap();
assert_eq!(from_typed, typed);
}
#[test]
fn into_own_transport_revalidates_typed_value() {
let malformed = TransportProtocol::new("ws://owner.example.com", Protocol::Https);
assert!(matches!(
IntoOwnTransport::into_own_transport(malformed),
Err(TransportValidationError::SchemeMismatch { .. })
));
}
#[test]
fn into_own_transport_rejects_unsupported_str_scheme() {
assert!(matches!(
IntoOwnTransport::into_own_transport("ws://owner.example.com"),
Err(TransportValidationError::SchemeMismatch { .. })
));
}
}