use super::{
CutPoint, MAX_STEPS, Outbox, TIMESTAMP, frame, self_attestation, unframe, would_block,
};
use crate::handshake;
use crate::protocol::{ArkToHost, HostToArk, ark_to_host};
use crate::session::Session;
use crate::{
Attestation, CRYPTO_DOMAIN_WIRE, CRYPTO_DOMAIN_WIRE_ARK_TO_HOST,
CRYPTO_DOMAIN_WIRE_HOST_TO_ARK, Error, MAX_FRAME_SIZE, MAX_MESSAGE_SIZE,
};
use darkbio_crypto::{cbor, cose, xdsa, xhpke};
use std::cell::RefCell;
use std::collections::VecDeque;
use std::fmt;
use std::io::{self, Read};
use std::rc::Rc;
const PROBE_ID: u64 = u64::MAX;
#[derive(Clone, Debug, PartialEq, Eq)]
#[cfg_attr(feature = "fuzz", derive(arbitrary::Arbitrary))]
pub enum Step {
Reset,
ResetPair,
Hello,
HelloReplay,
HelloBadKey,
Ack,
AckReplay,
AckTampered,
AckBadAuth,
AckBadSigner,
AckBadPayload,
AckBadEncap,
Request(u8),
RequestReplay,
RequestTampered,
Garbage,
Junk(Vec<u8>),
Truncated(u8),
Partial,
Oversized,
Yield,
Interrupt,
Break,
Heal,
Cut { point: CutPoint, then_broken: bool },
Chunk(u8),
Batch(u8),
}
impl Step {
fn queues_frames(&self) -> bool {
!matches!(
self,
Step::Yield
| Step::Interrupt
| Step::Break
| Step::Heal
| Step::Cut { .. }
| Step::Chunk(_)
| Step::Batch(_)
)
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum State {
#[default]
Idle,
AwaitHello,
AwaitAck,
Established,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct Summary {
pub state: State, pub dropped: usize, pub fragments: usize, pub handshakes: usize, pub delivered: usize, pub replies: usize, pub reads: usize, }
enum Frame {
Empty,
Hello(Box<Keys>),
Ack,
Request(u64),
Garbage,
Junk,
}
enum Partial {
None,
Hello(Box<Keys>),
Junk,
}
enum Emit {
Dropped,
Fragment,
ArkHello(Box<Keys>),
Reply(u64, Rc<RefCell<Session>>),
}
impl fmt::Debug for Emit {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Emit::Dropped => write!(f, "Dropped"),
Emit::Fragment => write!(f, "Fragment"),
Emit::ArkHello(_) => write!(f, "ArkHello"),
Emit::Reply(id, _) => write!(f, "Reply({id})"),
}
}
}
enum Tail {
Fragment,
Body(Emit),
}
enum Payload {
Frame(Emit),
Signal,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Outcome {
Absorbed,
Message(u64),
Undecodable,
Yield,
Terminated,
}
#[derive(Clone)]
struct Keys {
signer: xdsa::SecretKey,
crypto: xhpke::SecretKey,
}
impl Keys {
fn generate() -> Self {
Self {
signer: xdsa::SecretKey::generate(),
crypto: xhpke::SecretKey::generate(),
}
}
fn hello(&self) -> Vec<u8> {
cbor::encode(&handshake::HostHello {
host_signer: self.signer.public_key(),
host_crypto: self.crypto.public_key(),
})
.unwrap()
}
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum AckFlaw {
Tampered,
Auth,
Signer,
Payload,
Encap,
}
struct Pending {
keys: Keys,
ark_crypto: xhpke::PublicKey,
receiver: xhpke::Receiver,
}
impl Pending {
fn ack(self, identity: &xdsa::PublicKey) -> (Vec<u8>, Session) {
let (sender, encap) = self
.ark_crypto
.new_sender(CRYPTO_DOMAIN_WIRE_HOST_TO_ARK)
.unwrap();
let ack = cose::seal_at(
&handshake::HostAck {
h2a_encap: encap.to_vec(),
},
&handshake::HostAckAuth {
ark_signer: identity.clone(),
ark_crypto: self.ark_crypto.clone(),
},
&self.keys.signer,
&self.ark_crypto,
CRYPTO_DOMAIN_WIRE,
TIMESTAMP,
)
.unwrap();
let session = Session {
sender,
receiver: self.receiver,
};
(ack, session)
}
fn bad_ack(&self, identity: &xdsa::PublicKey, flaw: AckFlaw) -> Vec<u8> {
let (_, encap) = self
.ark_crypto
.new_sender(CRYPTO_DOMAIN_WIRE_HOST_TO_ARK)
.unwrap();
let auth = handshake::HostAckAuth {
ark_signer: identity.clone(),
ark_crypto: match flaw {
AckFlaw::Auth => xhpke::SecretKey::generate().public_key(),
_ => self.ark_crypto.clone(),
},
};
let stranger = xdsa::SecretKey::generate();
let signer = match flaw {
AckFlaw::Signer => &stranger,
_ => &self.keys.signer,
};
let mut sealed = match flaw {
AckFlaw::Payload => cose::seal_at(
&(vec![1u8], vec![2u8]),
&auth,
signer,
&self.ark_crypto,
CRYPTO_DOMAIN_WIRE,
TIMESTAMP,
),
_ => cose::seal_at(
&handshake::HostAck {
h2a_encap: match flaw {
AckFlaw::Encap => vec![0x42; 3],
_ => encap.to_vec(),
},
},
&auth,
signer,
&self.ark_crypto,
CRYPTO_DOMAIN_WIRE,
TIMESTAMP,
),
}
.unwrap();
if flaw == AckFlaw::Tampered {
*sealed.last_mut().unwrap() ^= 0xff;
}
sealed
}
}
pub struct Client {
steps: VecDeque<Step>,
identity: xdsa::PublicKey, outbox: Outbox, bytes: Vec<u8>, chunk: usize, batch: usize, broken: bool, cut: Option<CutPoint>, flush_fails: bool,
state: State, partial: Partial, resync: bool, tail: Option<Tail>, emits: Vec<Emit>, outcome: Outcome,
pending: Option<Pending>, session: Option<Rc<RefCell<Session>>>,
last_hello: Option<(Vec<u8>, Keys)>, last_ack: Option<Vec<u8>>, last_request: Option<Vec<u8>>, last_valid: Option<Vec<u8>>,
summary: Summary,
}
impl Client {
fn new(steps: &[Step], identity: xdsa::PublicKey, outbox: Outbox) -> Self {
Self {
steps: steps.iter().take(MAX_STEPS).cloned().collect(),
identity,
outbox,
bytes: Vec::new(),
chunk: 0,
batch: 0,
broken: false,
cut: None,
flush_fails: false,
state: State::Idle,
partial: Partial::None,
resync: false,
tail: None,
emits: Vec::new(),
outcome: Outcome::Absorbed,
pending: None,
session: None,
last_hello: None,
last_ack: None,
last_request: None,
last_valid: None,
summary: Summary::default(),
}
}
fn execute(&mut self, step: Step) {
match step {
Step::Reset => {
self.deliver(Frame::Empty);
self.bytes.push(0x00);
}
Step::ResetPair => {
self.deliver(Frame::Empty);
self.deliver(Frame::Empty);
self.bytes.extend([0x00, 0x00]);
}
Step::Hello => {
let keys = Keys::generate();
let framed = frame(&keys.hello());
self.record(&framed);
self.last_hello = Some((framed.clone(), keys.clone()));
self.deliver(Frame::Hello(Box::new(keys)));
self.bytes.extend(framed);
}
Step::HelloReplay => {
if let Some((framed, keys)) = self.last_hello.clone() {
self.deliver(Frame::Hello(Box::new(keys)));
self.bytes.extend(framed);
}
}
Step::HelloBadKey => {
let signer = xdsa::SecretKey::generate().public_key().to_bytes().to_vec();
let hello = cbor::encode(&(signer, vec![0xffu8; xhpke::PUBLIC_KEY_SIZE])).unwrap();
self.junk(&hello);
}
Step::Ack => match self.pending.take() {
Some(pending) => {
let (ack, session) = pending.ack(&self.identity);
let framed = frame(&ack);
self.record(&framed);
self.last_ack = Some(framed.clone());
self.session = Some(Rc::new(RefCell::new(session)));
self.deliver(Frame::Ack);
self.bytes.extend(framed);
}
None => self.junk(b"ack without a pending server hello"),
},
Step::AckReplay => {
if let Some(framed) = self.last_ack.clone() {
self.deliver(Frame::Junk);
self.bytes.extend(framed);
}
}
Step::AckTampered => self.bad_ack(AckFlaw::Tampered),
Step::AckBadAuth => self.bad_ack(AckFlaw::Auth),
Step::AckBadSigner => self.bad_ack(AckFlaw::Signer),
Step::AckBadPayload => self.bad_ack(AckFlaw::Payload),
Step::AckBadEncap => self.bad_ack(AckFlaw::Encap),
Step::Request(tag) => match self.session.as_ref() {
Some(session) => {
let id = tag as u64;
let request = HostToArk {
id: Some(id),
content: None,
};
let packet = session
.borrow_mut()
.seal(&request, &mut Vec::new())
.unwrap();
let framed = frame(&packet);
self.record(&framed);
self.last_request = Some(framed.clone());
self.deliver(Frame::Request(id));
self.bytes.extend(framed);
}
None => self.junk(b"request without a session"),
},
Step::RequestReplay => {
if let Some(framed) = self.last_request.clone() {
self.deliver(Frame::Junk);
self.bytes.extend(framed);
}
}
Step::RequestTampered => match self.session.as_ref() {
Some(session) => {
let request = HostToArk {
id: Some(0),
content: None,
};
let mut packet = session
.borrow_mut()
.seal(&request, &mut Vec::new())
.unwrap();
*packet.last_mut().unwrap() ^= 0xff;
self.junk(&packet);
}
None => self.junk(b"tampered request without a session"),
},
Step::Garbage => match self.session.as_ref() {
Some(session) => {
let packet = session.borrow_mut().sender.seal(&[0x07], &[]).unwrap();
let framed = frame(&packet);
self.record(&framed);
self.deliver(Frame::Garbage);
self.bytes.extend(framed);
}
None => self.junk(b"garbage without a session"),
},
Step::Junk(mut junk) => {
for byte in junk.iter_mut() {
if *byte == 0 {
*byte = 1;
}
}
if junk.is_empty() {
junk.push(1);
}
self.deliver(Frame::Junk);
self.bytes.extend(junk);
self.bytes.push(0x00);
}
Step::Truncated(n) => {
if let Some(valid) = self.last_valid.clone() {
let keep = match valid.len() {
0..=1 => 1,
len => 1 + n as usize % (len - 1),
};
self.deliver(Frame::Junk);
self.bytes.extend(&valid[..keep]);
self.bytes.push(0x00);
}
}
Step::Partial => {
let keys = Keys::generate();
let mut framed = frame(&keys.hello());
framed.pop();
self.bytes.extend(framed);
self.partial = match self.partial {
Partial::None => Partial::Hello(Box::new(keys)),
_ => Partial::Junk,
};
}
Step::Oversized => {
self.partial = Partial::None;
self.bytes.resize(self.bytes.len() + MAX_FRAME_SIZE + 1, 1);
self.bytes.push(0x00);
}
Step::Yield | Step::Interrupt => {
unreachable!("yields and interrupts are handled by the reader")
}
Step::Break => self.set_broken(true),
Step::Heal => self.set_broken(false),
Step::Cut { point, then_broken } => {
self.cut = Some(point);
self.outbox.set_cut(point);
if then_broken {
self.set_broken(true);
}
}
Step::Chunk(n) => self.chunk = n as usize,
Step::Batch(n) => self.batch = n as usize,
}
}
fn bad_ack(&mut self, flaw: AckFlaw) {
match self.pending.as_ref() {
Some(pending) => {
let ack = pending.bad_ack(&self.identity, flaw);
self.junk(&ack);
}
None => self.junk(b"bad ack without a pending server hello"),
}
}
fn junk(&mut self, text: &[u8]) {
self.deliver(Frame::Junk);
self.bytes.extend(frame(text));
}
fn record(&mut self, framed: &[u8]) {
self.last_valid = Some(framed[..framed.len() - 1].to_vec());
}
fn set_broken(&mut self, broken: bool) {
self.outbox.set_broken(broken);
self.broken = broken;
}
fn deliver(&mut self, frame: Frame) {
let frame = match std::mem::replace(&mut self.partial, Partial::None) {
Partial::None => frame,
Partial::Hello(keys) => match frame {
Frame::Empty => Frame::Hello(keys),
_ => Frame::Junk,
},
Partial::Junk => Frame::Junk,
};
match (self.state, frame) {
(_, Frame::Empty) => {
self.forget();
self.state = State::AwaitHello;
}
(State::AwaitHello, Frame::Hello(keys)) => {
if self.send(Payload::Frame(Emit::ArkHello(keys))) {
self.state = State::AwaitAck;
} else {
self.forget();
self.state = State::Idle;
self.send(Payload::Signal);
}
}
(State::AwaitAck, Frame::Ack) => {
self.state = State::Established;
}
(State::Established, Frame::Request(id)) => {
self.outcome = Outcome::Message(id);
}
(State::Established, Frame::Garbage) => {
self.outcome = Outcome::Undecodable;
}
_ => {
self.forget();
self.state = State::Idle;
self.send(Payload::Signal);
}
}
}
fn send(&mut self, payload: Payload) -> bool {
let mut sent = !self.resync || self.write_zero();
if sent {
sent = match payload {
Payload::Frame(emit) => self.write_frame(emit),
Payload::Signal => self.write_zero(),
};
}
if sent {
sent = !std::mem::take(&mut self.flush_fails);
}
self.resync = !sent;
sent
}
fn write_zero(&mut self) -> bool {
if self.cut == Some(CutPoint::Start) {
self.cut = None;
return false;
}
if self.broken {
return false;
}
self.zero_out();
true
}
fn write_frame(&mut self, emit: Emit) -> bool {
match self.cut.take() {
Some(CutPoint::Start) => false,
Some(CutPoint::Middle(_)) => {
self.tail = Some(Tail::Fragment);
false
}
Some(CutPoint::Delimiter) => {
self.tail = Some(Tail::Body(emit));
false
}
Some(CutPoint::Flush) => {
self.flush_fails = true;
self.emits.push(emit);
true
}
None if self.broken => false,
None => {
self.emits.push(emit);
true
}
}
}
fn zero_out(&mut self) {
self.emits.push(match self.tail.take() {
None => Emit::Dropped,
Some(Tail::Fragment) => Emit::Fragment,
Some(Tail::Body(emit)) => emit,
});
}
fn forget(&mut self) {
self.session = None;
self.pending = None;
}
fn interrupt(&mut self, outcome: Outcome) {
if matches!(self.state, State::AwaitHello | State::AwaitAck) {
self.forget();
self.state = State::Idle;
}
self.outcome = outcome;
}
fn surfaced(&mut self, outcome: Outcome) {
assert_eq!(self.outcome, outcome, "model vs server");
self.outcome = Outcome::Absorbed;
}
fn sync(&mut self) {
assert_eq!(
self.outcome,
Outcome::Absorbed,
"server read on past a step it should have surfaced"
);
let frames = self.outbox.take_frames();
let emits = std::mem::take(&mut self.emits);
assert_eq!(frames.len(), emits.len(), "model expected {emits:?}");
assert_eq!(
self.outbox.has_tail(),
self.tail.is_some(),
"unterminated frame"
);
for (frame, emit) in frames.iter().zip(emits) {
match emit {
Emit::Dropped => {
assert!(
frame.is_empty(),
"expected an empty frame, server emitted {} bytes",
frame.len()
);
self.summary.dropped += 1;
}
Emit::Fragment => {
assert!(
!frame.is_empty(),
"expected a cut frame, server emitted an empty one"
);
self.summary.fragments += 1;
}
Emit::ArkHello(keys) => {
self.receive_hello(frame, *keys);
self.summary.handshakes += 1;
}
Emit::Reply(id, session) => {
let msg: ArkToHost = session
.borrow_mut()
.open(&unframe(frame))
.expect("reply failed to open");
assert_eq!(msg.id, Some(id));
self.summary.replies += 1;
}
}
}
}
fn receive_hello(&mut self, frame: &[u8], keys: Keys) {
let auth = handshake::ArkHelloAuth {
host_signer: keys.signer.public_key(),
host_crypto: keys.crypto.public_key(),
};
let sign1 = cose::decrypt(&unframe(frame), &auth, &keys.crypto, CRYPTO_DOMAIN_WIRE)
.expect("server hello failed to decrypt");
let hello: handshake::ArkHello =
cose::verify(&sign1, &auth, &self.identity, CRYPTO_DOMAIN_WIRE, None)
.expect("server hello signature invalid");
let encap: [u8; xhpke::ENCAP_KEY_SIZE] = hello
.a2h_encap
.try_into()
.expect("server hello encap size invalid");
let receiver = keys
.crypto
.new_receiver(&encap, CRYPTO_DOMAIN_WIRE_ARK_TO_HOST)
.unwrap();
if self.state == State::AwaitAck {
self.pending = Some(Pending {
keys,
ark_crypto: hello.ark_crypto,
receiver,
});
}
}
}
struct Feed(Rc<RefCell<Client>>);
impl Read for Feed {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
let mut client = self.0.borrow_mut();
if client.bytes.is_empty() {
client.sync();
client.batch = 0;
loop {
match client.steps.pop_front() {
None => {
client.interrupt(Outcome::Terminated);
return Ok(0);
}
Some(Step::Yield) => {
client.interrupt(Outcome::Yield);
return Err(would_block());
}
Some(Step::Interrupt) => return Err(io::ErrorKind::Interrupted.into()),
Some(step) => client.execute(step),
}
if !client.bytes.is_empty() {
break;
}
}
while client.batch > 1
&& client.outcome == Outcome::Absorbed
&& client.steps.front().is_some_and(Step::queues_frames)
{
let step = client.steps.pop_front().unwrap();
client.execute(step);
client.batch -= 1;
}
}
let mut n = buf.len().min(client.bytes.len());
if client.chunk > 0 {
n = n.min(client.chunk);
}
buf[..n].copy_from_slice(&client.bytes[..n]);
client.bytes.drain(..n);
client.summary.reads += 1;
Ok(n)
}
}
type Server = crate::Server<Feed, Outbox, Attestation>;
fn check_session(server: &mut Server, client: &Client) {
let established = client.state == State::Established;
let oversized = vec![0x42; MAX_MESSAGE_SIZE + 1];
let refused = server.send_message(ArkToHost {
id: Some(PROBE_ID),
err: None,
content: Some(ark_to_host::Content::Develop(oversized)),
});
match refused {
Err(Error::PacketTooLarge(_)) => {
assert!(established, "server has a session the model does not")
}
Err(Error::EncryptionFailed(_)) => {
assert!(!established, "server lacks the session the model has")
}
other => panic!("unexpected oversized send result: {other:?}"),
}
}
fn send(server: &mut Server, client: &mut Client, id: u64) {
let established = client.state == State::Established;
let expected = established.then(|| {
let session = client
.session
.clone()
.expect("established without a session");
let sent = client.send(Payload::Frame(Emit::Reply(id, session)));
if !sent {
client.forget();
client.state = State::Idle;
client.send(Payload::Signal);
}
sent
});
let sent = server.send_message(ArkToHost {
id: Some(id),
err: None,
content: None,
});
match (expected, sent) {
(Some(true), Ok(())) => {}
(Some(false), Err(Error::SendFailed(_))) => {}
(None, Err(Error::EncryptionFailed(_))) => {}
(expected, sent) => panic!("model expected {expected:?}, server returned {sent:?}"),
}
}
pub fn run(steps: &[Step]) -> Summary {
#[cfg(feature = "fuzz")]
super::seed::seed(super::seed::SERVER_PROTOCOL, steps);
let signer = xdsa::SecretKey::generate();
let attestation = self_attestation(&signer);
let outbox = Outbox::default();
let client = Rc::new(RefCell::new(Client::new(
steps,
signer.public_key(),
outbox.clone(),
)));
let mut server = Server::new_at(Feed(client.clone()), outbox, signer, attestation, TIMESTAMP);
loop {
match server.next_message() {
Ok(msg) => {
let id = msg.id.expect("delivered message without an id");
let mut client = client.borrow_mut();
client.surfaced(Outcome::Message(id));
client.summary.delivered += 1;
send(&mut server, &mut client, id);
}
Err(Error::PacketDecodingFailed(_)) => {
client.borrow_mut().surfaced(Outcome::Undecodable);
}
Err(Error::RecvFailed(err)) if err.kind() == io::ErrorKind::WouldBlock => {
let mut client = client.borrow_mut();
client.surfaced(Outcome::Yield);
check_session(&mut server, &client);
send(&mut server, &mut client, PROBE_ID);
}
Err(Error::Terminated) => {
client.borrow_mut().surfaced(Outcome::Terminated);
break;
}
Err(err) => panic!("unexpected error from the server: {err}"),
}
}
let mut client = client.borrow_mut();
client.sync();
check_session(&mut server, &client);
client.summary.state = client.state;
client.summary
}
#[cfg(test)]
mod tests;