use crate::{error::ProtocolError, options::SessionOptions, proto::parser::Response};
const LOG_TARGET: &str = "yosemite::proto::session";
#[derive(Debug, PartialEq, Eq, Clone)]
enum StreamKind {
Accept,
Connect,
Forward,
}
#[derive(Debug, PartialEq, Eq, Clone)]
enum StreamState {
Uninitialized,
Handshaking,
Handshaked,
Pending(StreamKind),
}
#[derive(Debug, PartialEq, Eq, Clone)]
enum SessionState {
Uninitialized,
Handshaking,
Handshaked,
SessionCreatePending,
Active {
destination: String,
stream_state: StreamState,
},
Poisoned,
}
#[derive(Clone)]
pub struct SessionController {
options: SessionOptions,
state: SessionState,
}
impl SessionController {
pub fn new(options: SessionOptions) -> Result<Self, ProtocolError> {
Ok(Self {
options,
state: SessionState::Uninitialized,
})
}
pub fn handshake_session(&mut self) -> Result<Vec<u8>, ProtocolError> {
match std::mem::replace(&mut self.state, SessionState::Poisoned) {
SessionState::Uninitialized => {
tracing::trace!(
target: LOG_TARGET,
nickname = %self.options.nickname,
"send handshake for session",
);
self.state = SessionState::Handshaking;
Ok(String::from("HELLO VERSION\n").into_bytes())
}
state => {
tracing::warn!(
target: LOG_TARGET,
?state,
"cannot create session, invalid state",
);
debug_assert!(false);
Err(ProtocolError::InvalidState)
}
}
}
pub fn create_session(&mut self, command: String) -> Result<Vec<u8>, ProtocolError> {
match std::mem::replace(&mut self.state, SessionState::Poisoned) {
SessionState::Handshaked => {
tracing::trace!(
target: LOG_TARGET,
nickname = %self.options.nickname,
destination = ?self.options.destination,
"create new session",
);
self.state = SessionState::SessionCreatePending;
Ok(command.into_bytes())
}
state => {
tracing::warn!(
target: LOG_TARGET,
?state,
"cannot create session, invalid state",
);
debug_assert!(false);
Err(ProtocolError::InvalidState)
}
}
}
pub fn handshake_stream(&mut self) -> Result<Vec<u8>, ProtocolError> {
match std::mem::replace(&mut self.state, SessionState::Poisoned) {
SessionState::Active {
destination,
stream_state: StreamState::Uninitialized,
} => {
tracing::trace!(
target: LOG_TARGET,
nickname = %self.options.nickname,
"send handshake for stream",
);
self.state = SessionState::Active {
destination,
stream_state: StreamState::Handshaking,
};
Ok(String::from("HELLO VERSION\n").into_bytes())
}
state => {
tracing::warn!(
target: LOG_TARGET,
?state,
"cannot create session, invalid state",
);
debug_assert!(false);
Err(ProtocolError::InvalidState)
}
}
}
pub fn create_stream(&mut self, remote_destination: &str) -> Result<Vec<u8>, ProtocolError> {
match std::mem::replace(&mut self.state, SessionState::Poisoned) {
SessionState::Active {
destination,
stream_state: StreamState::Handshaked,
} => {
tracing::info!(
target: LOG_TARGET,
nickname = %self.options.nickname,
remote_destination = %format!("{}...", &destination[..10]),
"open stream to remote destination",
);
self.state = SessionState::Active {
destination,
stream_state: StreamState::Pending(StreamKind::Connect),
};
Ok(format!(
"STREAM CONNECT ID={} DESTINATION={} SILENT=false\n",
self.options.nickname, remote_destination
)
.into_bytes())
}
state => {
tracing::warn!(
target: LOG_TARGET,
?state,
"cannot create session, invalid state",
);
debug_assert!(false);
Err(ProtocolError::InvalidState)
}
}
}
pub fn accept_stream(&mut self) -> Result<Vec<u8>, ProtocolError> {
match std::mem::replace(&mut self.state, SessionState::Poisoned) {
SessionState::Active {
destination,
stream_state: StreamState::Handshaked,
} => {
tracing::trace!(
target: LOG_TARGET,
nickname = %self.options.nickname,
"start listening for virtual stream",
);
self.state = SessionState::Active {
destination,
stream_state: StreamState::Pending(StreamKind::Accept),
};
Ok(
format!("STREAM ACCEPT ID={} SILENT=false\n", self.options.nickname)
.into_bytes(),
)
}
state => {
tracing::warn!(
target: LOG_TARGET,
?state,
"cannot create session, invalid state",
);
debug_assert!(false);
Err(ProtocolError::InvalidState)
}
}
}
pub fn forward_stream(&mut self, port: u16) -> Result<Vec<u8>, ProtocolError> {
match std::mem::replace(&mut self.state, SessionState::Poisoned) {
SessionState::Active {
destination,
stream_state: StreamState::Handshaked,
} => {
tracing::trace!(
target: LOG_TARGET,
nickname = %self.options.nickname,
?port,
"forward incoming connections",
);
self.state = SessionState::Active {
destination,
stream_state: StreamState::Pending(StreamKind::Forward),
};
Ok(format!(
"STREAM FORWARD ID={} PORT={port} SILENT={}\n",
self.options.nickname,
self.options.silent_forward.to_string(),
)
.into_bytes())
}
state => {
tracing::warn!(
target: LOG_TARGET,
?state,
"cannot create session, invalid state",
);
debug_assert!(false);
Err(ProtocolError::InvalidState)
}
}
}
pub fn handle_response(&mut self, response: &str) -> Result<(), ProtocolError> {
match std::mem::replace(&mut self.state, SessionState::Poisoned) {
SessionState::Handshaking => match Response::parse(response) {
Some(Response::Hello {
version: Ok(version),
}) => {
tracing::trace!(
target: LOG_TARGET,
nickname = %self.options.nickname,
%version,
"session handshake done",
);
self.state = SessionState::Handshaked;
Ok(())
}
Some(Response::Hello {
version: Err(error),
}) => return Err(ProtocolError::Router(error)),
None => {
tracing::warn!(
target: LOG_TARGET,
nickname = %self.options.nickname,
?response,
"invalid response from router session `HELLO`",
);
return Err(ProtocolError::InvalidMessage);
}
Some(response) => {
tracing::warn!(
nickname = %self.options.nickname,
?response,
"unexpected response from router session `HELLO`",
);
return Err(ProtocolError::InvalidState);
}
},
SessionState::SessionCreatePending => match Response::parse(response) {
Some(Response::Session {
destination: Ok(destination),
}) => {
tracing::info!(
target: LOG_TARGET,
nickname = %self.options.nickname,
"session created",
);
self.state = SessionState::Active {
destination,
stream_state: StreamState::Uninitialized,
};
Ok(())
}
Some(Response::Session {
destination: Err(error),
}) => return Err(ProtocolError::Router(error)),
None => {
tracing::warn!(
target: LOG_TARGET,
nickname = %self.options.nickname,
?response,
"invalid response from router `SESSION CREATE`",
);
return Err(ProtocolError::InvalidMessage);
}
Some(response) => {
tracing::warn!(
nickname = %self.options.nickname,
?response,
"unexpected response from router to `SESSION CREATE`",
);
return Err(ProtocolError::InvalidState);
}
},
SessionState::Active {
destination,
stream_state: StreamState::Handshaking,
} => match Response::parse(response) {
Some(Response::Hello {
version: Ok(version),
}) => {
tracing::trace!(
target: LOG_TARGET,
nickname = %self.options.nickname,
%version,
"stream handshake done",
);
self.state = SessionState::Active {
destination,
stream_state: StreamState::Handshaked,
};
Ok(())
}
Some(Response::Hello {
version: Err(error),
}) => return Err(ProtocolError::Router(error)),
None => {
tracing::warn!(
target: LOG_TARGET,
nickname = %self.options.nickname,
?response,
"invalid response from router stream `HELLO`",
);
return Err(ProtocolError::InvalidMessage);
}
Some(response) => {
tracing::warn!(
nickname = %self.options.nickname,
?response,
"unexpected response from router stream `HELLO`",
);
return Err(ProtocolError::InvalidState);
}
},
SessionState::Active {
destination,
stream_state: StreamState::Pending(direction),
} => match Response::parse(response) {
Some(Response::Stream { result: Ok(()) }) => {
tracing::info!(
target: LOG_TARGET,
nickname = %self.options.nickname,
?direction,
"stream status ok",
);
self.state = SessionState::Active {
destination,
stream_state: StreamState::Uninitialized,
};
Ok(())
}
Some(Response::Stream { result: Err(error) }) =>
return Err(ProtocolError::Router(error)),
None => {
tracing::warn!(
target: LOG_TARGET,
nickname = %self.options.nickname,
?response,
?direction,
"invalid response from router to `STREAM CREATE`",
);
return Err(ProtocolError::InvalidMessage);
}
Some(response) => {
tracing::warn!(
nickname = %self.options.nickname,
?response,
?direction,
"unexpected response from router to `STREAM CREATE`",
);
return Err(ProtocolError::InvalidState);
}
},
state => {
tracing::warn!(
target: LOG_TARGET,
?state,
"cannot handle response, invalid state",
);
debug_assert!(false);
Err(ProtocolError::InvalidState)
}
}
}
pub fn destination(&self) -> &str {
let SessionState::Active { destination, .. } = &self.state else {
panic!("invalid state");
};
&destination
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn open_virtual_stream() {
let mut controller = SessionController::new(Default::default()).unwrap();
assert_eq!(controller.state, SessionState::Uninitialized);
assert_eq!(
controller.handshake_session(),
Ok(String::from("HELLO VERSION\n").into_bytes())
);
assert_eq!(controller.state, SessionState::Handshaking);
assert!(controller.handle_response("HELLO REPLY RESULT=OK VERSION=3.3\n").is_ok());
assert_eq!(controller.state, SessionState::Handshaked);
let command = "SESSION CREATE \
STYLE=STREAM \
ID=nickname \
DESTINATION=TRANSIENT \
SIGNATURE_TYPE=7 \
i2cp.leaseSetEncType=4\n";
assert!(controller.create_session(command.to_string()).is_ok());
assert_eq!(controller.state, SessionState::SessionCreatePending);
assert!(controller
.handle_response("SESSION STATUS RESULT=OK DESTINATION=I2P_DESTINATION\n")
.is_ok());
match &controller.state {
SessionState::Active { destination, .. }
if destination.as_str() == "I2P_DESTINATION" => {}
state => panic!("invalid state: {state:?}"),
}
assert!(controller.handshake_stream().is_ok());
let SessionState::Active {
stream_state: StreamState::Handshaking,
..
} = controller.state
else {
panic!("invalid state");
};
assert!(controller.handle_response("HELLO REPLY RESULT=OK VERSION=3.3\n").is_ok());
let SessionState::Active {
stream_state: StreamState::Handshaked,
..
} = controller.state
else {
panic!("invalid state");
};
assert!(controller.create_stream("destination").is_ok());
let SessionState::Active {
stream_state: StreamState::Pending(StreamKind::Connect),
..
} = controller.state
else {
panic!("invalid state");
};
assert!(controller.handle_response("STREAM STATUS RESULT=OK\n").is_ok());
let SessionState::Active {
stream_state: StreamState::Uninitialized,
..
} = controller.state
else {
panic!("invalid state");
};
}
#[test]
fn accept_virtual_stream() {
let mut controller = SessionController::new(Default::default()).unwrap();
assert_eq!(controller.state, SessionState::Uninitialized);
assert_eq!(
controller.handshake_session(),
Ok(String::from("HELLO VERSION\n").into_bytes())
);
assert_eq!(controller.state, SessionState::Handshaking);
assert!(controller.handle_response("HELLO REPLY RESULT=OK VERSION=3.3\n").is_ok());
assert_eq!(controller.state, SessionState::Handshaked);
let command = "SESSION CREATE \
STYLE=STREAM \
ID=nickname \
DESTINATION=TRANSIENT \
SIGNATURE_TYPE=7 \
i2cp.leaseSetEncType=4\n";
assert!(controller.create_session(command.to_string()).is_ok());
assert_eq!(controller.state, SessionState::SessionCreatePending);
assert!(controller
.handle_response("SESSION STATUS RESULT=OK DESTINATION=I2P_DESTINATION\n")
.is_ok());
match &controller.state {
SessionState::Active { destination, .. }
if destination.as_str() == "I2P_DESTINATION" => {}
state => panic!("invalid state: {state:?}"),
}
assert!(controller.handshake_stream().is_ok());
let SessionState::Active {
stream_state: StreamState::Handshaking,
..
} = controller.state
else {
panic!("invalid state");
};
assert!(controller.handle_response("HELLO REPLY RESULT=OK VERSION=3.3\n").is_ok());
let SessionState::Active {
stream_state: StreamState::Handshaked,
..
} = controller.state
else {
panic!("invalid state");
};
assert!(controller.accept_stream().is_ok());
let SessionState::Active {
stream_state: StreamState::Pending(StreamKind::Accept),
..
} = controller.state
else {
panic!("invalid state");
};
assert!(controller.handle_response("STREAM STATUS RESULT=OK\n").is_ok());
let SessionState::Active {
stream_state: StreamState::Uninitialized,
..
} = controller.state
else {
panic!("invalid state");
};
}
}