use crate::service::Transport;
use bytes::Bytes;
use std::future::Future;
use std::net::SocketAddr;
use std::pin::Pin;
pub const H2C_PREFACE: &[u8] = b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n";
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Preface {
Http1,
Http2,
NeedMore,
}
#[inline]
pub fn is_h2c_preface(buf: &[u8]) -> Preface {
let n = buf.len().min(H2C_PREFACE.len());
if buf[..n] != H2C_PREFACE[..n] {
return Preface::Http1;
}
if buf.len() >= H2C_PREFACE.len() {
Preface::Http2
} else {
Preface::NeedMore
}
}
pub trait H2Fallback {
fn handle(
&self,
io: Box<dyn Transport>,
buffered: Bytes,
peer: Option<SocketAddr>,
) -> Pin<Box<dyn Future<Output = ()>>>;
}
pub trait UpgradeConsumer {
fn handle(&self, upgraded: crate::service::Upgraded) -> Pin<Box<dyn Future<Output = ()>>>;
}
#[cfg(feature = "tls")]
#[derive(Clone, Debug)]
pub struct TlsConfig {
pub cert_chain: Vec<Vec<u8>>,
pub key_der: Vec<u8>,
pub alpn: Vec<Vec<u8>>,
}
#[cfg(feature = "tls")]
impl TlsConfig {
pub fn new(cert_chain: Vec<Vec<u8>>, key_der: Vec<u8>) -> Self {
Self {
cert_chain,
key_der,
alpn: vec![b"http/1.1".to_vec()],
}
}
pub fn with_h2(mut self) -> Self {
if !self.alpn.iter().any(|p| p == b"h2") {
self.alpn.insert(0, b"h2".to_vec());
}
self
}
pub fn server_config(&self) -> Result<std::sync::Arc<rustls::ServerConfig>, TlsError> {
use rustls::pki_types::{CertificateDer, PrivateKeyDer};
if self.cert_chain.is_empty() {
return Err(TlsError::NoCertificate);
}
let certs: Vec<CertificateDer<'static>> = self
.cert_chain
.iter()
.map(|c| CertificateDer::from(c.clone()))
.collect();
let key = PrivateKeyDer::try_from(self.key_der.clone())
.map_err(|_| TlsError::InvalidPrivateKey)?;
let provider = std::sync::Arc::new(rustls::crypto::ring::default_provider());
let mut cfg = rustls::ServerConfig::builder_with_provider(provider)
.with_safe_default_protocol_versions()
.map_err(|e| TlsError::Rustls(e.to_string()))?
.with_no_client_auth()
.with_single_cert(certs, key)
.map_err(|e| TlsError::Rustls(e.to_string()))?;
cfg.alpn_protocols = self.alpn.clone();
Ok(std::sync::Arc::new(cfg))
}
}
#[cfg(feature = "tls")]
#[derive(Debug, thiserror::Error)]
pub enum TlsError {
#[error("no certificate supplied")]
NoCertificate,
#[error("invalid private key")]
InvalidPrivateKey,
#[error("rustls: {0}")]
Rustls(String),
}
#[cfg(feature = "tls")]
#[inline]
pub fn negotiated_h2(conn: &rustls::ServerConnection) -> bool {
conn.alpn_protocol() == Some(b"h2")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn detects_http1() {
assert_eq!(is_h2c_preface(b"GET / HTTP/1.1\r\n"), Preface::Http1);
assert_eq!(is_h2c_preface(b"POST /x HTTP/1.1\r\n"), Preface::Http1);
}
#[test]
fn detects_the_preface() {
assert_eq!(is_h2c_preface(H2C_PREFACE), Preface::Http2);
let mut extended = H2C_PREFACE.to_vec();
extended.extend_from_slice(&[0, 0, 0, 4, 0, 0, 0, 0, 0]);
assert_eq!(is_h2c_preface(&extended), Preface::Http2);
}
#[test]
fn needs_more_on_a_short_prefix() {
assert_eq!(is_h2c_preface(b""), Preface::NeedMore);
assert_eq!(is_h2c_preface(b"P"), Preface::NeedMore);
assert_eq!(is_h2c_preface(b"PRI * HTTP/2"), Preface::NeedMore);
assert_eq!(
is_h2c_preface(&H2C_PREFACE[..H2C_PREFACE.len() - 1]),
Preface::NeedMore
);
}
#[test]
fn a_divergent_prefix_is_http1() {
assert_eq!(is_h2c_preface(b"PRX"), Preface::Http1);
assert_eq!(is_h2c_preface(b"PATCH / HTTP/1.1\r\n"), Preface::Http1);
assert_eq!(is_h2c_preface(b"PROPFIND / HTTP/1.1\r\n"), Preface::Http1);
let mut nearly = H2C_PREFACE.to_vec();
let last = nearly.len() - 1;
nearly[last] = b'X';
assert_eq!(is_h2c_preface(&nearly), Preface::Http1);
}
#[test]
fn every_prefix_length_is_classified_consistently() {
for i in 0..H2C_PREFACE.len() {
assert_eq!(
is_h2c_preface(&H2C_PREFACE[..i]),
Preface::NeedMore,
"prefix of length {i} must be NeedMore"
);
}
assert_eq!(is_h2c_preface(H2C_PREFACE), Preface::Http2);
}
#[cfg(feature = "tls")]
#[test]
fn tls_config_defaults_to_http11_alpn() {
let c = TlsConfig::new(vec![vec![1, 2, 3]], vec![4, 5, 6]);
assert_eq!(c.alpn, vec![b"http/1.1".to_vec()]);
}
#[cfg(feature = "tls")]
#[test]
fn with_h2_prefers_h2_and_is_idempotent() {
let c = TlsConfig::new(vec![vec![1]], vec![2]).with_h2();
assert_eq!(c.alpn, vec![b"h2".to_vec(), b"http/1.1".to_vec()]);
let c = c.with_h2();
assert_eq!(
c.alpn,
vec![b"h2".to_vec(), b"http/1.1".to_vec()],
"calling it twice must not duplicate the protocol"
);
}
#[cfg(feature = "tls")]
#[test]
fn server_config_rejects_an_empty_chain() {
let c = TlsConfig::new(vec![], vec![1, 2, 3]);
assert!(matches!(c.server_config(), Err(TlsError::NoCertificate)));
}
#[cfg(feature = "tls")]
#[test]
fn server_config_rejects_a_bogus_key() {
let c = TlsConfig::new(vec![vec![1, 2, 3]], vec![]);
assert!(c.server_config().is_err());
}
}