#[cfg(test)]
pub mod replay;
use base64::prelude::*;
use std::cell::Cell;
use std::fmt::Debug;
use std::path::PathBuf;
pub const ENV: &str = "WIRE_VECTORS";
#[derive(PartialEq, Eq)]
pub struct Vector {
scenario: String, script: Vec<String>, write_failures: bool, identity: Vec<u8>, attestation: Vec<u8>, server_keys: Vec<Vec<u8>>, trace: Vec<Event>, }
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Event {
Handshake { xdsa: Vec<u8>, xhpke: Vec<u8> },
Send { message: Vec<u8> },
Retain,
SendRetained { message: Vec<u8> },
Recv,
Ok { message: Option<Vec<u8>> },
Err { kind: String },
Session { established: bool },
Read { bytes: Vec<u8>, chunk: usize },
ReadFailed { error: ReadError },
Write { bytes: Vec<u8>, failed: bool },
FlushFailed,
WriteTimedOut { bytes: Vec<u8> },
FlushTimedOut,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ReadError {
Eof,
Failed,
Interrupted,
TimedOut,
}
impl ReadError {
pub fn name(self) -> &'static str {
match self {
ReadError::Eof => "eof",
ReadError::Failed => "failed",
ReadError::Interrupted => "interrupted",
ReadError::TimedOut => "timed_out",
}
}
pub fn parse(name: &str) -> Self {
match name {
"eof" => ReadError::Eof,
"failed" => ReadError::Failed,
"interrupted" => ReadError::Interrupted,
"timed_out" => ReadError::TimedOut,
other => panic!("unknown read error {other}"),
}
}
}
thread_local! {
static RUNS: Cell<usize> = const { Cell::new(0) };
}
pub fn scenario() -> Option<String> {
if !cfg!(test) && std::env::var_os(ENV).is_none() {
return None;
}
let thread = std::thread::current();
let base = thread
.name()
.unwrap_or("scenario")
.rsplit("::")
.next()
.unwrap();
let base = base
.strip_prefix("test_scripted_")
.or_else(|| base.strip_prefix("test_"))
.unwrap_or(base);
let runs = RUNS.with(|runs| {
runs.set(runs.get() + 1);
runs.get()
});
Some(match runs {
1 => base.to_string(),
n => format!("{base}-{n}"),
})
}
impl Vector {
pub fn open<S: Debug>(
scenario: Option<String>,
steps: &[S],
write_failures: bool,
identity: Vec<u8>,
attestation: Vec<u8>,
) -> Option<Self> {
Some(Self {
scenario: scenario?,
script: steps.iter().map(|step| format!("{step:?}")).collect(),
write_failures,
identity,
attestation,
server_keys: Vec::new(),
trace: Vec::new(),
})
}
pub fn server_key(&mut self, xhpke: Vec<u8>) {
self.server_keys.push(xhpke);
}
pub fn log(&mut self, event: Event) {
self.trace.push(event);
}
pub fn write(&self) {
let Some(root) = std::env::var_os(ENV) else {
return;
};
let dir = PathBuf::from(root).join("client");
std::fs::create_dir_all(&dir).expect("failed to create the vector directory");
std::fs::write(dir.join(format!("{}.json", self.scenario)), self.json())
.expect("failed to write the vector");
}
pub fn json(&self) -> String {
let script: Vec<String> = self
.script
.iter()
.map(|step| serde_json::to_string(step).unwrap())
.collect();
let mut lines = vec![
"{".to_string(),
format!(
" \"scenario\": {},",
serde_json::to_string(&self.scenario).unwrap()
),
format!(" \"script\": [{}],", script.join(", ")),
format!(" \"write_failures\": {},", self.write_failures),
" \"server\": {".to_string(),
format!(
" \"identity\": \"{}\",",
BASE64_STANDARD.encode(&self.identity)
),
format!(
" \"attestation\": \"{}\",",
BASE64_STANDARD.encode(&self.attestation)
),
" \"xhpke\": [".to_string(),
];
let keys = self
.server_keys
.iter()
.map(|key| format!(" \"{}\"", BASE64_STANDARD.encode(key)));
lines.extend(listed(keys));
lines.extend([" ]", " },", " \"trace\": ["].map(String::from));
lines.extend(listed(
self.trace
.iter()
.map(|event| format!(" {}", event.json())),
));
lines.extend([" ]", "}", ""].map(String::from));
lines.join("\n")
}
}
impl Event {
fn json(&self) -> String {
let fields = match self {
Event::Handshake { xdsa, xhpke } => format!(
"\"event\": \"handshake\", \"xdsa\": \"{}\", \"xhpke\": \"{}\"",
BASE64_STANDARD.encode(xdsa),
BASE64_STANDARD.encode(xhpke)
),
Event::Send { message } if message.len() > crate::transport::MAX_MESSAGE_SIZE => {
format!("\"event\": \"send\", {}", payload(message))
}
Event::Send { message } => format!(
"\"event\": \"send\", \"message\": \"{}\"",
BASE64_STANDARD.encode(message)
),
Event::Retain => "\"event\": \"retain\"".to_string(),
Event::SendRetained { message } => format!(
"\"event\": \"send_retained\", \"message\": \"{}\"",
BASE64_STANDARD.encode(message)
),
Event::Recv => "\"event\": \"recv\"".to_string(),
Event::Ok { message: None } => "\"event\": \"ok\"".to_string(),
Event::Ok {
message: Some(message),
} => format!(
"\"event\": \"ok\", \"message\": \"{}\"",
BASE64_STANDARD.encode(message)
),
Event::Err { kind } => format!("\"event\": \"error\", \"kind\": \"{kind}\""),
Event::Session { established } => {
format!("\"event\": \"session\", \"established\": {established}")
}
Event::Read { bytes, chunk: 0 } => format!("\"event\": \"read\", {}", payload(bytes)),
Event::Read { bytes, chunk } => {
format!(
"\"event\": \"read\", {}, \"chunk\": {chunk}",
payload(bytes)
)
}
Event::ReadFailed { error } => {
format!("\"event\": \"read\", \"error\": \"{}\"", error.name())
}
Event::Write {
bytes,
failed: false,
} => format!("\"event\": \"write\", {}", payload(bytes)),
Event::Write {
bytes,
failed: true,
} => format!("\"event\": \"write\", {}, \"failed\": true", payload(bytes)),
Event::FlushFailed => "\"event\": \"flush_failed\"".to_string(),
Event::WriteTimedOut { bytes } => {
format!("\"event\": \"write_timed_out\", {}", payload(bytes))
}
Event::FlushTimedOut => "\"event\": \"flush_timed_out\"".to_string(),
};
format!("{{{fields}}}")
}
}
fn payload(bytes: &[u8]) -> String {
let mut runs: Vec<(u8, usize)> = Vec::new();
for &byte in bytes {
match runs.last_mut() {
Some((last, count)) if *last == byte => *count += 1,
_ => runs.push((byte, 1)),
}
}
if runs.len() * 8 < bytes.len() {
let runs: Vec<String> = runs
.iter()
.map(|(byte, count)| format!("[{byte}, {count}]"))
.collect();
return format!("\"runs\": [{}]", runs.join(", "));
}
format!("\"bytes\": \"{}\"", BASE64_STANDARD.encode(bytes))
}
fn listed(items: impl Iterator<Item = String>) -> impl Iterator<Item = String> {
let mut items = items.peekable();
std::iter::from_fn(move || {
let mut item = items.next()?;
if items.peek().is_some() {
item.push(',');
}
Some(item)
})
}