use std::collections::VecDeque;
use std::fmt;
use std::io;
use std::time::Duration;
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
use tokio::net::TcpStream;
use tokio::time::timeout;
use super::framing::{CONTROL_LINK_SAME, LineBuffer, SCHEME_SEPARATOR, redact};
use super::{EncodedRequest, StreamOpen, Transport, TransportError, TransportProperties};
use crate::protocol::request::{BindSession, CreateSession, PROTOCOL_VERSION, TlcpRequest as _};
const CONTENT_TYPE: &str = "text/plain; charset=UTF-8";
const DEFAULT_PORT_PLAIN: u16 = 80;
const DEFAULT_PORT_TLS: u16 = 443;
const MAX_HEAD_BYTES: usize = 64 * 1024;
const MAX_HEADERS: usize = 64;
const MAX_CHUNK_HEADER_BYTES: usize = 8 * 1024;
const MAX_CONTROL_BODY_BYTES: usize = 1024 * 1024;
const READ_BUFFER_BYTES: usize = 16 * 1024;
const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(30);
const CONTROL_TIMEOUT: Duration = Duration::from_secs(30);
enum Conn {
Plain(TcpStream),
Tls(Box<tokio_native_tls::TlsStream<TcpStream>>),
}
impl Conn {
async fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
match self {
Self::Plain(stream) => stream.read(buf).await,
Self::Tls(stream) => stream.read(buf).await,
}
}
async fn write_all(&mut self, buf: &[u8]) -> io::Result<()> {
match self {
Self::Plain(stream) => stream.write_all(buf).await,
Self::Tls(stream) => stream.write_all(buf).await,
}
}
async fn flush(&mut self) -> io::Result<()> {
match self {
Self::Plain(stream) => stream.flush().await,
Self::Tls(stream) => stream.flush().await,
}
}
}
impl fmt::Debug for Conn {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Plain(_) => f.write_str("Conn::Plain"),
Self::Tls(_) => f.write_str("Conn::Tls"),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct Endpoint {
tls: bool,
host: String,
port: u16,
authority: String,
redacted: String,
}
#[cold]
#[inline(never)]
fn connect_error(target: &str, reason: impl Into<String>) -> TransportError {
TransportError::Connect {
target: target.to_owned(),
source: Box::new(io::Error::new(io::ErrorKind::InvalidInput, reason.into())),
}
}
fn bare_authority(authority: &str) -> &str {
let host = authority
.split_once('/')
.map_or(authority, |(head, _)| head);
host.rsplit_once('@').map_or(host, |(_, host)| host)
}
fn split_host_port(
authority: &str,
tls: bool,
redacted: &str,
) -> Result<(String, u16), TransportError> {
let default = if tls {
DEFAULT_PORT_TLS
} else {
DEFAULT_PORT_PLAIN
};
if let Some(rest) = authority.strip_prefix('[') {
let Some((host, tail)) = rest.split_once(']') else {
return Err(connect_error(redacted, "malformed IPv6 host: missing `]`"));
};
let port = match tail.strip_prefix(':') {
Some(port) => port
.parse()
.map_err(|_| connect_error(redacted, format!("invalid port `{port}`")))?,
None if tail.is_empty() => default,
None => return Err(connect_error(redacted, "malformed IPv6 authority")),
};
return Ok((host.to_owned(), port));
}
match authority.rsplit_once(':') {
Some((host, port)) if !host.is_empty() => {
let port = port
.parse()
.map_err(|_| connect_error(redacted, format!("invalid port `{port}`")))?;
Ok((host.to_owned(), port))
}
_ => Ok((authority.to_owned(), default)),
}
}
fn resolve(address: &str, control_link: Option<&str>) -> Result<Endpoint, TransportError> {
let address = address.trim();
let redacted = redact(address);
let Some((scheme, rest)) = address.split_once(SCHEME_SEPARATOR) else {
return Err(connect_error(
&redacted,
"server address must start with http://, https://, ws:// or wss://",
));
};
let tls = match scheme.to_ascii_lowercase().as_str() {
"http" | "ws" => false,
"https" | "wss" => true,
other => {
return Err(connect_error(
&redacted,
format!("unsupported scheme `{other}`; expected http, https, ws or wss"),
));
}
};
let configured = bare_authority(rest);
let authority = match control_link {
Some(link) if !link.trim().is_empty() && link.trim() != CONTROL_LINK_SAME => {
let link = link.trim();
let link = link
.split_once(SCHEME_SEPARATOR)
.map_or(link, |(_, rest)| rest);
bare_authority(link)
}
_ => configured,
};
let authority = authority.trim_end_matches('/');
if authority.is_empty() {
return Err(connect_error(&redacted, "server address has no host"));
}
let redacted_authority = redact(&format!("{scheme}{SCHEME_SEPARATOR}{authority}"));
let (host, port) = split_host_port(authority, tls, &redacted_authority)?;
Ok(Endpoint {
tls,
host,
port,
authority: authority.to_owned(),
redacted: redacted_authority,
})
}
#[derive(Debug)]
enum BodyFraming {
Length(u64),
Chunked(Chunked),
UntilClose,
}
#[derive(Debug)]
struct Head {
code: u16,
framing: BodyFraming,
leftover: Vec<u8>,
}
fn parse_head(buf: &[u8]) -> Result<Option<Head>, TransportError> {
let mut headers = [httparse::EMPTY_HEADER; MAX_HEADERS];
let mut response = httparse::Response::new(&mut headers);
match response.parse(buf) {
Ok(httparse::Status::Complete(head_len)) => {
let code = response
.code
.ok_or_else(|| TransportError::MalformedFrame {
reason: "HTTP response head had no status code".to_owned(),
})?;
let framing = framing_from_headers(response.headers)?;
let leftover = buf.get(head_len..).unwrap_or_default().to_vec();
Ok(Some(Head {
code,
framing,
leftover,
}))
}
Ok(httparse::Status::Partial) => Ok(None),
Err(error) => Err(TransportError::MalformedFrame {
reason: format!("invalid HTTP response head: {error}"),
}),
}
}
fn framing_from_headers(headers: &[httparse::Header]) -> Result<BodyFraming, TransportError> {
let mut chunked = false;
let mut length: Option<u64> = None;
for header in headers {
if header.name.eq_ignore_ascii_case("transfer-encoding") {
let value = String::from_utf8_lossy(header.value).to_ascii_lowercase();
if value.split(',').any(|coding| coding.trim() == "chunked") {
chunked = true;
}
} else if header.name.eq_ignore_ascii_case("content-length") {
let text =
std::str::from_utf8(header.value).map_err(|_| TransportError::MalformedFrame {
reason: "Content-Length was not ASCII".to_owned(),
})?;
let value = text
.trim()
.parse::<u64>()
.map_err(|_| TransportError::MalformedFrame {
reason: format!("invalid Content-Length `{text}`"),
})?;
length = Some(value);
}
}
if chunked {
Ok(BodyFraming::Chunked(Chunked::default()))
} else if let Some(value) = length {
Ok(BodyFraming::Length(value))
} else {
Ok(BodyFraming::UntilClose)
}
}
#[derive(Debug, Default, PartialEq, Eq)]
enum ChunkState {
#[default]
Size,
Data(u64),
DataCrlf,
Trailer,
Done,
}
#[derive(Debug, Default)]
struct Chunked {
state: ChunkState,
}
impl Chunked {
#[must_use]
#[inline]
fn is_done(&self) -> bool {
matches!(self.state, ChunkState::Done)
}
fn decode(&mut self, pending: &mut Vec<u8>, out: &mut Vec<u8>) -> Result<bool, TransportError> {
let mut pos = 0usize;
loop {
match self.state {
ChunkState::Size => {
let rest = pending.get(pos..).unwrap_or_default();
let Some(eol) = find_crlf(rest) else {
check_header_bound(rest.len())?;
break;
};
let line = rest.get(..eol).unwrap_or_default();
let size = parse_chunk_size(line)?;
pos = pos.saturating_add(eol).saturating_add(2);
self.state = if size == 0 {
ChunkState::Trailer
} else {
ChunkState::Data(size)
};
}
ChunkState::Data(remaining) => {
let rest = pending.get(pos..).unwrap_or_default();
let take = remaining.min(rest.len() as u64);
let taken = usize::try_from(take).unwrap_or(rest.len());
if let Some(chunk) = rest.get(..taken) {
out.extend_from_slice(chunk);
}
pos = pos.saturating_add(taken);
let left = remaining.saturating_sub(take);
if left == 0 {
self.state = ChunkState::DataCrlf;
} else {
self.state = ChunkState::Data(left);
break;
}
}
ChunkState::DataCrlf => {
let rest = pending.get(pos..).unwrap_or_default();
if rest.len() < 2 {
break;
}
pos = pos.saturating_add(2);
self.state = ChunkState::Size;
}
ChunkState::Trailer => {
let rest = pending.get(pos..).unwrap_or_default();
let Some(eol) = find_crlf(rest) else {
check_header_bound(rest.len())?;
break;
};
if eol == 0 {
pos = pos.saturating_add(2);
self.state = ChunkState::Done;
} else {
pos = pos.saturating_add(eol).saturating_add(2);
}
}
ChunkState::Done => break,
}
}
pending.drain(..pos.min(pending.len()));
Ok(self.is_done())
}
}
fn check_header_bound(held: usize) -> Result<(), TransportError> {
if held > MAX_CHUNK_HEADER_BYTES {
return Err(TransportError::Capacity {
limit_bytes: MAX_CHUNK_HEADER_BYTES,
reason: format!("{held} bytes of a chunk framing line with no terminator"),
});
}
Ok(())
}
#[must_use]
fn find_crlf(bytes: &[u8]) -> Option<usize> {
bytes.windows(2).position(|pair| pair == b"\r\n")
}
fn parse_chunk_size(line: &[u8]) -> Result<u64, TransportError> {
let hex = match line.iter().position(|&byte| byte == b';') {
Some(index) => line.get(..index).unwrap_or_default(),
None => line,
};
let text = std::str::from_utf8(hex)
.map_err(|_| TransportError::MalformedFrame {
reason: "chunk size was not ASCII".to_owned(),
})?
.trim();
u64::from_str_radix(text, 16).map_err(|_| TransportError::MalformedFrame {
reason: format!("invalid chunk size `{text}`"),
})
}
#[derive(Debug)]
enum BodyRead {
Progress(Vec<u8>),
Ended(Vec<u8>),
}
fn decode_body(
framing: &mut BodyFraming,
pending: &mut Vec<u8>,
out: &mut Vec<u8>,
) -> Result<bool, TransportError> {
match framing {
BodyFraming::Length(remaining) => {
let take = (*remaining).min(pending.len() as u64);
let taken = usize::try_from(take).unwrap_or(pending.len());
if let Some(chunk) = pending.get(..taken) {
out.extend_from_slice(chunk);
}
pending.drain(..taken.min(pending.len()));
*remaining = remaining.saturating_sub(take);
Ok(*remaining == 0)
}
BodyFraming::UntilClose => {
out.append(pending);
Ok(false)
}
BodyFraming::Chunked(chunked) => chunked.decode(pending, out),
}
}
#[derive(Debug)]
struct StreamConn {
socket: Conn,
framing: BodyFraming,
pending: Vec<u8>,
}
impl StreamConn {
async fn read_chunk(&mut self) -> Result<BodyRead, TransportError> {
let mut out = Vec::new();
if decode_body(&mut self.framing, &mut self.pending, &mut out)? {
return Ok(BodyRead::Ended(out));
}
if !out.is_empty() {
return Ok(BodyRead::Progress(out));
}
let mut buffer = [0u8; READ_BUFFER_BYTES];
let read = self.socket.read(&mut buffer).await.map_err(|error| {
TransportError::ConnectionLost {
reason: error.to_string(),
}
})?;
if read == 0 {
return match &self.framing {
BodyFraming::UntilClose | BodyFraming::Length(0) => Ok(BodyRead::Ended(Vec::new())),
BodyFraming::Length(_) => Err(TransportError::ConnectionLost {
reason: "stream connection closed before its Content-Length".to_owned(),
}),
BodyFraming::Chunked(chunked) if chunked.is_done() => {
Ok(BodyRead::Ended(Vec::new()))
}
BodyFraming::Chunked(_) => Err(TransportError::ConnectionLost {
reason: "chunked stream connection closed mid-body".to_owned(),
}),
};
}
if let Some(slice) = buffer.get(..read) {
self.pending.extend_from_slice(slice);
}
let complete = decode_body(&mut self.framing, &mut self.pending, &mut out)?;
Ok(if complete {
BodyRead::Ended(out)
} else {
BodyRead::Progress(out)
})
}
}
async fn dial(endpoint: &Endpoint) -> Result<Conn, TransportError> {
let connect = |source: Box<dyn std::error::Error + Send + Sync>| TransportError::Connect {
target: endpoint.redacted.clone(),
source,
};
let tcp = TcpStream::connect((endpoint.host.as_str(), endpoint.port))
.await
.map_err(|error| connect(Box::new(error)))?;
let _ = tcp.set_nodelay(true);
if endpoint.tls {
let connector =
native_tls::TlsConnector::new().map_err(|error| connect(Box::new(error)))?;
let connector = tokio_native_tls::TlsConnector::from(connector);
let tls = connector
.connect(&endpoint.host, tcp)
.await
.map_err(|error| connect(Box::new(error)))?;
Ok(Conn::Tls(Box::new(tls)))
} else {
Ok(Conn::Plain(tcp))
}
}
async fn write_request(
socket: &mut Conn,
target: &str,
authority: &str,
body: &str,
) -> io::Result<()> {
let head = format!(
"POST {target} HTTP/1.1\r\n\
Host: {authority}\r\n\
Content-Type: {CONTENT_TYPE}\r\n\
Content-Length: {}\r\n\
Connection: close\r\n\
\r\n",
body.len()
);
socket.write_all(head.as_bytes()).await?;
socket.write_all(body.as_bytes()).await?;
socket.flush().await
}
async fn read_head(socket: &mut Conn) -> Result<Head, TransportError> {
let mut buffer = Vec::new();
let mut scratch = [0u8; READ_BUFFER_BYTES];
loop {
if let Some(head) = parse_head(&buffer)? {
return Ok(head);
}
if buffer.len() > MAX_HEAD_BYTES {
return Err(TransportError::Capacity {
limit_bytes: MAX_HEAD_BYTES,
reason: format!(
"{} bytes of an unterminated HTTP response head",
buffer.len()
),
});
}
let read =
socket
.read(&mut scratch)
.await
.map_err(|error| TransportError::ConnectionLost {
reason: error.to_string(),
})?;
if read == 0 {
return Err(TransportError::ConnectionLost {
reason: "connection closed before the HTTP response head".to_owned(),
});
}
if let Some(slice) = scratch.get(..read) {
buffer.extend_from_slice(slice);
}
}
}
async fn open_post(
endpoint: &Endpoint,
target: &str,
body: &str,
) -> Result<(Conn, Head), TransportError> {
let mut socket = dial(endpoint).await?;
write_request(&mut socket, target, &endpoint.authority, body)
.await
.map_err(|error| TransportError::ConnectionLost {
reason: error.to_string(),
})?;
let head = read_head(&mut socket).await?;
Ok((socket, head))
}
#[must_use]
#[inline]
const fn is_success(code: u16) -> bool {
code >= 200 && code < 300
}
fn split_body_lines(body: &[u8]) -> Result<Vec<String>, TransportError> {
let mut lines = Vec::new();
for segment in body.split(|&byte| byte == b'\n') {
let line = segment.strip_suffix(b"\r").unwrap_or(segment);
if line.is_empty() {
continue;
}
let line = std::str::from_utf8(line).map_err(|error| TransportError::MalformedFrame {
reason: format!("a control response line was not valid UTF-8: {error}"),
})?;
lines.push(line.to_owned());
}
Ok(lines)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum HttpMode {
Streaming,
Polling,
}
pub(crate) struct HttpTransport {
mode: HttpMode,
address: String,
control_link: Option<String>,
stream: Option<StreamConn>,
lines: LineBuffer,
control_lines: VecDeque<String>,
pending_final: Option<Result<String, TransportError>>,
}
impl fmt::Debug for HttpTransport {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("HttpTransport")
.field("mode", &self.mode)
.field("target", &redact(&self.address))
.field("control_link", &self.control_link.as_deref().map(redact))
.field("stream_open", &self.stream.is_some())
.field("queued_control_lines", &self.control_lines.len())
.finish()
}
}
impl HttpTransport {
#[must_use = "a transport does nothing until a stream is opened"]
pub(crate) fn try_new(
address: impl Into<String>,
mode: HttpMode,
) -> Result<Self, TransportError> {
let address = address.into();
resolve(&address, None)?;
Ok(Self {
mode,
address,
control_link: None,
stream: None,
lines: LineBuffer::default(),
control_lines: VecDeque::new(),
pending_final: None,
})
}
fn endpoint(&self) -> Result<Endpoint, TransportError> {
resolve(&self.address, self.control_link.as_deref())
}
fn reset_stream(&mut self) {
self.stream = None;
self.lines.clear();
self.pending_final = None;
}
fn take_final_line(&mut self) {
self.pending_final = self.lines.flush_partial();
}
async fn control_round_trip(
&self,
request: &EncodedRequest,
) -> Result<Vec<String>, TransportError> {
let endpoint = self.endpoint().map_err(|error| TransportError::Send {
name: request.name,
reason: error.to_string(),
})?;
let target = format!("{}?LS_protocol={PROTOCOL_VERSION}", request.path);
let (mut socket, head) = open_post(&endpoint, &target, &request.parameters)
.await
.map_err(|error| TransportError::Send {
name: request.name,
reason: error.to_string(),
})?;
if !is_success(head.code) {
return Err(TransportError::Send {
name: request.name,
reason: format!(
"server answered the control request with HTTP {}",
head.code
),
});
}
let mut framing = head.framing;
let mut pending = head.leftover;
let mut body = Vec::new();
loop {
let mut out = Vec::new();
let complete = match decode_body(&mut framing, &mut pending, &mut out) {
Ok(complete) => complete,
Err(error) => {
return Err(TransportError::Send {
name: request.name,
reason: error.to_string(),
});
}
};
body.extend_from_slice(&out);
if body.len() > MAX_CONTROL_BODY_BYTES {
return Err(TransportError::Send {
name: request.name,
reason: format!(
"control response exceeded this client's limit of {MAX_CONTROL_BODY_BYTES} bytes"
),
});
}
if complete {
break;
}
if out.is_empty() {
let mut scratch = [0u8; READ_BUFFER_BYTES];
let read =
socket
.read(&mut scratch)
.await
.map_err(|error| TransportError::Send {
name: request.name,
reason: error.to_string(),
})?;
if read == 0 {
break;
}
if let Some(slice) = scratch.get(..read) {
pending.extend_from_slice(slice);
}
}
}
split_body_lines(&body).map_err(|error| TransportError::Send {
name: request.name,
reason: error.to_string(),
})
}
}
impl Transport for HttpTransport {
fn properties(&self) -> TransportProperties {
match self.mode {
HttpMode::Streaming => TransportProperties {
control_shares_stream: false,
ends_on_content_length: true,
is_polling: false,
},
HttpMode::Polling => TransportProperties {
control_shares_stream: false,
ends_on_content_length: false,
is_polling: true,
},
}
}
fn set_control_link(&mut self, host: Option<&str>) {
self.control_link = match host {
Some(host) if !host.trim().is_empty() && host.trim() != CONTROL_LINK_SAME => {
Some(host.trim().to_owned())
}
_ => None,
};
tracing::debug!(
control_link = self
.control_link
.as_deref()
.map(redact)
.as_deref()
.unwrap_or("<configured address>"),
"control link recorded; applies to the next HTTP connection"
);
}
async fn open_stream(&mut self, request: StreamOpen) -> Result<(), TransportError> {
let (target, body) = match &request {
StreamOpen::Create(create) => {
let body = create.http_body().map_err(|_| TransportError::Send {
name: CreateSession::NAME,
reason: "session creation parameters cannot be encoded".to_owned(),
})?;
(CreateSession::http_target(), body)
}
StreamOpen::Bind(bind) => {
let body = bind.http_body().map_err(|error| TransportError::Send {
name: BindSession::NAME,
reason: error.to_string(),
})?;
(BindSession::http_target(), body)
}
};
let endpoint = self.endpoint()?;
self.reset_stream();
let (socket, head) = timeout(HANDSHAKE_TIMEOUT, open_post(&endpoint, &target, &body))
.await
.map_err(|_| {
connect_error(&endpoint.redacted, "stream connection handshake timed out")
})??;
if !is_success(head.code) {
return Err(connect_error(
&endpoint.redacted,
format!("server answered the stream request with HTTP {}", head.code),
));
}
self.stream = Some(StreamConn {
socket,
framing: head.framing,
pending: head.leftover,
});
Ok(())
}
async fn next_line(&mut self) -> Option<Result<String, TransportError>> {
loop {
if let Some(line) = self.control_lines.pop_front() {
return Some(Ok(line));
}
if let Some(line) = self.lines.pop_line() {
return Some(line);
}
if let Some(line) = self.pending_final.take() {
return Some(line);
}
let has_stream = self.stream.is_some();
if !has_stream {
return None;
}
let outcome = match self.stream.as_mut() {
Some(stream) => stream.read_chunk().await,
None => return None,
};
match outcome {
Ok(BodyRead::Progress(bytes)) => {
if let Err(error) = self.lines.push_bytes(&bytes) {
self.reset_stream();
return Some(Err(error));
}
}
Ok(BodyRead::Ended(bytes)) => {
let pushed = self.lines.push_bytes(&bytes);
self.stream = None;
if let Err(error) = pushed {
self.reset_stream();
return Some(Err(error));
}
self.take_final_line();
}
Err(error) => {
self.reset_stream();
return Some(Err(error));
}
}
}
}
async fn send_control(&mut self, request: EncodedRequest) -> Result<(), TransportError> {
let lines = timeout(CONTROL_TIMEOUT, self.control_round_trip(&request))
.await
.map_err(|_| TransportError::Send {
name: request.name,
reason: "control request timed out".to_owned(),
})??;
self.control_lines.extend(lines);
Ok(())
}
async fn close(&mut self) -> Result<(), TransportError> {
self.reset_stream();
self.control_lines.clear();
Ok(())
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::unwrap_used)]
#![allow(clippy::indexing_slicing)]
use super::*;
use std::net::SocketAddr;
use tokio::net::TcpListener;
#[test]
fn test_resolve_defaults_the_port_by_scheme() {
let http = resolve("http://push.example.com", None).unwrap();
assert_eq!(
(http.tls, http.host.as_str(), http.port),
(false, "push.example.com", 80)
);
let https = resolve("https://push.example.com", None).unwrap();
assert_eq!(
(https.tls, https.host.as_str(), https.port),
(true, "push.example.com", 443)
);
}
#[test]
fn test_resolve_accepts_ws_schemes_as_synonyms() {
assert!(!resolve("ws://h", None).unwrap().tls);
assert!(resolve("wss://h", None).unwrap().tls);
}
#[test]
fn test_resolve_reads_an_explicit_port() {
let endpoint = resolve("http://push.example.com:8080/lightstreamer", None).unwrap();
assert_eq!(
(endpoint.host.as_str(), endpoint.port),
("push.example.com", 8080)
);
assert_eq!(endpoint.authority, "push.example.com:8080");
}
#[test]
fn test_resolve_rejects_a_missing_or_unknown_scheme() {
assert!(matches!(
resolve("push.example.com", None),
Err(TransportError::Connect { .. })
));
assert!(matches!(
resolve("ftp://push.example.com", None),
Err(TransportError::Connect { .. })
));
}
#[test]
fn test_resolve_rejects_a_bad_port() {
assert!(matches!(
resolve("http://push.example.com:notaport", None),
Err(TransportError::Connect { .. })
));
}
#[test]
fn test_resolve_control_link_replaces_authority_and_keeps_tls() {
let endpoint = resolve("https://push.example.com", Some("node7.example.com:8443")).unwrap();
assert!(endpoint.tls, "the TLS scheme is inherited, not the link's");
assert_eq!(
(endpoint.host.as_str(), endpoint.port),
("node7.example.com", 8443)
);
}
#[test]
fn test_resolve_control_link_star_keeps_the_configured_address() {
let endpoint = resolve("https://push.example.com", Some(CONTROL_LINK_SAME)).unwrap();
assert_eq!(endpoint.host, "push.example.com");
}
#[test]
fn test_resolve_control_link_cannot_downgrade_to_cleartext() {
let endpoint =
resolve("https://push.example.com", Some("http://node7.example.com")).unwrap();
assert!(endpoint.tls);
assert_eq!(endpoint.host, "node7.example.com");
}
#[test]
fn test_resolve_strips_userinfo_from_the_redacted_target() {
let endpoint = resolve("https://user:secret@push.example.com", None).unwrap();
assert!(
!endpoint.redacted.contains("secret"),
"{}",
endpoint.redacted
);
assert_eq!(endpoint.host, "push.example.com");
}
#[test]
fn test_resolve_handles_a_bracketed_ipv6_literal() {
let endpoint = resolve("http://[::1]:8080", None).unwrap();
assert_eq!((endpoint.host.as_str(), endpoint.port), ("::1", 8080));
let no_port = resolve("https://[2001:db8::1]", None).unwrap();
assert_eq!((no_port.host.as_str(), no_port.port), ("2001:db8::1", 443));
}
#[test]
fn test_parse_head_is_partial_until_the_blank_line() {
assert!(
parse_head(b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n")
.unwrap()
.is_none()
);
}
#[test]
fn test_parse_head_reads_status_and_content_length() {
let head = parse_head(b"HTTP/1.1 200 OK\r\nContent-Length: 9\r\n\r\nREQOK,1\r\n")
.unwrap()
.unwrap();
assert_eq!(head.code, 200);
assert!(matches!(head.framing, BodyFraming::Length(9)));
assert_eq!(head.leftover, b"REQOK,1\r\n");
}
#[test]
fn test_parse_head_detects_chunked() {
let head = parse_head(b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n")
.unwrap()
.unwrap();
assert!(matches!(head.framing, BodyFraming::Chunked(_)));
}
#[test]
fn test_parse_head_without_framing_is_until_close() {
let head = parse_head(b"HTTP/1.1 200 OK\r\n\r\n").unwrap().unwrap();
assert!(matches!(head.framing, BodyFraming::UntilClose));
}
#[test]
fn test_parse_head_rejects_a_non_numeric_content_length() {
assert!(matches!(
parse_head(b"HTTP/1.1 200 OK\r\nContent-Length: eight\r\n\r\n"),
Err(TransportError::MalformedFrame { .. })
));
}
#[test]
fn test_chunked_decodes_a_whole_body() {
let mut chunked = Chunked::default();
let mut pending = b"7\r\nREQOK,1\r\n0\r\n\r\n".to_vec();
let mut out = Vec::new();
let done = chunked.decode(&mut pending, &mut out).unwrap();
assert!(done);
assert_eq!(out, b"REQOK,1");
assert!(pending.is_empty());
}
#[test]
fn test_chunked_reassembles_across_reads() {
let mut chunked = Chunked::default();
let mut out = Vec::new();
let mut pending = b"5\r\nU,1,".to_vec();
assert!(!chunked.decode(&mut pending, &mut out).unwrap());
assert_eq!(out, b"U,1,");
pending.extend_from_slice(b"1\r\n0\r\n\r\n");
assert!(chunked.decode(&mut pending, &mut out).unwrap());
assert_eq!(out, b"U,1,1");
}
#[test]
fn test_chunked_handles_multiple_chunks() {
let mut chunked = Chunked::default();
let mut pending = b"3\r\nU,1\r\n4\r\n,1,x\r\n0\r\n\r\n".to_vec();
let mut out = Vec::new();
assert!(chunked.decode(&mut pending, &mut out).unwrap());
assert_eq!(out, b"U,1,1,x");
}
#[test]
fn test_chunked_ignores_a_chunk_extension() {
let mut chunked = Chunked::default();
let mut pending = b"7;ext=1\r\nREQOK,9\r\n0\r\n\r\n".to_vec();
let mut out = Vec::new();
assert!(chunked.decode(&mut pending, &mut out).unwrap());
assert_eq!(out, b"REQOK,9");
}
#[test]
fn test_chunked_rejects_an_invalid_size() {
let mut chunked = Chunked::default();
let mut pending = b"zz\r\n".to_vec();
let mut out = Vec::new();
assert!(matches!(
chunked.decode(&mut pending, &mut out),
Err(TransportError::MalformedFrame { .. })
));
}
#[test]
fn test_chunked_refuses_an_unterminated_size_line() {
let mut chunked = Chunked::default();
let mut pending = vec![b'a'; MAX_CHUNK_HEADER_BYTES + 1];
let mut out = Vec::new();
assert!(matches!(
chunked.decode(&mut pending, &mut out),
Err(TransportError::Capacity { .. })
));
}
#[test]
fn test_decode_body_length_stops_at_the_content_length() {
let mut framing = BodyFraming::Length(7);
let mut pending = b"REQOK,1\r\nspurious".to_vec();
let mut out = Vec::new();
assert!(decode_body(&mut framing, &mut pending, &mut out).unwrap());
assert_eq!(out, b"REQOK,1");
}
#[test]
fn test_decode_body_until_close_takes_everything() {
let mut framing = BodyFraming::UntilClose;
let mut pending = b"U,1,1,x\r\n".to_vec();
let mut out = Vec::new();
assert!(!decode_body(&mut framing, &mut pending, &mut out).unwrap());
assert_eq!(out, b"U,1,1,x\r\n");
assert!(pending.is_empty());
}
#[test]
fn test_split_body_lines_keeps_an_unterminated_last_line() {
assert_eq!(split_body_lines(b"REQOK,1").unwrap(), vec!["REQOK,1"]);
assert_eq!(
split_body_lines(b"REQERR,1,17,denied\r\n").unwrap(),
vec!["REQERR,1,17,denied"]
);
}
#[test]
fn test_split_body_lines_drops_empty_lines() {
assert_eq!(
split_body_lines(b"\r\nREQOK,1\r\n\r\n").unwrap(),
vec!["REQOK,1"]
);
}
#[test]
fn test_properties_declare_streaming_behaviour() {
let transport =
HttpTransport::try_new("https://push.example.com", HttpMode::Streaming).unwrap();
assert_eq!(
transport.properties(),
TransportProperties {
control_shares_stream: false,
ends_on_content_length: true,
is_polling: false,
}
);
}
#[test]
fn test_properties_declare_polling_behaviour() {
let transport =
HttpTransport::try_new("https://push.example.com", HttpMode::Polling).unwrap();
assert_eq!(
transport.properties(),
TransportProperties {
control_shares_stream: false,
ends_on_content_length: false,
is_polling: true,
}
);
}
#[test]
fn test_try_new_rejects_an_unusable_address() {
assert!(HttpTransport::try_new("push.example.com", HttpMode::Streaming).is_err());
}
#[test]
fn test_set_control_link_records_and_clears() {
let mut transport =
HttpTransport::try_new("https://push.example.com", HttpMode::Streaming).unwrap();
transport.set_control_link(Some("node4.example.com"));
assert_eq!(transport.control_link.as_deref(), Some("node4.example.com"));
transport.set_control_link(Some(CONTROL_LINK_SAME));
assert_eq!(transport.control_link, None);
transport.set_control_link(Some("node4.example.com"));
transport.set_control_link(None);
assert_eq!(transport.control_link, None);
}
#[tokio::test]
async fn test_close_without_a_stream_is_a_no_op() {
let mut transport =
HttpTransport::try_new("https://push.example.com", HttpMode::Streaming).unwrap();
assert!(transport.close().await.is_ok());
assert!(transport.next_line().await.is_none());
}
async fn bind_loopback() -> (TcpListener, SocketAddr) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
(listener, address)
}
async fn read_request(socket: &mut TcpStream) -> String {
let mut buffer = Vec::new();
let mut scratch = [0u8; 4096];
loop {
let head_end = buffer.windows(4).position(|w| w == b"\r\n\r\n");
if let Some(end) = head_end {
let head = String::from_utf8_lossy(&buffer[..end]);
let content_length = head
.lines()
.find_map(|line| {
line.strip_prefix("Content-Length: ")
.or_else(|| line.strip_prefix("content-length: "))
})
.and_then(|value| value.trim().parse::<usize>().ok())
.unwrap_or(0);
if buffer.len() >= end + 4 + content_length {
return String::from_utf8_lossy(&buffer).into_owned();
}
}
let read = socket.read(&mut scratch).await.unwrap();
if read == 0 {
return String::from_utf8_lossy(&buffer).into_owned();
}
buffer.extend_from_slice(&scratch[..read]);
}
}
#[tokio::test]
async fn test_streaming_reads_the_body_line_by_line_then_ends() {
let (listener, address) = bind_loopback().await;
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let request = read_request(&mut socket).await;
socket
.write_all(
b"HTTP/1.1 200 OK\r\nContent-Length: 53\r\n\r\n\
CONOK,S1,50000,5000,*\r\nSERVNAME,LS\r\nU,1,1,x\r\nLOOP,0\r\n",
)
.await
.unwrap();
socket.flush().await.unwrap();
drop(socket);
request
});
let mut transport =
HttpTransport::try_new(format!("http://{address}"), HttpMode::Streaming).unwrap();
transport
.open_stream(StreamOpen::Create(Box::default()))
.await
.unwrap();
assert_eq!(
transport.next_line().await.unwrap().unwrap(),
"CONOK,S1,50000,5000,*"
);
assert_eq!(transport.next_line().await.unwrap().unwrap(), "SERVNAME,LS");
assert_eq!(transport.next_line().await.unwrap().unwrap(), "U,1,1,x");
assert_eq!(transport.next_line().await.unwrap().unwrap(), "LOOP,0");
assert!(
transport.next_line().await.is_none(),
"the body ends the stream"
);
let request = server.await.unwrap();
assert!(
request.starts_with(
"POST /lightstreamer/create_session.txt?LS_protocol=TLCP-2.5.0 HTTP/1.1\r\n"
),
"unexpected request line: {request:?}"
);
assert!(
request.contains("LS_cid="),
"the body carries the request parameters"
);
}
#[tokio::test]
async fn test_content_length_bounded_body_flushes_its_last_line() {
let (listener, address) = bind_loopback().await;
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let _ = read_request(&mut socket).await;
socket
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 9\r\n\r\nU,1,1,abc")
.await
.unwrap();
socket.flush().await.unwrap();
drop(socket);
});
let mut transport =
HttpTransport::try_new(format!("http://{address}"), HttpMode::Streaming).unwrap();
transport
.open_stream(StreamOpen::Create(Box::default()))
.await
.unwrap();
assert_eq!(transport.next_line().await.unwrap().unwrap(), "U,1,1,abc");
assert!(transport.next_line().await.is_none());
server.await.unwrap();
}
#[tokio::test]
async fn test_polling_cycle_ends_the_stream() {
let (listener, address) = bind_loopback().await;
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let _ = read_request(&mut socket).await;
socket
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 9\r\n\r\nU,1,1,x\r\n")
.await
.unwrap();
socket.flush().await.unwrap();
drop(socket);
});
let mut transport =
HttpTransport::try_new(format!("http://{address}"), HttpMode::Polling).unwrap();
transport
.open_stream(StreamOpen::Create(Box::default()))
.await
.unwrap();
assert_eq!(transport.next_line().await.unwrap().unwrap(), "U,1,1,x");
assert!(
transport.next_line().await.is_none(),
"one poll cycle, then rebind"
);
server.await.unwrap();
}
#[tokio::test]
async fn test_streaming_decodes_a_chunked_body() {
let (listener, address) = bind_loopback().await;
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let _ = read_request(&mut socket).await;
socket
.write_all(
b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n\
17\r\nCONOK,S1,50000,5000,*\r\n\r\n8\r\nLOOP,0\r\n\r\n0\r\n\r\n",
)
.await
.unwrap();
socket.flush().await.unwrap();
drop(socket);
});
let mut transport =
HttpTransport::try_new(format!("http://{address}"), HttpMode::Streaming).unwrap();
transport
.open_stream(StreamOpen::Create(Box::default()))
.await
.unwrap();
assert_eq!(
transport.next_line().await.unwrap().unwrap(),
"CONOK,S1,50000,5000,*"
);
assert_eq!(transport.next_line().await.unwrap().unwrap(), "LOOP,0");
assert!(transport.next_line().await.is_none());
server.await.unwrap();
}
#[tokio::test]
async fn test_control_response_is_forwarded_onto_the_line_stream() {
let (stream_listener, stream_address) = bind_loopback().await;
let (control_listener, control_address) = bind_loopback().await;
let stream_server = tokio::spawn(async move {
let (mut socket, _) = stream_listener.accept().await.unwrap();
let _ = read_request(&mut socket).await;
socket
.write_all(
b"HTTP/1.1 200 OK\r\nContent-Length: 23\r\n\r\nCONOK,S1,50000,5000,*\r\n",
)
.await
.unwrap();
socket.flush().await.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
});
let control_server = tokio::spawn(async move {
let (mut socket, _) = control_listener.accept().await.unwrap();
let request = read_request(&mut socket).await;
socket
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 9\r\n\r\nREQOK,7\r\n")
.await
.unwrap();
socket.flush().await.unwrap();
drop(socket);
request
});
let mut transport =
HttpTransport::try_new(format!("http://{stream_address}"), HttpMode::Streaming)
.unwrap();
transport
.open_stream(StreamOpen::Create(Box::default()))
.await
.unwrap();
assert_eq!(
transport.next_line().await.unwrap().unwrap(),
"CONOK,S1,50000,5000,*"
);
transport.set_control_link(Some(&control_address.to_string()));
transport
.send_control(EncodedRequest {
name: "control",
path: "/lightstreamer/control.txt",
parameters: "LS_reqId=7&LS_op=add&LS_subId=1&LS_group=item1&LS_schema=x&LS_mode=MERGE&LS_session=S1".to_owned(),
})
.await
.unwrap();
assert_eq!(transport.next_line().await.unwrap().unwrap(), "REQOK,7");
let control_request = control_server.await.unwrap();
assert!(
control_request
.starts_with("POST /lightstreamer/control.txt?LS_protocol=TLCP-2.5.0 HTTP/1.1\r\n"),
"the control POST went to the control link: {control_request:?}"
);
assert!(control_request.contains("LS_reqId=7"));
stream_server.await.unwrap();
}
#[tokio::test]
async fn test_control_request_non_success_status_is_a_send_error() {
let (listener, address) = bind_loopback().await;
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let _ = read_request(&mut socket).await;
socket
.write_all(b"HTTP/1.1 503 Service Unavailable\r\nContent-Length: 0\r\n\r\n")
.await
.unwrap();
socket.flush().await.unwrap();
drop(socket);
});
let mut transport =
HttpTransport::try_new(format!("http://{address}"), HttpMode::Streaming).unwrap();
let error = transport
.send_control(EncodedRequest {
name: "control",
path: "/lightstreamer/control.txt",
parameters: "LS_reqId=1&LS_op=delete&LS_subId=1&LS_session=S1".to_owned(),
})
.await
.unwrap_err();
assert!(
matches!(error, TransportError::Send { .. }),
"got {error:?}"
);
server.await.unwrap();
}
#[tokio::test]
async fn test_oversized_response_head_is_refused() {
let (listener, address) = bind_loopback().await;
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let _ = read_request(&mut socket).await;
socket.write_all(b"HTTP/1.1 200 OK\r\n").await.unwrap();
let filler = vec![b'x'; MAX_HEAD_BYTES + 4096];
socket.write_all(b"X-Filler: ").await.unwrap();
let _ = socket.write_all(&filler).await;
let _ = socket.flush().await;
});
let mut transport =
HttpTransport::try_new(format!("http://{address}"), HttpMode::Streaming).unwrap();
let error = transport
.open_stream(StreamOpen::Create(Box::default()))
.await
.unwrap_err();
assert!(
matches!(error, TransportError::Capacity { .. }),
"got {error:?}"
);
drop(server);
}
#[tokio::test]
async fn test_stream_open_non_success_status_is_a_connect_error() {
let (listener, address) = bind_loopback().await;
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let _ = read_request(&mut socket).await;
socket
.write_all(b"HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\n\r\n")
.await
.unwrap();
socket.flush().await.unwrap();
drop(socket);
});
let mut transport =
HttpTransport::try_new(format!("http://{address}"), HttpMode::Streaming).unwrap();
let error = transport
.open_stream(StreamOpen::Create(Box::default()))
.await
.unwrap_err();
assert!(
matches!(error, TransportError::Connect { .. }),
"got {error:?}"
);
server.await.unwrap();
}
#[tokio::test]
async fn test_truncated_content_length_body_is_connection_lost() {
let (listener, address) = bind_loopback().await;
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let _ = read_request(&mut socket).await;
socket
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 100\r\n\r\nU,1,1,x\r\n")
.await
.unwrap();
socket.flush().await.unwrap();
drop(socket);
});
let mut transport =
HttpTransport::try_new(format!("http://{address}"), HttpMode::Streaming).unwrap();
transport
.open_stream(StreamOpen::Create(Box::default()))
.await
.unwrap();
assert_eq!(transport.next_line().await.unwrap().unwrap(), "U,1,1,x");
let error = transport.next_line().await.unwrap().unwrap_err();
assert!(
matches!(error, TransportError::ConnectionLost { .. }),
"got {error:?}"
);
server.await.unwrap();
}
}