use std::io::{Cursor, Read, Write};
use std::net::{TcpListener, TcpStream};
use crate::ws_codec::{FrameDecoder, FrameEncoder, Message};
use crate::ws_handshake::{server_handshake, HandshakeError};
pub enum ReadOutcome {
Text(String),
Binary(Vec<u8>),
Control,
Pending,
Closed,
}
pub struct ReplayStream {
pub stream: TcpStream,
pub replay: Cursor<Vec<u8>>,
}
impl ReplayStream {
pub fn new(stream: TcpStream, peeked: Vec<u8>) -> Self {
ReplayStream {
stream,
replay: Cursor::new(peeked),
}
}
}
impl Read for ReplayStream {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
if self.replay.position() < self.replay.get_ref().len() as u64 {
return self.replay.read(buf);
}
self.stream.read(buf)
}
}
impl Write for ReplayStream {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.stream.write(buf)
}
fn flush(&mut self) -> std::io::Result<()> {
self.stream.flush()
}
}
pub struct WsServerConnection {
pub stream: ReplayStream,
pub decoder: FrameDecoder,
pub encoder: FrameEncoder,
}
impl WsServerConnection {
pub fn accept(stream: TcpStream, peeked: Vec<u8>) -> Result<Self, HandshakeError> {
let mut replay = ReplayStream::new(stream, peeked);
server_handshake(&mut replay)?;
Ok(Self {
stream: replay,
decoder: FrameDecoder::new(),
encoder: FrameEncoder::new(),
})
}
pub fn write_text(&mut self, data: &str) -> Result<(), ()> {
let frame = self.encoder.encode_text(data);
self.stream.write_all(frame).map_err(|_| ())?;
self.stream.flush().map_err(|_| ())
}
pub fn read(&mut self) -> ReadOutcome {
let header = match self.decoder.decode_frame(&mut self.stream) {
Ok(Some(h)) => h,
Ok(None) => return ReadOutcome::Pending,
Err(ref e)
if e.kind() == std::io::ErrorKind::WouldBlock
|| e.kind() == std::io::ErrorKind::TimedOut =>
{
return ReadOutcome::Pending;
}
Err(_) => return ReadOutcome::Closed,
};
let payload = self.decoder.take_payload(&header);
match Message::from_frame(header.opcode, payload) {
Message::Text(text) => ReadOutcome::Text(text),
Message::Binary(data) => ReadOutcome::Binary(data),
Message::Ping(_) | Message::Pong(_) => ReadOutcome::Control,
Message::Close(_, _) => ReadOutcome::Closed,
}
}
}
pub fn read_message(conn: &mut WsServerConnection) -> Result<Option<String>, ()> {
loop {
match conn.read() {
ReadOutcome::Text(text) => return Ok(Some(text)),
ReadOutcome::Binary(data) => {
return Ok(Some(String::from_utf8_lossy(&data).into_owned()))
}
ReadOutcome::Control | ReadOutcome::Pending => continue,
ReadOutcome::Closed => return Ok(None),
}
}
}
pub fn bind_nonblocking(addr: &str, port: u16) -> std::io::Result<TcpListener> {
let listener = TcpListener::bind((addr, port))?;
listener.set_nonblocking(true)?;
Ok(listener)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ws_codec::{apply_mask, Opcode};
use std::io::Cursor;
use std::thread;
#[test]
fn read_message_text_frame() {
let mut frame = Vec::new();
frame.push(0x81); frame.push(0x05); frame.extend_from_slice(b"hello");
let _stream = TcpListener::bind("127.0.0.1:0").unwrap();
let mut decoder = FrameDecoder::new();
let mut cursor = Cursor::new(frame.clone());
let header = decoder.decode_frame(&mut cursor).unwrap().unwrap();
assert_eq!(header.opcode, Opcode::Text);
let payload = decoder.take_payload(&header);
assert_eq!(payload, b"hello");
}
#[test]
fn read_message_control_ping_returns_retry() {
let mut frame = Vec::new();
frame.push(0x89); frame.push(0x00);
let mut decoder = FrameDecoder::new();
let mut cursor = Cursor::new(frame);
let header = decoder.decode_frame(&mut cursor).unwrap().unwrap();
let payload = decoder.take_payload(&header);
let msg = Message::from_frame(header.opcode, payload);
assert!(matches!(msg, Message::Ping(_)));
}
#[test]
fn end_to_end_server_pushes_text_client_decodes() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let payload_text = "{\"method\":\"X\",\"params\":{}}";
let server_handle = thread::spawn(move || {
let (mut stream, _) = listener.accept().unwrap();
stream.set_nonblocking(false).ok();
let mut buf = [0u8; 8192];
let n = stream.read(&mut buf).unwrap();
let peeked = buf[..n].to_vec();
let mut conn = WsServerConnection::accept(stream, peeked).unwrap();
conn.write_text(payload_text).unwrap();
thread::sleep(Duration::from_millis(50));
});
let mut client = TcpStream::connect(addr).unwrap();
let key = "dGhlIHNhbXBsZSBub25jZQ==";
let req = format!(
"GET / HTTP/1.1\r\nHost: t\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\
Sec-WebSocket-Key: {}\r\nSec-WebSocket-Version: 13\r\n\r\n",
key
);
client.write_all(req.as_bytes()).unwrap();
client.flush().unwrap();
let mut got = Vec::new();
let mut byte = [0u8; 1];
while client.read(&mut byte).unwrap() > 0 {
got.push(byte[0]);
if got.ends_with(b"\r\n\r\n") {
break;
}
}
let mut decoder = FrameDecoder::new();
let header = decoder.decode_frame(&mut client).unwrap().unwrap();
assert_eq!(header.opcode, Opcode::Text);
let payload = decoder.take_payload(&header);
assert_eq!(std::str::from_utf8(&payload).unwrap(), payload_text);
server_handle.join().unwrap();
}
use std::time::Duration;
#[test]
fn masked_client_frame_decoded_by_server_decoder() {
let mut frame = Vec::new();
frame.push(0x81); let key = [0x42u8, 0x11, 0xee, 0x07];
let payload = b"hi";
frame.push(0x80 | payload.len() as u8); frame.extend_from_slice(&key);
let mut masked = payload.to_vec();
apply_mask(&mut masked, &key);
frame.extend_from_slice(&masked);
let mut decoder = FrameDecoder::new();
let mut cursor = Cursor::new(frame);
let header = decoder.decode_frame(&mut cursor).unwrap().unwrap();
assert!(header.mask);
let mask_key = decoder.take_mask();
let mut payload = decoder.take_payload(&header);
apply_mask(&mut payload, &mask_key);
assert_eq!(payload, b"hi");
}
}