use std::collections::HashMap;
use chia_bls::SecretKey;
use chia_protocol::Bytes32;
use crate::constants::MAX_CONCURRENT_STREAMS;
use crate::envelope::{DigMessageEnvelope, InteractionShape, StreamFrame, StreamHeader};
use crate::error::{MessageError, Result};
use crate::replay::ReplayGuard;
use crate::seal::{open_message, seal_message, SealParams};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Role {
Initiator,
Responder,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Handshake {
Opening,
Established,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StreamState {
Opening,
Open,
HalfClosedLocal,
HalfClosedRemote,
Closed,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Accepted {
Established,
Data,
Credit(u64),
RemoteClosed,
CloseAcked,
Reset,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum StreamEvent {
Opened,
Established,
Data(Vec<u8>),
CreditGranted(u64),
RemoteClosed,
CloseAcked,
PeerReset,
}
#[derive(Debug, PartialEq, Eq)]
pub enum StreamAccept {
Event(StreamEvent),
Dropped {
cause: MessageError,
},
Reset {
frame: Box<DigMessageEnvelope>,
cause: MessageError,
},
}
#[derive(Debug)]
pub struct StreamSession {
role: Role,
handshake: Handshake,
send_seq: u64,
recv_seq: u64,
send_credit: u64,
recv_window_remaining: u64,
local_closed: bool,
remote_closed: bool,
reset: bool,
}
impl StreamSession {
fn initiator(recv_window: u32) -> Self {
Self {
role: Role::Initiator,
handshake: Handshake::Opening,
send_seq: 0,
recv_seq: 0,
send_credit: 0,
recv_window_remaining: u64::from(recv_window),
local_closed: false,
remote_closed: false,
reset: false,
}
}
fn responder(granted_credit: u32) -> Self {
Self {
role: Role::Responder,
handshake: Handshake::Opening,
send_seq: 0,
recv_seq: 0,
send_credit: u64::from(granted_credit),
recv_window_remaining: 0,
local_closed: false,
remote_closed: false,
reset: false,
}
}
#[must_use]
pub fn state(&self) -> StreamState {
if self.reset || (self.local_closed && self.remote_closed) {
return StreamState::Closed;
}
if self.handshake == Handshake::Opening {
return StreamState::Opening;
}
match (self.local_closed, self.remote_closed) {
(true, false) => StreamState::HalfClosedLocal,
(false, true) => StreamState::HalfClosedRemote,
_ => StreamState::Open,
}
}
#[must_use]
pub fn is_closed(&self) -> bool {
self.state() == StreamState::Closed
}
#[must_use]
pub fn send_credit(&self) -> u64 {
self.send_credit
}
fn build_open_ack(&mut self, recv_window: u32) -> Result<StreamHeader> {
if self.role != Role::Responder || self.handshake != Handshake::Opening {
return Err(MessageError::StreamProtocol(
"OPEN_ACK only from a responder mid-handshake",
));
}
self.handshake = Handshake::Established;
self.recv_window_remaining = u64::from(recv_window);
Ok(header(StreamFrame::OpenAck, 0, recv_window))
}
fn build_data(&mut self) -> Result<StreamHeader> {
if self.handshake != Handshake::Established {
return Err(MessageError::StreamProtocol(
"DATA before the stream is established",
));
}
if self.local_closed {
return Err(MessageError::StreamProtocol("DATA after our own CLOSE"));
}
if self.send_credit == 0 {
return Err(MessageError::StreamProtocol(
"DATA exceeds the granted credit window",
));
}
self.send_credit -= 1;
let seq = self.send_seq;
self.send_seq += 1;
Ok(header(StreamFrame::Data, seq, 0))
}
fn build_credit(&mut self, n: u32) -> Result<StreamHeader> {
if self.handshake != Handshake::Established {
return Err(MessageError::StreamProtocol(
"CREDIT before the stream is established",
));
}
self.recv_window_remaining = self.recv_window_remaining.saturating_add(u64::from(n));
Ok(header(StreamFrame::Credit, 0, n))
}
fn build_close(&mut self) -> Result<StreamHeader> {
if self.local_closed {
return Err(MessageError::StreamProtocol("CLOSE after our own CLOSE"));
}
self.local_closed = true;
Ok(header(StreamFrame::Close, 0, 0))
}
fn build_reset(&mut self) -> StreamHeader {
self.reset = true;
header(StreamFrame::Reset, 0, 0)
}
fn on_recv(&mut self, frame: StreamFrame, hdr: StreamHeader) -> Result<Accepted> {
if self.reset {
return Err(MessageError::StreamProtocol("frame after RESET"));
}
match frame {
StreamFrame::Open => Err(MessageError::StreamProtocol(
"duplicate OPEN for a live stream",
)),
StreamFrame::OpenAck => {
if self.role != Role::Initiator || self.handshake != Handshake::Opening {
return Err(MessageError::StreamProtocol("unexpected OPEN_ACK"));
}
self.handshake = Handshake::Established;
self.send_credit = u64::from(hdr.window);
Ok(Accepted::Established)
}
StreamFrame::Data => {
if self.handshake != Handshake::Established {
return Err(MessageError::StreamProtocol(
"DATA before the stream is established",
));
}
if self.remote_closed {
return Err(MessageError::StreamProtocol("DATA after the peer's CLOSE"));
}
if hdr.seq != self.recv_seq {
return Err(MessageError::StreamProtocol(
"out-of-order / gap / replayed DATA seq",
));
}
if self.recv_window_remaining == 0 {
return Err(MessageError::StreamProtocol(
"DATA exceeds the credit window we granted",
));
}
self.recv_window_remaining -= 1;
self.recv_seq += 1;
Ok(Accepted::Data)
}
StreamFrame::Credit => {
if self.handshake != Handshake::Established {
return Err(MessageError::StreamProtocol(
"CREDIT before the stream is established",
));
}
self.send_credit = self.send_credit.saturating_add(u64::from(hdr.window));
Ok(Accepted::Credit(u64::from(hdr.window)))
}
StreamFrame::Close => {
if self.handshake != Handshake::Established {
return Err(MessageError::StreamProtocol(
"CLOSE before the stream is established",
));
}
if self.remote_closed {
return Err(MessageError::StreamProtocol(
"duplicate CLOSE from the peer",
));
}
self.remote_closed = true;
Ok(Accepted::RemoteClosed)
}
StreamFrame::CloseAck => {
if !self.local_closed {
return Err(MessageError::StreamProtocol("CLOSE_ACK without our CLOSE"));
}
Ok(Accepted::CloseAcked)
}
StreamFrame::Reset => {
self.reset = true;
Ok(Accepted::Reset)
}
}
}
}
fn header(frame: StreamFrame, seq: u64, window: u32) -> StreamHeader {
StreamHeader {
frame: frame.as_u8(),
seq,
window,
}
}
pub struct StreamEndpoint<'a> {
identity_sk: &'a SecretKey,
local_did: Bytes32,
local_epoch: u32,
peer_did: Bytes32,
peer_pub: &'a [u8; 48],
message_type: u32,
send_counter: u64,
sessions: HashMap<Bytes32, StreamSession>,
max_concurrent: usize,
guard: ReplayGuard,
}
impl<'a> StreamEndpoint<'a> {
#[must_use]
pub fn new(
identity_sk: &'a SecretKey,
local_did: Bytes32,
local_epoch: u32,
peer_did: Bytes32,
peer_pub: &'a [u8; 48],
message_type: u32,
) -> Self {
Self {
identity_sk,
local_did,
local_epoch,
peer_did,
peer_pub,
message_type,
send_counter: 0,
sessions: HashMap::new(),
max_concurrent: MAX_CONCURRENT_STREAMS,
guard: ReplayGuard::new(),
}
}
#[must_use]
pub fn with_max_concurrent(mut self, max: usize) -> Self {
self.max_concurrent = max;
self
}
#[must_use]
pub fn stream_count(&self) -> usize {
self.sessions.len()
}
#[must_use]
pub fn session(&self, correlation_id: Bytes32) -> Option<&StreamSession> {
self.sessions.get(&correlation_id)
}
pub fn open(
&mut self,
correlation_id: Bytes32,
recv_window: u32,
now_ms: u64,
expires_at: u64,
) -> Result<DigMessageEnvelope> {
if self.sessions.contains_key(&correlation_id) {
return Err(MessageError::StreamProtocol(
"correlation_id already in use",
));
}
if self.sessions.len() >= self.max_concurrent {
return Err(MessageError::StreamLimit {
cap: self.max_concurrent,
});
}
self.sessions
.insert(correlation_id, StreamSession::initiator(recv_window));
let hdr = header(StreamFrame::Open, 0, recv_window);
self.seal_frame(correlation_id, hdr, &[], now_ms, expires_at)
}
pub fn open_ack(
&mut self,
correlation_id: Bytes32,
recv_window: u32,
now_ms: u64,
expires_at: u64,
) -> Result<DigMessageEnvelope> {
let hdr = self
.session_mut(correlation_id)?
.build_open_ack(recv_window)?;
self.seal_frame(correlation_id, hdr, &[], now_ms, expires_at)
}
pub fn send_data(
&mut self,
correlation_id: Bytes32,
payload: &[u8],
now_ms: u64,
expires_at: u64,
) -> Result<DigMessageEnvelope> {
let hdr = self.session_mut(correlation_id)?.build_data()?;
self.seal_frame(correlation_id, hdr, payload, now_ms, expires_at)
}
pub fn grant_credit(
&mut self,
correlation_id: Bytes32,
n: u32,
now_ms: u64,
expires_at: u64,
) -> Result<DigMessageEnvelope> {
let hdr = self.session_mut(correlation_id)?.build_credit(n)?;
self.seal_frame(correlation_id, hdr, &[], now_ms, expires_at)
}
pub fn close(
&mut self,
correlation_id: Bytes32,
now_ms: u64,
expires_at: u64,
) -> Result<DigMessageEnvelope> {
let hdr = self.session_mut(correlation_id)?.build_close()?;
let env = self.seal_frame(correlation_id, hdr, &[], now_ms, expires_at)?;
self.drop_if_closed(correlation_id);
Ok(env)
}
pub fn reset(
&mut self,
correlation_id: Bytes32,
now_ms: u64,
expires_at: u64,
) -> Result<DigMessageEnvelope> {
let hdr = self.session_mut(correlation_id)?.build_reset();
let env = self.seal_frame(correlation_id, hdr, &[], now_ms, expires_at)?;
self.sessions.remove(&correlation_id);
Ok(env)
}
pub fn accept(
&mut self,
envelope: &DigMessageEnvelope,
resolve_sender_pub: impl Fn(Bytes32, u32) -> Option<[u8; 48]>,
now_ms: u64,
) -> Result<StreamAccept> {
let correlation_id = envelope.correlation_id;
let opened = match open_message(
self.identity_sk,
envelope,
&resolve_sender_pub,
&mut self.guard,
now_ms,
) {
Ok(opened) => opened,
Err(cause) => return Ok(StreamAccept::Dropped { cause }),
};
if opened.sender != self.peer_did {
return Ok(StreamAccept::Dropped {
cause: MessageError::StreamProtocol(
"authenticated sender is not this stream's peer",
),
});
}
let Some(hdr) = envelope.stream else {
return Ok(StreamAccept::Dropped {
cause: MessageError::StreamProtocol("stream event on a non-stream envelope"),
});
};
let Some(frame) = StreamFrame::from_u8(hdr.frame) else {
return Ok(StreamAccept::Dropped {
cause: MessageError::StreamProtocol("unknown stream frame kind"),
});
};
if opened.shape != InteractionShape::StreamFrame {
return Ok(StreamAccept::Dropped {
cause: MessageError::StreamProtocol("stream frame with a non-stream shape"),
});
}
if frame == StreamFrame::Reset {
return Ok(if self.sessions.remove(&correlation_id).is_some() {
StreamAccept::Event(StreamEvent::PeerReset)
} else {
StreamAccept::Dropped {
cause: MessageError::StreamProtocol("RESET for an unknown stream"),
}
});
}
if frame == StreamFrame::Open {
return self.accept_open(correlation_id, hdr, now_ms);
}
if !self.sessions.contains_key(&correlation_id) {
return Ok(StreamAccept::Dropped {
cause: MessageError::StreamProtocol("frame for an unknown stream"),
});
}
let transition = self
.sessions
.get_mut(&correlation_id)
.expect("presence checked above")
.on_recv(frame, hdr);
match transition {
Ok(accepted) => {
let event = to_event(accepted, opened.payload);
self.drop_if_closed(correlation_id);
Ok(StreamAccept::Event(event))
}
Err(cause) => {
self.sessions.remove(&correlation_id);
self.reset_response(correlation_id, now_ms, cause)
}
}
}
fn accept_open(
&mut self,
correlation_id: Bytes32,
hdr: StreamHeader,
now_ms: u64,
) -> Result<StreamAccept> {
if self.sessions.remove(&correlation_id).is_some() {
return self.reset_response(
correlation_id,
now_ms,
MessageError::StreamProtocol("duplicate OPEN for a live stream"),
);
}
if self.sessions.len() >= self.max_concurrent {
return self.reset_response(
correlation_id,
now_ms,
MessageError::StreamLimit {
cap: self.max_concurrent,
},
);
}
self.sessions
.insert(correlation_id, StreamSession::responder(hdr.window));
Ok(StreamAccept::Event(StreamEvent::Opened))
}
fn reset_response(
&mut self,
correlation_id: Bytes32,
now_ms: u64,
cause: MessageError,
) -> Result<StreamAccept> {
let hdr = header(StreamFrame::Reset, 0, 0);
let frame = self.seal_frame(correlation_id, hdr, &[], now_ms, 0)?;
Ok(StreamAccept::Reset {
frame: Box::new(frame),
cause,
})
}
fn seal_frame(
&mut self,
correlation_id: Bytes32,
hdr: StreamHeader,
payload: &[u8],
now_ms: u64,
expires_at: u64,
) -> Result<DigMessageEnvelope> {
let counter = self.send_counter;
self.send_counter += 1;
let params = SealParams {
sender_sk: self.identity_sk,
sender: self.local_did,
sender_epoch: self.local_epoch,
recipient: self.peer_did,
recipient_pub: self.peer_pub,
message_type: self.message_type,
shape: InteractionShape::StreamFrame,
correlation_id,
stream: Some(hdr),
counter,
timestamp_ms: now_ms,
expires_at,
payload,
};
seal_message(¶ms)
}
fn session_mut(&mut self, correlation_id: Bytes32) -> Result<&mut StreamSession> {
self.sessions
.get_mut(&correlation_id)
.ok_or(MessageError::StreamProtocol("unknown stream"))
}
fn drop_if_closed(&mut self, correlation_id: Bytes32) {
if self
.sessions
.get(&correlation_id)
.is_some_and(StreamSession::is_closed)
{
self.sessions.remove(&correlation_id);
}
}
}
fn to_event(accepted: Accepted, payload: Vec<u8>) -> StreamEvent {
match accepted {
Accepted::Established => StreamEvent::Established,
Accepted::Data => StreamEvent::Data(payload),
Accepted::Credit(n) => StreamEvent::CreditGranted(n),
Accepted::RemoteClosed => StreamEvent::RemoteClosed,
Accepted::CloseAcked => StreamEvent::CloseAcked,
Accepted::Reset => StreamEvent::PeerReset,
}
}
#[cfg(test)]
mod tests {
use super::*;
use dig_identity::{derive_identity_sk, master_secret_key_from_seed, public_key_bytes};
use sha2::{Digest, Sha256};
const NOW: u64 = 1_700_000_000_000;
fn sk(label: &str) -> SecretKey {
let seed: [u8; 32] = Sha256::digest(label.as_bytes()).into();
derive_identity_sk(&master_secret_key_from_seed(&seed))
}
fn cid(tag: &str) -> Bytes32 {
Bytes32::new(Sha256::digest(tag.as_bytes()).into())
}
#[test]
fn initiator_handshake_then_data() {
let mut s = StreamSession::initiator(4);
assert_eq!(s.state(), StreamState::Opening);
assert!(s.build_data().is_err());
assert_eq!(
s.on_recv(StreamFrame::OpenAck, header(StreamFrame::OpenAck, 0, 2)),
Ok(Accepted::Established)
);
assert_eq!(s.state(), StreamState::Open);
assert_eq!(s.send_credit(), 2);
assert!(s.build_data().is_ok());
assert!(s.build_data().is_ok());
assert!(s.build_data().is_err());
}
#[test]
fn outbound_data_seq_increments() {
let mut s = StreamSession::initiator(0);
s.on_recv(StreamFrame::OpenAck, header(StreamFrame::OpenAck, 0, 3))
.unwrap();
assert_eq!(s.build_data().unwrap().seq, 0);
assert_eq!(s.build_data().unwrap().seq, 1);
assert_eq!(s.build_data().unwrap().seq, 2);
}
#[test]
fn responder_rejects_out_of_order_recv_seq() {
let mut s = StreamSession::responder(0);
s.build_open_ack(4).unwrap();
assert_eq!(
s.on_recv(StreamFrame::Data, header(StreamFrame::Data, 0, 0)),
Ok(Accepted::Data)
);
assert!(s
.on_recv(StreamFrame::Data, header(StreamFrame::Data, 2, 0))
.is_err());
assert!(s
.on_recv(StreamFrame::Data, header(StreamFrame::Data, 0, 0))
.is_err());
}
#[test]
fn recv_credit_window_bounds_inbound_data() {
let mut s = StreamSession::responder(0);
s.build_open_ack(1).unwrap(); assert_eq!(
s.on_recv(StreamFrame::Data, header(StreamFrame::Data, 0, 0)),
Ok(Accepted::Data)
);
assert!(s
.on_recv(StreamFrame::Data, header(StreamFrame::Data, 1, 0))
.is_err());
}
#[test]
fn credit_frame_relieves_send_backpressure() {
let mut s = StreamSession::initiator(0);
s.on_recv(StreamFrame::OpenAck, header(StreamFrame::OpenAck, 0, 1))
.unwrap();
s.build_data().unwrap();
assert!(s.build_data().is_err(), "credit exhausted");
assert_eq!(
s.on_recv(StreamFrame::Credit, header(StreamFrame::Credit, 0, 2)),
Ok(Accepted::Credit(2))
);
assert!(s.build_data().is_ok());
assert!(s.build_data().is_ok());
assert!(s.build_data().is_err());
}
#[test]
fn bidirectional_half_close() {
let mut s = StreamSession::initiator(4);
s.on_recv(StreamFrame::OpenAck, header(StreamFrame::OpenAck, 0, 4))
.unwrap();
s.build_close().unwrap();
assert_eq!(s.state(), StreamState::HalfClosedLocal);
assert!(s.build_data().is_err(), "no DATA after our CLOSE");
assert_eq!(
s.on_recv(StreamFrame::Data, header(StreamFrame::Data, 0, 0)),
Ok(Accepted::Data)
);
assert_eq!(
s.on_recv(StreamFrame::Close, header(StreamFrame::Close, 0, 0)),
Ok(Accepted::RemoteClosed)
);
assert_eq!(s.state(), StreamState::Closed);
assert!(s.is_closed());
}
#[test]
fn reset_aborts_from_any_state() {
let mut s = StreamSession::initiator(4);
assert_eq!(
s.on_recv(StreamFrame::Reset, header(StreamFrame::Reset, 0, 0)),
Ok(Accepted::Reset)
);
assert_eq!(s.state(), StreamState::Closed);
assert!(s
.on_recv(StreamFrame::Data, header(StreamFrame::Data, 0, 0))
.is_err());
}
struct Pair {
a_sk: SecretKey,
a_did: Bytes32,
a_pub: [u8; 48],
b_sk: SecretKey,
b_did: Bytes32,
b_pub: [u8; 48],
}
fn pair(tag: &str) -> Pair {
let a_sk = sk(&format!("{tag}/a"));
let b_sk = sk(&format!("{tag}/b"));
Pair {
a_pub: public_key_bytes(&a_sk),
b_pub: public_key_bytes(&b_sk),
a_did: cid(&format!("{tag}/a-did")),
b_did: cid(&format!("{tag}/b-did")),
a_sk,
b_sk,
}
}
const MT: u32 = 0x0000_0200;
#[test]
fn full_stream_round_trip_open_data_close() {
let p = pair("rt");
let mut alice = StreamEndpoint::new(&p.a_sk, p.a_did, 0, p.b_did, &p.b_pub, MT);
let mut bob = StreamEndpoint::new(&p.b_sk, p.b_did, 0, p.a_did, &p.a_pub, MT);
let sender_is_bob = |_d: Bytes32, _e: u32| Some(p.b_pub);
let sender_is_alice = |_d: Bytes32, _e: u32| Some(p.a_pub);
let stream = cid("rt/stream");
let open = alice.open(stream, 4, NOW, 0).unwrap();
assert!(matches!(
bob.accept(&open, sender_is_alice, NOW).unwrap(),
StreamAccept::Event(StreamEvent::Opened)
));
let ack = bob.open_ack(stream, 4, NOW, 0).unwrap();
assert!(matches!(
alice.accept(&ack, sender_is_bob, NOW).unwrap(),
StreamAccept::Event(StreamEvent::Established)
));
let d0 = alice.send_data(stream, b"hello ", NOW, 0).unwrap();
let d1 = alice.send_data(stream, b"world", NOW, 0).unwrap();
assert_eq!(
bob.accept(&d0, sender_is_alice, NOW).unwrap(),
StreamAccept::Event(StreamEvent::Data(b"hello ".to_vec()))
);
assert_eq!(
bob.accept(&d1, sender_is_alice, NOW).unwrap(),
StreamAccept::Event(StreamEvent::Data(b"world".to_vec()))
);
let close = alice.close(stream, NOW, 0).unwrap();
assert!(matches!(
bob.accept(&close, sender_is_alice, NOW).unwrap(),
StreamAccept::Event(StreamEvent::RemoteClosed)
));
}
#[test]
fn concurrent_stream_cap_rejects_the_nth_plus_one_open() {
let p = pair("cap");
let mut bob =
StreamEndpoint::new(&p.b_sk, p.b_did, 0, p.a_did, &p.a_pub, MT).with_max_concurrent(2);
let mut alice = StreamEndpoint::new(&p.a_sk, p.a_did, 0, p.b_did, &p.b_pub, MT);
let sender_is_alice = |_d: Bytes32, _e: u32| Some(p.a_pub);
for i in 0..2 {
let open = alice.open(cid(&format!("cap/{i}")), 1, NOW, 0).unwrap();
assert!(matches!(
bob.accept(&open, sender_is_alice, NOW).unwrap(),
StreamAccept::Event(StreamEvent::Opened)
));
}
assert_eq!(bob.stream_count(), 2);
let open3 = alice.open(cid("cap/3"), 1, NOW, 0).unwrap();
match bob.accept(&open3, sender_is_alice, NOW).unwrap() {
StreamAccept::Reset { cause, .. } => {
assert!(matches!(cause, MessageError::StreamLimit { cap: 2 }));
}
other => panic!("expected a StreamLimit RESET, got {other:?}"),
}
assert_eq!(
bob.stream_count(),
2,
"the rejected OPEN created no session"
);
}
#[test]
fn failed_verify_frame_is_dropped_never_reset() {
let p = pair("badverify");
let mut alice = StreamEndpoint::new(&p.a_sk, p.a_did, 0, p.b_did, &p.b_pub, MT);
let mut bob = StreamEndpoint::new(&p.b_sk, p.b_did, 0, p.a_did, &p.a_pub, MT);
let sender_is_alice = |_d: Bytes32, _e: u32| Some(p.a_pub);
let stream = cid("badverify/s");
let mut open = alice.open(stream, 4, NOW, 0).unwrap();
let last = open.sealed.ciphertext.len() - 1;
open.sealed.ciphertext[last] ^= 0x01;
assert_eq!(
bob.accept(&open, sender_is_alice, NOW).unwrap(),
StreamAccept::Dropped {
cause: MessageError::OpenFailed
}
);
assert_eq!(bob.stream_count(), 0);
}
#[test]
fn frame_for_unknown_stream_is_dropped_never_reset() {
let p = pair("proto");
let mut alice = StreamEndpoint::new(&p.a_sk, p.a_did, 0, p.b_did, &p.b_pub, MT);
let mut bob = StreamEndpoint::new(&p.b_sk, p.b_did, 0, p.a_did, &p.a_pub, MT);
let sender_is_alice = |_d: Bytes32, _e: u32| Some(p.a_pub);
let stream = cid("proto/s");
alice.open(stream, 4, NOW, 0).unwrap();
alice
.sessions
.get_mut(&stream)
.unwrap()
.on_recv(StreamFrame::OpenAck, header(StreamFrame::OpenAck, 0, 4))
.unwrap();
let data = alice.send_data(stream, b"early", NOW, 0).unwrap();
assert!(matches!(
bob.accept(&data, sender_is_alice, NOW).unwrap(),
StreamAccept::Dropped { .. }
));
}
#[test]
fn replayed_data_on_a_live_stream_is_dropped_and_stream_survives() {
let p = pair("replay");
let mut alice = StreamEndpoint::new(&p.a_sk, p.a_did, 0, p.b_did, &p.b_pub, MT);
let mut bob = StreamEndpoint::new(&p.b_sk, p.b_did, 0, p.a_did, &p.a_pub, MT);
let sender_is_alice = |_d: Bytes32, _e: u32| Some(p.a_pub);
let stream = cid("replay/s");
let open = alice.open(stream, 4, NOW, 0).unwrap();
bob.accept(&open, sender_is_alice, NOW).unwrap();
let ack = bob.open_ack(stream, 4, NOW, 0).unwrap();
alice.accept(&ack, |_d, _e| Some(p.b_pub), NOW).unwrap();
let data = alice.send_data(stream, b"once", NOW, 0).unwrap();
assert_eq!(
bob.accept(&data, sender_is_alice, NOW).unwrap(),
StreamAccept::Event(StreamEvent::Data(b"once".to_vec()))
);
assert_eq!(
bob.accept(&data, sender_is_alice, NOW).unwrap(),
StreamAccept::Dropped {
cause: MessageError::Replay
}
);
assert_eq!(bob.stream_count(), 1);
assert_eq!(bob.session(stream).unwrap().state(), StreamState::Open);
let next = alice.send_data(stream, b"twice", NOW, 0).unwrap();
assert_eq!(
bob.accept(&next, sender_is_alice, NOW).unwrap(),
StreamAccept::Event(StreamEvent::Data(b"twice".to_vec()))
);
}
#[test]
fn inbound_reset_for_unknown_stream_does_not_beget_a_reset() {
let p = pair("resetstorm");
let mut alice = StreamEndpoint::new(&p.a_sk, p.a_did, 0, p.b_did, &p.b_pub, MT);
let mut bob = StreamEndpoint::new(&p.b_sk, p.b_did, 0, p.a_did, &p.a_pub, MT);
let sender_is_alice = |_d: Bytes32, _e: u32| Some(p.a_pub);
let ghost = cid("resetstorm/ghost");
alice.open(ghost, 1, NOW, 0).unwrap();
let reset = alice.reset(ghost, NOW, 0).unwrap();
assert!(matches!(
bob.accept(&reset, sender_is_alice, NOW).unwrap(),
StreamAccept::Dropped { .. }
));
assert_eq!(bob.stream_count(), 0);
}
#[test]
fn protocol_violation_on_established_stream_still_resets() {
let p = pair("legitreset");
let mut alice = StreamEndpoint::new(&p.a_sk, p.a_did, 0, p.b_did, &p.b_pub, MT);
let mut bob = StreamEndpoint::new(&p.b_sk, p.b_did, 0, p.a_did, &p.a_pub, MT);
let sender_is_alice = |_d: Bytes32, _e: u32| Some(p.a_pub);
let sender_is_bob = |_d: Bytes32, _e: u32| Some(p.b_pub);
let stream = cid("legitreset/s");
let open = alice.open(stream, 4, NOW, 0).unwrap();
bob.accept(&open, sender_is_alice, NOW).unwrap();
let ack = bob.open_ack(stream, 4, NOW, 0).unwrap();
alice.accept(&ack, sender_is_bob, NOW).unwrap();
alice.sessions.get_mut(&stream).unwrap().send_seq = 1;
let bad = alice.send_data(stream, b"gap", NOW, 0).unwrap();
match bob.accept(&bad, sender_is_alice, NOW).unwrap() {
StreamAccept::Reset { cause, .. } => {
assert!(matches!(cause, MessageError::StreamProtocol(_)));
}
other => panic!("expected a RESET on the authenticated violation, got {other:?}"),
}
assert_eq!(bob.stream_count(), 0, "the violated session is dropped");
}
#[test]
fn authenticated_frame_from_a_non_peer_sender_is_hard_dropped() {
let p = pair("senderbind");
let eve_sk = sk("senderbind/eve");
let eve_did = cid("senderbind/eve-did");
let eve_pub = public_key_bytes(&eve_sk);
let mut bob = StreamEndpoint::new(&p.b_sk, p.b_did, 0, p.a_did, &p.a_pub, MT);
let mut eve = StreamEndpoint::new(&eve_sk, eve_did, 0, p.b_did, &p.b_pub, MT);
let permissive_resolver = move |d: Bytes32, _e: u32| {
if d == eve_did {
Some(eve_pub)
} else {
Some(p.a_pub)
}
};
let stream = cid("senderbind/s");
let open = eve.open(stream, 4, NOW, 0).unwrap();
assert_eq!(
bob.accept(&open, permissive_resolver, NOW).unwrap(),
StreamAccept::Dropped {
cause: MessageError::StreamProtocol(
"authenticated sender is not this stream's peer"
)
}
);
assert_eq!(
bob.stream_count(),
0,
"no session opened for the wrong peer"
);
}
#[test]
fn duplicate_open_for_a_live_stream_drops_the_local_session_too() {
let p = pair("dupopen");
let mut alice = StreamEndpoint::new(&p.a_sk, p.a_did, 0, p.b_did, &p.b_pub, MT);
let mut bob = StreamEndpoint::new(&p.b_sk, p.b_did, 0, p.a_did, &p.a_pub, MT);
let sender_is_alice = |_d: Bytes32, _e: u32| Some(p.a_pub);
let stream = cid("dupopen/s");
let open = alice.open(stream, 4, NOW, 0).unwrap();
assert!(matches!(
bob.accept(&open, sender_is_alice, NOW).unwrap(),
StreamAccept::Event(StreamEvent::Opened)
));
assert_eq!(bob.stream_count(), 1);
alice.sessions.remove(&stream);
let dup_open = alice.open(stream, 4, NOW, 0).unwrap();
match bob.accept(&dup_open, sender_is_alice, NOW).unwrap() {
StreamAccept::Reset { cause, .. } => {
assert!(matches!(cause, MessageError::StreamProtocol(_)));
}
other => panic!("expected a RESET on the duplicate OPEN, got {other:?}"),
}
assert_eq!(
bob.stream_count(),
0,
"the stale local session must not survive the duplicate OPEN"
);
}
#[test]
fn every_frame_uses_a_unique_ephemeral() {
let p = pair("uniq");
let mut alice = StreamEndpoint::new(&p.a_sk, p.a_did, 0, p.b_did, &p.b_pub, MT);
let stream = cid("uniq/s");
let mut kems = Vec::new();
kems.push(alice.open(stream, 100, NOW, 0).unwrap().sealed.kem_enc);
alice
.sessions
.get_mut(&stream)
.unwrap()
.on_recv(StreamFrame::OpenAck, header(StreamFrame::OpenAck, 0, 100))
.unwrap();
for _ in 0..16 {
kems.push(
alice
.send_data(stream, b"x", NOW, 0)
.unwrap()
.sealed
.kem_enc,
);
}
let unique: std::collections::HashSet<_> = kems.iter().collect();
assert_eq!(
unique.len(),
kems.len(),
"every frame must use a fresh ephemeral"
);
}
#[test]
fn credit_grant_close_ack_and_peer_reset_round_trip() {
let p = pair("credit");
let mut alice = StreamEndpoint::new(&p.a_sk, p.a_did, 0, p.b_did, &p.b_pub, MT);
let mut bob = StreamEndpoint::new(&p.b_sk, p.b_did, 0, p.a_did, &p.a_pub, MT);
let sender_is_alice = |_d: Bytes32, _e: u32| Some(p.a_pub);
let sender_is_bob = |_d: Bytes32, _e: u32| Some(p.b_pub);
let stream = cid("credit/s");
let open = alice.open(stream, 1, NOW, 0).unwrap();
bob.accept(&open, sender_is_alice, NOW).unwrap();
let ack = bob.open_ack(stream, 1, NOW, 0).unwrap();
alice.accept(&ack, sender_is_bob, NOW).unwrap();
let credit = bob.grant_credit(stream, 3, NOW, 0).unwrap();
assert_eq!(
alice.accept(&credit, sender_is_bob, NOW).unwrap(),
StreamAccept::Event(StreamEvent::CreditGranted(3))
);
assert_eq!(alice.session(stream).unwrap().send_credit(), 4);
let reset = bob.reset(stream, NOW, 0).unwrap();
assert_eq!(
alice.accept(&reset, sender_is_bob, NOW).unwrap(),
StreamAccept::Event(StreamEvent::PeerReset)
);
assert_eq!(alice.stream_count(), 0);
assert_eq!(
bob.stream_count(),
0,
"reset drops the sender's session too"
);
}
#[test]
fn close_ack_is_delivered() {
let p = pair("closeack");
let mut alice = StreamEndpoint::new(&p.a_sk, p.a_did, 0, p.b_did, &p.b_pub, MT);
let mut bob = StreamEndpoint::new(&p.b_sk, p.b_did, 0, p.a_did, &p.a_pub, MT);
let sender_is_alice = |_d: Bytes32, _e: u32| Some(p.a_pub);
let sender_is_bob = |_d: Bytes32, _e: u32| Some(p.b_pub);
let stream = cid("closeack/s");
let open = alice.open(stream, 2, NOW, 0).unwrap();
bob.accept(&open, sender_is_alice, NOW).unwrap();
let ack = bob.open_ack(stream, 2, NOW, 0).unwrap();
alice.accept(&ack, sender_is_bob, NOW).unwrap();
let close = alice.close(stream, NOW, 0).unwrap();
bob.accept(&close, sender_is_alice, NOW).unwrap();
let close_ack = bob
.seal_frame(stream, header(StreamFrame::CloseAck, 0, 0), &[], NOW, 0)
.unwrap();
assert_eq!(
alice.accept(&close_ack, sender_is_bob, NOW).unwrap(),
StreamAccept::Event(StreamEvent::CloseAcked)
);
}
#[test]
fn endpoint_send_side_error_branches() {
let p = pair("errs");
let mut alice = StreamEndpoint::new(&p.a_sk, p.a_did, 0, p.b_did, &p.b_pub, MT);
let stream = cid("errs/s");
assert!(matches!(
alice.send_data(stream, b"x", NOW, 0),
Err(MessageError::StreamProtocol(_))
));
alice.open(stream, 1, NOW, 0).unwrap();
assert!(matches!(
alice.open(stream, 1, NOW, 0),
Err(MessageError::StreamProtocol(_))
));
let mut tight =
StreamEndpoint::new(&p.a_sk, p.a_did, 0, p.b_did, &p.b_pub, MT).with_max_concurrent(1);
tight.open(cid("errs/only"), 1, NOW, 0).unwrap();
assert!(matches!(
tight.open(cid("errs/2"), 1, NOW, 0),
Err(MessageError::StreamLimit { cap: 1 })
));
}
#[test]
fn half_closed_remote_state_then_local_close() {
let mut s = StreamSession::initiator(4);
s.on_recv(StreamFrame::OpenAck, header(StreamFrame::OpenAck, 0, 4))
.unwrap();
s.on_recv(StreamFrame::Close, header(StreamFrame::Close, 0, 0))
.unwrap();
assert_eq!(s.state(), StreamState::HalfClosedRemote);
assert!(
s.build_data().is_ok(),
"we may still send after the peer's CLOSE"
);
s.build_close().unwrap();
assert_eq!(s.state(), StreamState::Closed);
}
}