#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum Role {
Client,
Server,
}
impl Role {
pub(crate) const fn peer(self) -> Self {
match self {
Self::Client => Self::Server,
Self::Server => Self::Client,
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum State {
Idle,
SendResponse,
SendBody,
Done,
MightSwitchProtocol,
SwitchedProtocol,
MustClose,
Closed,
Error,
}
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
pub(crate) struct SwitchProposals {
connect: bool,
upgrade: bool,
}
impl SwitchProposals {
pub(crate) const NONE: Self = Self {
connect: false,
upgrade: false,
};
#[cfg(test)]
const CONNECT: Self = Self {
connect: true,
upgrade: false,
};
#[cfg(test)]
const UPGRADE: Self = Self {
connect: false,
upgrade: true,
};
const fn contains(self, kind: SwitchKind) -> bool {
match kind {
SwitchKind::Connect => self.connect,
SwitchKind::Upgrade => self.upgrade,
}
}
pub(crate) const fn from_flags(connect: bool, upgrade: bool) -> Self {
Self { connect, upgrade }
}
const fn is_empty(self) -> bool {
!self.connect && !self.upgrade
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) enum SwitchKind {
Connect,
Upgrade,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) enum StateEvent {
Request(SwitchProposals),
InformationalResponse,
Response,
Data,
EndOfMessage,
ProtocolSwitch(SwitchKind),
ConnectionClosed,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) struct InvalidTransition;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) struct ConnectionState {
client: State,
server: State,
keep_alive: bool,
proposals: SwitchProposals,
}
impl ConnectionState {
pub(crate) const fn new() -> Self {
Self {
client: State::Idle,
server: State::Idle,
keep_alive: true,
proposals: SwitchProposals::NONE,
}
}
pub(crate) fn state(self, role: Role) -> State {
match role {
Role::Client => self.client,
Role::Server => self.server,
}
}
pub(crate) fn process_event(
&mut self,
role: Role,
event: StateEvent,
) -> Result<(), InvalidTransition> {
let mut next = *self;
let current = match role {
Role::Client => next.client,
Role::Server => next.server,
};
let transitioned = direct_transition(role, current, event).ok_or(InvalidTransition)?;
match (role, event) {
(Role::Client, StateEvent::Request(proposals)) => {
if next.server != State::Idle {
return Err(InvalidTransition);
}
next.proposals = proposals;
next.server = State::SendResponse;
}
(Role::Server, StateEvent::ProtocolSwitch(kind)) => {
if !next.proposals.contains(kind) {
return Err(InvalidTransition);
}
}
(Role::Server, StateEvent::Response) => {
if next.proposals.contains(SwitchKind::Connect) {
next.keep_alive = false;
}
next.proposals = SwitchProposals::NONE;
if current == State::Idle {
next.keep_alive = false;
}
}
_ => {}
}
match role {
Role::Client => next.client = transitioned,
Role::Server => next.server = transitioned,
}
next.stabilize();
*self = next;
Ok(())
}
pub(crate) fn disable_keep_alive(&mut self) {
self.keep_alive = false;
self.stabilize();
}
pub(crate) const fn keep_alive(self) -> bool {
self.keep_alive
}
pub(crate) fn process_error(&mut self, role: Role) {
match role {
Role::Client => self.client = State::Error,
Role::Server => self.server = State::Error,
}
self.stabilize();
}
pub(crate) fn start_next_cycle(&mut self) -> Result<(), InvalidTransition> {
if self.client != State::Done
|| self.server != State::Done
|| !self.keep_alive
|| !self.proposals.is_empty()
{
return Err(InvalidTransition);
}
*self = Self::new();
Ok(())
}
fn stabilize(&mut self) {
loop {
let before = *self;
if self.client == State::Done && !self.proposals.is_empty() {
self.client = State::MightSwitchProtocol;
}
if self.client == State::MightSwitchProtocol && self.server == State::SwitchedProtocol {
self.client = State::SwitchedProtocol;
} else if self.client == State::MightSwitchProtocol && self.proposals.is_empty() {
self.client = State::Done;
}
if self.client == State::SwitchedProtocol && self.server == State::SwitchedProtocol {
self.proposals = SwitchProposals::NONE;
}
if !self.keep_alive {
if self.client == State::Done {
self.client = State::MustClose;
}
if self.server == State::Done {
self.server = State::MustClose;
}
}
match (self.client, self.server) {
(State::Closed, State::Done | State::Idle) | (State::Error, State::Done) => {
self.server = State::MustClose
}
(State::Done | State::Idle, State::Closed) | (State::Done, State::Error) => {
self.client = State::MustClose
}
_ => {}
}
if *self == before {
return;
}
}
}
}
fn direct_transition(role: Role, state: State, event: StateEvent) -> Option<State> {
match role {
Role::Client => match state {
State::Idle => match event {
StateEvent::Request(_) => Some(State::SendBody),
StateEvent::ConnectionClosed => Some(State::Closed),
event => reject(event),
},
State::SendBody => match event {
StateEvent::Data => Some(State::SendBody),
StateEvent::EndOfMessage => Some(State::Done),
event => reject(event),
},
State::Done | State::MustClose | State::Closed => match event {
StateEvent::ConnectionClosed => Some(State::Closed),
event => reject(event),
},
State::SendResponse
| State::MightSwitchProtocol
| State::SwitchedProtocol
| State::Error => reject(event),
},
Role::Server => match state {
State::Idle => match event {
StateEvent::Response => Some(State::SendBody),
StateEvent::ConnectionClosed => Some(State::Closed),
event => reject(event),
},
State::SendResponse => match event {
StateEvent::InformationalResponse => Some(State::SendResponse),
StateEvent::Response => Some(State::SendBody),
StateEvent::ProtocolSwitch(_) => Some(State::SwitchedProtocol),
event => reject(event),
},
State::SendBody => match event {
StateEvent::Data => Some(State::SendBody),
StateEvent::EndOfMessage => Some(State::Done),
event => reject(event),
},
State::Done | State::MustClose | State::Closed => match event {
StateEvent::ConnectionClosed => Some(State::Closed),
event => reject(event),
},
State::MightSwitchProtocol | State::SwitchedProtocol | State::Error => reject(event),
},
}
}
fn reject(event: StateEvent) -> Option<State> {
match event {
StateEvent::Request(_)
| StateEvent::InformationalResponse
| StateEvent::Response
| StateEvent::Data
| StateEvent::EndOfMessage
| StateEvent::ProtocolSwitch(_)
| StateEvent::ConnectionClosed => None,
}
}
#[cfg(test)]
mod tests {
use super::*;
use Role::*;
use State::*;
use StateEvent as Event;
use StateEvent::*;
const STATES: [State; 9] = [
State::Idle,
State::SendResponse,
State::SendBody,
State::Done,
State::MightSwitchProtocol,
State::SwitchedProtocol,
State::MustClose,
State::Closed,
State::Error,
];
const EVENTS: [StateEvent; 8] = [
StateEvent::Request(SwitchProposals::NONE),
StateEvent::InformationalResponse,
StateEvent::Response,
StateEvent::Data,
StateEvent::EndOfMessage,
StateEvent::ProtocolSwitch(SwitchKind::Connect),
StateEvent::ProtocolSwitch(SwitchKind::Upgrade),
StateEvent::ConnectionClosed,
];
const DIRECT_CONTRACT: &[(Role, State, StateEvent, State)] = &[
(Client, Idle, Request(SwitchProposals::NONE), SendBody),
(Client, Idle, ConnectionClosed, Closed),
(Client, SendBody, Data, SendBody),
(Client, SendBody, EndOfMessage, Done),
(Client, Done, ConnectionClosed, Closed),
(Client, MustClose, ConnectionClosed, Closed),
(Client, Closed, ConnectionClosed, Closed),
(Server, Idle, Response, SendBody),
(Server, Idle, ConnectionClosed, Closed),
(Server, SendResponse, InformationalResponse, SendResponse),
(Server, SendResponse, Response, SendBody),
(
Server,
SendResponse,
ProtocolSwitch(SwitchKind::Connect),
SwitchedProtocol,
),
(
Server,
SendResponse,
ProtocolSwitch(SwitchKind::Upgrade),
SwitchedProtocol,
),
(Server, SendBody, Data, SendBody),
(Server, SendBody, EndOfMessage, Done),
(Server, Done, ConnectionClosed, Closed),
(Server, MustClose, ConnectionClosed, Closed),
(Server, Closed, ConnectionClosed, Closed),
];
#[test]
fn direct_transitions_match_the_protocol_contract() {
for role in [Role::Client, Role::Server] {
for state in STATES {
for event in EVENTS {
let expected = DIRECT_CONTRACT.iter().find_map(
|&(case_role, case_state, case_event, next)| {
(role == case_role && state == case_state && event == case_event)
.then_some(next)
},
);
assert_eq!(
direct_transition(role, state, event),
expected,
"{role:?} {state:?} {event:?}"
);
}
}
}
}
#[test]
fn a_complete_exchange_can_be_reused() {
let mut state = ConnectionState::new();
state
.process_event(Role::Client, Event::Request(SwitchProposals::NONE))
.unwrap();
assert_eq!(
(state.client, state.server),
(State::SendBody, State::SendResponse)
);
state.process_event(Role::Client, Event::Data).unwrap();
state
.process_event(Role::Client, Event::EndOfMessage)
.unwrap();
state
.process_event(Role::Server, Event::InformationalResponse)
.unwrap();
state.process_event(Role::Server, Event::Response).unwrap();
state.process_event(Role::Server, Event::Data).unwrap();
state
.process_event(Role::Server, Event::EndOfMessage)
.unwrap();
assert_eq!((state.client, state.server), (State::Done, State::Done));
state.start_next_cycle().unwrap();
assert_eq!(state, ConnectionState::new());
}
#[test]
fn invalid_input_is_rejected_atomically() {
let mut state = ConnectionState::new();
state
.process_event(Role::Client, Event::Request(SwitchProposals::NONE))
.unwrap();
let before = state;
assert_eq!(
state.process_event(Role::Client, Event::Request(SwitchProposals::NONE)),
Err(InvalidTransition)
);
assert_eq!(state, before);
assert_eq!(
state.process_event(Role::Server, Event::ProtocolSwitch(SwitchKind::Upgrade)),
Err(InvalidTransition)
);
assert_eq!(state, before);
}
#[test]
fn disabling_keep_alive_forces_close_after_each_message() {
let mut state = ConnectionState::new();
state
.process_event(Role::Client, Event::Request(SwitchProposals::NONE))
.unwrap();
state.disable_keep_alive();
state
.process_event(Role::Client, Event::EndOfMessage)
.unwrap();
assert_eq!(state.client, State::MustClose);
state.process_event(Role::Server, Event::Response).unwrap();
state
.process_event(Role::Server, Event::EndOfMessage)
.unwrap();
assert_eq!(state.server, State::MustClose);
assert_eq!(state.start_next_cycle(), Err(InvalidTransition));
}
#[test]
fn proposed_upgrade_and_connect_can_switch_protocols() {
for (proposals, kind) in [
(SwitchProposals::UPGRADE, SwitchKind::Upgrade),
(SwitchProposals::CONNECT, SwitchKind::Connect),
] {
let mut state = ConnectionState::new();
state
.process_event(Role::Client, Event::Request(proposals))
.unwrap();
state.disable_keep_alive();
state
.process_event(Role::Server, Event::ProtocolSwitch(kind))
.unwrap();
assert_eq!(
(state.client, state.server),
(State::SendBody, State::SwitchedProtocol)
);
state
.process_event(Role::Client, Event::EndOfMessage)
.unwrap();
assert_eq!(
(state.client, state.server),
(State::SwitchedProtocol, State::SwitchedProtocol)
);
assert_eq!(state.proposals, SwitchProposals::NONE);
}
}
#[test]
fn switch_selection_must_match_one_of_the_request_proposals() {
for selected in [SwitchKind::Connect, SwitchKind::Upgrade] {
let mut state = ConnectionState::new();
state
.process_event(
Role::Client,
Event::Request(SwitchProposals::from_flags(true, true)),
)
.unwrap();
state
.process_event(Role::Server, Event::ProtocolSwitch(selected))
.unwrap();
assert_eq!(state.server, State::SwitchedProtocol);
}
for (proposed, selected) in [
(SwitchProposals::CONNECT, SwitchKind::Upgrade),
(SwitchProposals::UPGRADE, SwitchKind::Connect),
] {
let mut state = ConnectionState::new();
state
.process_event(Role::Client, Event::Request(proposed))
.unwrap();
let before = state;
assert_eq!(
state.process_event(Role::Server, Event::ProtocolSwitch(selected)),
Err(InvalidTransition)
);
assert_eq!(state, before);
}
}
#[test]
fn denied_switch_restores_the_close_rule() {
let mut state = ConnectionState::new();
state
.process_event(Role::Client, Event::Request(SwitchProposals::UPGRADE))
.unwrap();
state.disable_keep_alive();
state
.process_event(Role::Client, Event::EndOfMessage)
.unwrap();
assert_eq!(state.client, State::MightSwitchProtocol);
state.process_event(Role::Server, Event::Response).unwrap();
assert_eq!(
(state.client, state.server),
(State::MustClose, State::SendBody)
);
assert_eq!(state.proposals, SwitchProposals::NONE);
}
#[test]
fn direct_error_response_and_protocol_errors_require_close() {
let mut response = ConnectionState::new();
response
.process_event(Role::Server, Event::Response)
.unwrap();
assert!(!response.keep_alive);
response
.process_event(Role::Server, Event::EndOfMessage)
.unwrap();
assert_eq!(response.server, State::MustClose);
let mut error = ConnectionState::new();
error
.process_event(Role::Client, Event::Request(SwitchProposals::NONE))
.unwrap();
error
.process_event(Role::Client, Event::EndOfMessage)
.unwrap();
error.process_error(Role::Server);
assert_eq!(
(error.client, error.server),
(State::MustClose, State::Error)
);
let mut closed = ConnectionState {
client: State::Done,
server: State::Done,
keep_alive: true,
proposals: SwitchProposals::NONE,
};
closed
.process_event(Role::Server, Event::ConnectionClosed)
.unwrap();
assert_eq!(
(closed.client, closed.server),
(State::MustClose, State::Closed)
);
}
}