use std::path::{Path, PathBuf};
use std::time::Duration;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::UnixStream;
pub const MAX_RESPONSE_BYTES: usize = 8 * 1024 * 1024;
pub const EXCHANGE_TIMEOUT: Duration = Duration::from_secs(30);
#[derive(Debug)]
pub enum TransportError {
NotListening { path: PathBuf, detail: String },
Timeout,
Io(String),
Protocol(String),
TooLarge,
}
impl std::fmt::Display for TransportError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::NotListening { path, detail } => write!(
f,
"no Conflux owner is listening at '{}': {detail}",
path.display()
),
Self::Timeout => write!(f, "the owner did not answer within {EXCHANGE_TIMEOUT:?}"),
Self::Io(detail) => write!(f, "the connection to the owner failed: {detail}"),
Self::Protocol(detail) => write!(f, "the owner's response was not usable: {detail}"),
Self::TooLarge => write!(
f,
"the owner's response exceeded the {MAX_RESPONSE_BYTES}-byte client limit"
),
}
}
}
impl std::error::Error for TransportError {}
#[derive(Debug, Clone)]
pub struct HttpResponse {
pub status: u16,
pub body: Vec<u8>,
}
impl HttpResponse {
pub fn json<T: serde::de::DeserializeOwned>(&self) -> Result<T, TransportError> {
serde_json::from_slice(&self.body).map_err(|error| {
TransportError::Protocol(format!("response body is not the expected JSON: {error}"))
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TokenRejection {
CarriageReturn,
LineFeed,
ControlCharacter,
Delete,
}
impl TokenRejection {
pub fn as_str(self) -> &'static str {
match self {
Self::CarriageReturn => "a carriage return (CR)",
Self::LineFeed => "a line feed (LF)",
Self::ControlCharacter => "an ASCII control character",
Self::Delete => "the DEL character",
}
}
}
impl std::fmt::Display for TokenRejection {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
pub fn validate_token(value: &str) -> Result<(), TokenRejection> {
for byte in value.bytes() {
match byte {
b'\r' => return Err(TokenRejection::CarriageReturn),
b'\n' => return Err(TokenRejection::LineFeed),
0x7F => return Err(TokenRejection::Delete),
0x00..=0x1F => return Err(TokenRejection::ControlCharacter),
_ => {}
}
}
Ok(())
}
#[derive(Clone)]
pub struct UnixApiClient {
socket: PathBuf,
token: Option<String>,
}
impl std::fmt::Debug for UnixApiClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("UnixApiClient")
.field("socket", &self.socket)
.field(
"token",
&if self.token.is_some() {
"<redacted>"
} else {
"<none>"
},
)
.finish()
}
}
impl UnixApiClient {
pub fn new(socket: PathBuf, token: Option<String>) -> Result<Self, TokenRejection> {
if let Some(token) = &token {
validate_token(token)?;
}
Ok(Self { socket, token })
}
pub fn socket(&self) -> &Path {
&self.socket
}
#[cfg_attr(not(test), allow(dead_code))]
pub fn has_token(&self) -> bool {
self.token.is_some()
}
pub async fn get(&self, path_and_query: &str) -> Result<HttpResponse, TransportError> {
self.request("GET", path_and_query, None).await
}
pub async fn post_json(
&self,
path_and_query: &str,
body: &str,
) -> Result<HttpResponse, TransportError> {
self.request("POST", path_and_query, Some(body)).await
}
pub async fn put_json(
&self,
path_and_query: &str,
body: &str,
) -> Result<HttpResponse, TransportError> {
self.request("PUT", path_and_query, Some(body)).await
}
pub async fn delete(&self, path_and_query: &str) -> Result<HttpResponse, TransportError> {
self.request("DELETE", path_and_query, None).await
}
async fn request(
&self,
method: &str,
path_and_query: &str,
body: Option<&str>,
) -> Result<HttpResponse, TransportError> {
let exchange = self.exchange(method, path_and_query, body);
match tokio::time::timeout(EXCHANGE_TIMEOUT, exchange).await {
Ok(result) => result,
Err(_elapsed) => Err(TransportError::Timeout),
}
}
async fn exchange(
&self,
method: &str,
path_and_query: &str,
body: Option<&str>,
) -> Result<HttpResponse, TransportError> {
let mut stream = UnixStream::connect(&self.socket).await.map_err(|error| {
TransportError::NotListening {
path: self.socket.clone(),
detail: error.to_string(),
}
})?;
let request = self.encode_request(method, path_and_query, body);
stream
.write_all(request.as_bytes())
.await
.map_err(|error| TransportError::Io(error.to_string()))?;
stream
.flush()
.await
.map_err(|error| TransportError::Io(error.to_string()))?;
let mut raw = Vec::new();
let mut chunk = [0u8; 16 * 1024];
loop {
let read = stream
.read(&mut chunk)
.await
.map_err(|error| TransportError::Io(error.to_string()))?;
if read == 0 {
break;
}
if raw.len() + read > MAX_RESPONSE_BYTES {
return Err(TransportError::TooLarge);
}
raw.extend_from_slice(&chunk[..read]);
}
parse_response(&raw)
}
pub fn encode_request(&self, method: &str, path_and_query: &str, body: Option<&str>) -> String {
let mut request = format!(
"{method} {path_and_query} HTTP/1.1\r\n\
Host: localhost\r\n\
Accept: application/json\r\n\
User-Agent: cflx-client/{}\r\n\
Connection: close\r\n",
env!("CARGO_PKG_VERSION")
);
if let Some(token) = &self.token {
request.push_str(&format!("Authorization: Bearer {token}\r\n"));
}
match body {
Some(body) => {
request.push_str("Content-Type: application/json\r\n");
request.push_str(&format!("Content-Length: {}\r\n\r\n", body.len()));
request.push_str(body);
}
None => request.push_str("\r\n"),
}
request
}
}
pub fn parse_response(raw: &[u8]) -> Result<HttpResponse, TransportError> {
if raw.is_empty() {
return Err(TransportError::Protocol(
"the owner closed the connection without answering".to_string(),
));
}
let separator = find_subslice(raw, b"\r\n\r\n").ok_or_else(|| {
TransportError::Protocol("the response has no complete header block".to_string())
})?;
let head = std::str::from_utf8(&raw[..separator])
.map_err(|_| TransportError::Protocol("the response head is not UTF-8".to_string()))?;
let mut lines = head.split("\r\n");
let status_line = lines
.next()
.ok_or_else(|| TransportError::Protocol("the response has no status line".to_string()))?;
let status = parse_status(status_line)?;
let mut content_length: Option<usize> = None;
let mut chunked = false;
for line in lines {
let Some((name, value)) = line.split_once(':') else {
continue;
};
let name = name.trim().to_ascii_lowercase();
let value = value.trim();
match name.as_str() {
"content-length" => {
content_length = Some(value.parse().map_err(|_| {
TransportError::Protocol(format!("invalid Content-Length '{value}'"))
})?);
}
"transfer-encoding" if value.eq_ignore_ascii_case("chunked") => chunked = true,
_ => {}
}
}
let rest = &raw[separator + 4..];
let body = if chunked {
decode_chunked(rest)?
} else if let Some(length) = content_length {
if rest.len() < length {
return Err(TransportError::Protocol(
"the response body is shorter than its Content-Length".to_string(),
));
}
rest[..length].to_vec()
} else {
rest.to_vec()
};
if body.len() > MAX_RESPONSE_BYTES {
return Err(TransportError::TooLarge);
}
Ok(HttpResponse { status, body })
}
fn parse_status(status_line: &str) -> Result<u16, TransportError> {
let mut parts = status_line.split(' ');
let version = parts.next().unwrap_or_default();
if !version.starts_with("HTTP/1.") {
return Err(TransportError::Protocol(format!(
"unexpected status line '{status_line}'"
)));
}
parts
.next()
.and_then(|code| code.parse().ok())
.ok_or_else(|| TransportError::Protocol(format!("unexpected status line '{status_line}'")))
}
fn decode_chunked(mut rest: &[u8]) -> Result<Vec<u8>, TransportError> {
let mut body = Vec::new();
loop {
let line_end = find_subslice(rest, b"\r\n").ok_or_else(|| {
TransportError::Protocol("a chunk header is not terminated".to_string())
})?;
let header = std::str::from_utf8(&rest[..line_end])
.map_err(|_| TransportError::Protocol("a chunk header is not UTF-8".to_string()))?;
let size_text = header.split(';').next().unwrap_or_default().trim();
let size = usize::from_str_radix(size_text, 16)
.map_err(|_| TransportError::Protocol(format!("invalid chunk size '{size_text}'")))?;
rest = &rest[line_end + 2..];
if size == 0 {
return Ok(body);
}
if body.len() + size > MAX_RESPONSE_BYTES {
return Err(TransportError::TooLarge);
}
if rest.len() < size + 2 {
return Err(TransportError::Protocol(
"a chunk is shorter than its declared size".to_string(),
));
}
body.extend_from_slice(&rest[..size]);
rest = &rest[size + 2..];
}
}
fn find_subslice(haystack: &[u8], needle: &[u8]) -> Option<usize> {
haystack
.windows(needle.len())
.position(|window| window == needle)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Wake {
Activity,
Gap,
Idle,
}
const STATE_MARKER: &str = "\"category\":\"state\"";
const GAP_MARKER: &str = "\"category\":\"gap\"";
impl UnixApiClient {
pub async fn wake_on_activity(
&self,
after_sequence: u64,
instance_id: &str,
budget: Duration,
) -> Wake {
match tokio::time::timeout(
budget,
self.read_stream_until_activity(after_sequence, instance_id),
)
.await
{
Ok(Ok(wake)) => wake,
Ok(Err(_)) | Err(_) => Wake::Idle,
}
}
async fn read_stream_until_activity(
&self,
after_sequence: u64,
instance_id: &str,
) -> Result<Wake, TransportError> {
let path = format!(
"/api/v2/events?after_sequence={after_sequence}&instance_id={}",
encode_query_value(instance_id)
);
let mut stream = UnixStream::connect(&self.socket).await.map_err(|error| {
TransportError::NotListening {
path: self.socket.clone(),
detail: error.to_string(),
}
})?;
let request = self.encode_request("GET", &path, None);
stream
.write_all(request.as_bytes())
.await
.map_err(|error| TransportError::Io(error.to_string()))?;
let mut seen = String::new();
let mut chunk = [0u8; 8 * 1024];
loop {
let read = stream
.read(&mut chunk)
.await
.map_err(|error| TransportError::Io(error.to_string()))?;
if read == 0 {
return Ok(Wake::Idle);
}
seen.push_str(&String::from_utf8_lossy(&chunk[..read]));
if seen.contains(GAP_MARKER) {
return Ok(Wake::Gap);
}
if seen.contains(STATE_MARKER) {
return Ok(Wake::Activity);
}
if seen.len() > 64 * 1024 {
let tail = seen.split_off(seen.len() - 128);
seen = tail;
}
}
}
}
pub fn encode_query_value(value: &str) -> String {
let mut encoded = String::with_capacity(value.len());
for byte in value.bytes() {
match byte {
b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' | b'~' => {
encoded.push(byte as char)
}
other => encoded.push_str(&format!("%{other:02X}")),
}
}
encoded
}
#[cfg(test)]
mod tests {
use super::*;
fn client() -> UnixApiClient {
UnixApiClient::new(PathBuf::from("/tmp/cflx.sock"), Some("s3cret".to_string()))
.expect("a plain token is a valid header value")
}
#[test]
fn a_token_travels_only_in_the_authorization_header() {
let request = client().encode_request("GET", "/api/v2/state", None);
assert!(request.contains("Authorization: Bearer s3cret\r\n"));
let (head, _) = request.split_once("\r\n").unwrap();
assert_eq!(head, "GET /api/v2/state HTTP/1.1");
let without_auth: String = request
.lines()
.filter(|line| !line.starts_with("Authorization:"))
.collect();
assert!(!without_auth.contains("s3cret"), "{without_auth}");
}
#[test]
fn debugging_a_client_never_renders_its_token() {
let rendered = format!("{:?}", client());
assert!(!rendered.contains("s3cret"), "{rendered}");
assert!(rendered.contains("<redacted>"), "{rendered}");
let anonymous = format!(
"{:?}",
UnixApiClient::new(PathBuf::from("/tmp/cflx.sock"), None).expect("no token")
);
assert!(anonymous.contains("<none>"), "{anonymous}");
}
#[test]
fn a_tokenless_client_sends_no_authorization_header() {
let anonymous =
UnixApiClient::new(PathBuf::from("/tmp/cflx.sock"), None).expect("no token");
let request = anonymous.encode_request("GET", "/api/v2/health", None);
assert!(!request.contains("Authorization"));
assert!(!anonymous.has_token());
}
#[test]
fn a_body_carries_its_own_length_and_content_type() {
let request = client().encode_request("POST", "/api/v2/commands", Some(r#"{"a":1}"#));
assert!(request.contains("Content-Type: application/json\r\n"));
assert!(request.contains("Content-Length: 7\r\n"));
assert!(request.ends_with("\r\n\r\n{\"a\":1}"));
}
#[test]
fn a_content_length_response_is_parsed_exactly() {
let raw = b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 7\r\n\r\n{\"a\":1}trailing";
let response = parse_response(raw).expect("parses");
assert_eq!(response.status, 200);
assert_eq!(response.body, b"{\"a\":1}");
}
#[test]
fn a_chunked_response_is_reassembled() {
let raw = b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n4\r\n{\"a\"\r\n3\r\n:1}\r\n0\r\n\r\n";
let response = parse_response(raw).expect("parses");
assert_eq!(response.body, b"{\"a\":1}");
}
#[test]
fn a_bodyless_close_delimited_response_is_accepted() {
let raw = b"HTTP/1.1 204 No Content\r\nCache-Control: no-store\r\n\r\n";
let response = parse_response(raw).expect("parses");
assert_eq!(response.status, 204);
assert!(response.body.is_empty());
}
#[test]
fn a_truncated_or_nonsense_response_is_a_protocol_error_not_a_guess() {
for raw in [
&b""[..],
&b"HTTP/1.1 200 OK\r\nContent-Length: 9\r\n\r\nshort"[..],
&b"GARBAGE\r\n\r\n"[..],
&b"HTTP/1.1 200 OK\r\nContent-Length: nine\r\n\r\nshort"[..],
] {
assert!(
matches!(parse_response(raw), Err(TransportError::Protocol(_))),
"raw={:?}",
String::from_utf8_lossy(raw)
);
}
}
#[test]
fn an_oversized_body_is_refused_before_it_is_interpreted() {
let body = vec![b'x'; MAX_RESPONSE_BYTES + 1];
let mut raw =
format!("HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n", body.len()).into_bytes();
raw.extend_from_slice(&body);
assert!(matches!(
parse_response(&raw),
Err(TransportError::TooLarge)
));
}
#[test]
fn query_values_are_encoded_rather_than_interpolated() {
assert_eq!(encode_query_value("add-client-cli"), "add-client-cli");
assert_eq!(encode_query_value("a b&c=d"), "a%20b%26c%3Dd");
assert_eq!(encode_query_value("../etc"), "..%2Fetc");
}
#[tokio::test]
async fn an_absent_socket_reports_no_owner_rather_than_an_io_error() {
let tmp = tempfile::tempdir().unwrap();
let client = UnixApiClient::new(tmp.path().join("missing.sock"), None).expect("no token");
let error = client.get("/api/v2/health").await.expect_err("no owner");
assert!(matches!(error, TransportError::NotListening { .. }));
assert!(error.to_string().contains("no Conflux owner is listening"));
}
}