From 37717ea2e1b2d29b2ae0b62efbe725b9bf668ea0 Mon Sep 17 00:00:00 2001 From: tonytown Date: Sat, 3 Oct 2026 00:04:20 +0800 Subject: [PATCH] Migrate line protocol and nickname support Port the stable text protocol from feature to main: - line-based client commands MSG and NICK - server responses SYS, MSG, NICK, YOU and ERR - access-token authentication on connect - server-side nickname validation (length, charset, case-insensitive uniqueness) - per-client rate limiting and temporary IP bans - safe rejection of malformed input (empty messages, unknown commands, oversized lines) - parser/validation tests and PROTOCOL.md documentation File transfer, remote execution, LLM and chat-history changes are intentionally excluded. --- PROTOCOL.md | 63 ++++++ README.md | 10 +- src/client.rs | 108 ++++++++-- src/server.rs | 578 +++++++++++++++++++++++++++++++------------------- 4 files changed, 530 insertions(+), 229 deletions(-) create mode 100644 PROTOCOL.md diff --git a/PROTOCOL.md b/PROTOCOL.md new file mode 100644 index 0000000..46a5ee0 --- /dev/null +++ b/PROTOCOL.md @@ -0,0 +1,63 @@ +# HAT line protocol + +The server speaks a small, line-based text protocol over TCP. Every frame is a +single line terminated by `\n` (a trailing `\r` is accepted and stripped). + +## Connecting and authentication + +1. The client opens a TCP connection. +2. The server sends `SYS authenticating`. +3. The client sends the access token printed by the server as its first line. +4. On success the server sends `SYS connected; type /help to see available commands`. + On failure it sends `ERR invalid access token` and closes the connection. + +The token line is compared exactly (after trailing newline removal). A wrong or +missing token ends the connection. + +## Client to server + +| Command | Meaning | +| --- | --- | +| `MSG ` | Send a chat message. `` must not be empty. | +| `NICK ` | Request a nickname change. | + +Any other line is answered with `ERR unknown protocol command`. Blank lines are +ignored. Lines longer than 4 KiB are rejected with `ERR message is too long`. + +### Nickname rules + +- 1 to 16 characters. +- ASCII letters, digits, `_` and `-` only. +- Case-insensitive uniqueness: a nickname in use by another client is rejected. + +On success the server replies with `YOU ` and tells everyone else +`NICK `. On failure it replies with `ERR `. + +## Server to client + +| Message | Meaning | +| --- | --- | +| `SYS ` | Informational/server event (join, leave, welcome). | +| `MSG ` | A chat message from ``. | +| `NICK ` | Another client changed nickname. | +| `YOU ` | The recipient's own (possibly new) nickname. | +| `ERR ` | A rejected request or rate-limit notice. | + +Nicknames are assigned as `user-` on connect until changed. + +## Rate limiting and bans + +- At most one accepted message per 250 ms per client. +- Messages sent faster get `ERR you are sending messages too quickly` and count + as a violation. The violation counter resets after an accepted message. +- Reaching 10 violations bans the client's IP for 10 minutes with + `ERR too many rapid messages; you are banned for 10 minutes`, then the + connection is closed. +- A connection from a banned IP receives + `ERR temporarily banned; try again in seconds` and is closed. + +## Out of scope + +This document covers only the stable text protocol migrated to `main`. File +transfer (`PUT`/`GET`/`LS`), remote execution (`EXEC`) and LLM (`LLM`) commands +are not part of this protocol on `main`. diff --git a/README.md b/README.md index 65c9cf6..0602a26 100644 --- a/README.md +++ b/README.md @@ -22,7 +22,15 @@ Start the client in another terminal: cargo run --bin client ``` -The server prints an access token when it starts. Use that token with `/connect` in the client. +The server prints an access token when it starts. In the client, connect with +the token as the third argument: + +```text +/connect 127.0.0.1 6969 +``` + +Use `/nickname ` to change your nickname. See [PROTOCOL.md](PROTOCOL.md) +for the full line protocol. ## License MIT diff --git a/src/client.rs b/src/client.rs index c376516..910ebc5 100644 --- a/src/client.rs +++ b/src/client.rs @@ -56,6 +56,21 @@ fn cmd_disconnect(ctx: &mut Ctx, _args: &[&str]) { } } +fn cmd_nickname(ctx: &mut Ctx, args: &[&str]) { + if args.len() != 1 { + ctx.message("usage: /nickname "); + return; + } + match ctx.stream.as_mut() { + Some(stream) => { + if let Err(error) = stream.write_all(format!("NICK {}\n", args[0]).as_bytes()) { + ctx.message(error.to_string()); + } + } + None => ctx.message("not connected"), + } +} + const COMMANDS: &[Command] = &[ Command { name: "help", @@ -77,6 +92,11 @@ const COMMANDS: &[Command] = &[ description: "connect to server", run: cmd_connect, }, + Command { + name: "nickname", + description: "change your nickname", + run: cmd_nickname, + }, ]; fn handle_prompt(ctx: &mut Ctx, prompt: &str) { @@ -98,7 +118,8 @@ fn handle_prompt(ctx: &mut Ctx, prompt: &str) { match ctx.stream.as_mut() { Some(stream) => { - if let Err(error) = stream.write_all(input.as_bytes()) { + // Regular input is a chat message on the line protocol. + if let Err(error) = stream.write_all(format!("MSG {input}\n").as_bytes()) { ctx.message(error.to_string()); } } @@ -106,6 +127,25 @@ fn handle_prompt(ctx: &mut Ctx, prompt: &str) { } } +// Turn one server line into a chat entry. +fn handle_server_line(ctx: &mut Ctx, line: &str) { + let (kind, rest) = line.split_once(' ').unwrap_or((line, "")); + match kind { + "MSG" => { + let (nick, text) = rest.split_once(' ').unwrap_or((rest, "")); + ctx.message(format!("<{nick}> {text}")); + } + "NICK" => { + let (old, new) = rest.split_once(' ').unwrap_or((rest, "")); + ctx.message(format!("{old} is now known as {new}")); + } + "YOU" => ctx.message(format!("you are now {rest}")), + "SYS" => ctx.message(rest), + "ERR" => ctx.message(format!("error: {rest}")), + _ => ctx.message(line.to_owned()), + } +} + fn chat_window(stdout: &mut impl Write, chat: &[String], boundary: Rect) -> io::Result<()> { let n = chat.len(); let size = n.checked_sub(boundary.h).unwrap_or(0); @@ -123,19 +163,26 @@ fn cmd_connect(ctx: &mut Ctx, args: &[&str]) { ctx.message("You already connected "); return; } - //args is ip port + // args is ip port [token] if args.len() < 2 { - ctx.message("/connect ip port"); + ctx.message("/connect [token]"); return; } let addr = format!("{}:{}", args[0], args[1]); - let stream = match TcpStream::connect(&addr) { + let mut stream = match TcpStream::connect(&addr) { Ok(stream) => stream, Err(err) => { ctx.message(format!("failed to connect to {} : {}", addr, err)); return; } }; + // The server expects the access token as the first line. + if let Some(token) = args.get(2) + && let Err(error) = stream.write_all(format!("{token}\n").as_bytes()) + { + ctx.message(error.to_string()); + return; + } if let Err(err) = stream.set_nonblocking(true) { ctx.message(format!("failed to set noblock to {} : {}", addr, err)); return; @@ -163,7 +210,8 @@ fn run_client() -> io::Result<()> { stop: false, }; let mut prompt = String::new(); - let mut buffer = [0; 64]; + let mut pending = String::new(); + let mut buffer = [0; 4096]; while !ctx.stop { while poll(Duration::ZERO).unwrap_or(false) { @@ -203,15 +251,24 @@ fn run_client() -> io::Result<()> { Ok(0) => { ctx.message("disconnected"); ctx.stream = None; + pending.clear(); } Ok(n) => match from_utf8(&buffer[..n]) { - Ok(message) => ctx.message(message), + Ok(text) => { + // The protocol is line-based, so buffer partial reads. + pending.push_str(text); + while let Some(newline) = pending.find('\n') { + let line: String = pending.drain(..=newline).collect(); + handle_server_line(&mut ctx, line.trim_end_matches(['\r', '\n'])); + } + } Err(error) => ctx.message(format!("invalid server response: {error}")), }, Err(error) if error.kind() == ErrorKind::WouldBlock => {} Err(error) => { ctx.message(error.to_string()); ctx.stream = None; + pending.clear(); } } } @@ -241,6 +298,14 @@ fn run_client() -> io::Result<()> { mod tests { use super::*; + fn ctx() -> Ctx { + Ctx { + stream: None, + chat: Vec::new(), + stop: false, + } + } + #[test] fn command_table_contains_only_minimal_commands() { assert_eq!( @@ -248,18 +313,37 @@ mod tests { .iter() .map(|command| command.name) .collect::>(), - vec!["help", "quit", "disconnect", "connect"] + vec!["help", "quit", "disconnect", "connect", "nickname"] ); } #[test] fn unknown_command_is_reported() { - let mut ctx = Ctx { - stream: None, - chat: Vec::new(), - stop: false, - }; + let mut ctx = ctx(); handle_prompt(&mut ctx, "/missing"); assert_eq!(ctx.chat, vec!["unknown command: /missing"]); } + + #[test] + fn server_message_is_rendered_with_nickname() { + let mut ctx = ctx(); + handle_server_line(&mut ctx, "MSG alice hello there"); + assert_eq!(ctx.chat, vec![" hello there"]); + } + + #[test] + fn server_protocol_lines_are_rendered() { + let mut ctx = ctx(); + handle_server_line(&mut ctx, "YOU user-1234"); + handle_server_line(&mut ctx, "NICK user-1234 alice"); + handle_server_line(&mut ctx, "ERR nickname is already in use"); + assert_eq!( + ctx.chat, + vec![ + "you are now user-1234", + "user-1234 is now known as alice", + "error: nickname is already in use", + ] + ); + } } diff --git a/src/server.rs b/src/server.rs index c47affd..a7f73e6 100644 --- a/src/server.rs +++ b/src/server.rs @@ -1,135 +1,236 @@ use std::collections::HashMap; -use std::fmt::Write as OtherWrite; -use std::io::{Read, Write}; -use std::net::{IpAddr, SocketAddr, TcpListener, TcpStream}; -use std::result; -use std::str::from_utf8; -use std::sync::Arc; -use std::sync::mpsc::{Receiver, Sender, channel}; +use std::fmt::Write as _; +use std::io::{self, BufRead, BufReader, Write}; +use std::net::{IpAddr, Shutdown, SocketAddr, TcpListener, TcpStream}; +use std::sync::mpsc::{self, Receiver, Sender}; use std::thread; -use std::time::{Duration, SystemTime}; -use std::{fmt, usize}; +use std::time::{Duration, Instant}; -const SAFE_MODE: bool = true; -const BAN_LIMIT: Duration = Duration::from_secs(60 * 10); -const MES_FREQ: Duration = Duration::from_secs(1); -const BAN_FREQ: u32 = 10; -const TOKEN_LEN: usize = 16; +const LISTEN_ADDRESS: &str = "0.0.0.0:6969"; +const TOKEN_BYTES: usize = 16; +const MAX_LINE_BYTES: usize = 4 * 1024; +const MAX_NICKNAME_CHARS: usize = 16; +const MIN_MESSAGE_INTERVAL: Duration = Duration::from_millis(250); +const MAX_RATE_LIMIT_VIOLATIONS: u32 = 10; +const BAN_DURATION: Duration = Duration::from_secs(10 * 60); -type Result = result::Result; +struct Client { + nickname: String, + last_message_at: Option, + rate_limit_violations: u32, + outbound: Sender, +} + +// Messages the per-connection writer thread serializes onto the socket. +enum Outbound { + Line(String), + Shutdown, +} + +enum ServerEvent { + Connected { + address: SocketAddr, + outbound: Sender, + }, + Disconnected { + address: SocketAddr, + }, + LineReceived { + address: SocketAddr, + line: String, + }, +} + +fn send_line(mut stream: &TcpStream, line: &str) -> io::Result<()> { + writeln!(stream, "{line}") +} -struct Sensitive(T); +fn send_to(client: &Client, line: &str) { + let _ = client.outbound.send(Outbound::Line(line.to_owned())); +} -impl fmt::Display for Sensitive { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - if SAFE_MODE { - writeln!(f, "[REDACTED]") - } else { - writeln!(f, "{}", self.0) +fn broadcast(clients: &HashMap, except: SocketAddr, line: &str) { + for (address, client) in clients { + if *address != except { + send_to(client, line); } } } -struct Client { - conn: Arc, - last_message: SystemTime, - strike_count: u32, +// Validate a requested nickname. Returns a human-readable reason on failure. +fn nickname_error(name: &str) -> Option<&'static str> { + if name.is_empty() { + return Some("nickname cannot be empty"); + } + if name.chars().count() > MAX_NICKNAME_CHARS { + return Some("nickname is too long (maximum: 16 characters)"); + } + if !name + .chars() + .all(|character| character.is_ascii_alphanumeric() || matches!(character, '_' | '-')) + { + return Some("nickname may only contain letters, numbers, '_' and '-'"); + } + None } -enum Messages { - ClientConnected { - author: Arc, - }, - ClientDisconnected { - author_addr: SocketAddr, - }, - NewMessage { - author_addr: SocketAddr, - bytes: Vec, - }, +// A parsed client line. Unknown commands are reported so the caller can send +// an ERR without crashing the connection. +#[derive(Debug, PartialEq, Eq)] +enum ClientCommand<'a> { + Message(&'a str), + Nick(&'a str), + Unknown(&'a str), } -fn server(messages_receiver: Receiver) -> Result<()> { +// Parse a single client line. Blank lines are ignored. +fn parse_client_line(line: &str) -> Option> { + let line = line.trim(); + if line.is_empty() { + return None; + } + let (command, payload) = line.split_once(' ').unwrap_or((line, "")); + match command { + "MSG" => Some(ClientCommand::Message(payload.trim())), + "NICK" => Some(ClientCommand::Nick(payload.trim())), + other => Some(ClientCommand::Unknown(other)), + } +} + +fn remove_client(clients: &mut HashMap, address: SocketAddr) { + if let Some(client) = clients.remove(&address) { + println!("{address} ({}) disconnected", client.nickname); + broadcast( + clients, + address, + &format!("SYS {} left the chat", client.nickname), + ); + } +} + +fn run_server(events: Receiver) { let mut clients = HashMap::::new(); - let mut banned_mfs = HashMap::::new(); - loop { - let message = messages_receiver.recv().expect("Receiver is hang up !!"); - match message { - Messages::ClientConnected { author } => { - let author_address = author.peer_addr().expect("TODO: cache the ip addresss"); - let now = SystemTime::now(); - let banned_at = banned_mfs.remove(&author_address.ip()); - - banned_at.and_then(|banned_at| { - let diff = now - .duration_since(banned_at) - .expect("TODO: deal with the time"); - - if diff >= BAN_LIMIT { - None + let mut banned_until = HashMap::::new(); + + while let Ok(event) = events.recv() { + match event { + ServerEvent::Connected { address, outbound } => { + let now = Instant::now(); + banned_until.retain(|_, deadline| *deadline > now); + + if let Some(deadline) = banned_until.get(&address.ip()) { + let seconds = deadline.saturating_duration_since(now).as_secs() + 1; + let _ = outbound.send(Outbound::Line(format!( + "ERR temporarily banned; try again in {seconds} seconds" + ))); + let _ = outbound.send(Outbound::Shutdown); + continue; + } + + let nickname = format!("user-{}", address.port()); + let _ = outbound.send(Outbound::Line(format!("YOU {nickname}"))); + broadcast( + &clients, + address, + &format!("SYS {nickname} joined the chat"), + ); + println!("{address} connected as {nickname}"); + + clients.insert( + address, + Client { + nickname, + last_message_at: None, + rate_limit_violations: 0, + outbound, + }, + ); + } + ServerEvent::Disconnected { address } => remove_client(&mut clients, address), + ServerEvent::LineReceived { address, line } => { + let now = Instant::now(); + let (too_quick, should_ban) = { + let Some(client) = clients.get_mut(&address) else { + continue; + }; + let too_quick = client + .last_message_at + .is_some_and(|last| now.duration_since(last) < MIN_MESSAGE_INTERVAL); + client.last_message_at = Some(now); + if too_quick { + client.rate_limit_violations += 1; + send_to(client, "ERR you are sending messages too quickly"); + ( + true, + client.rate_limit_violations >= MAX_RATE_LIMIT_VIOLATIONS, + ) } else { - Some(banned_at) + client.rate_limit_violations = 0; + (false, false) } - }); + }; - if let Some(banned_at) = banned_at { - let diff = now - .duration_since(banned_at) - .expect("TODO: deal with the time"); - banned_mfs.insert(author_address.ip(), banned_at); - let mut author = author.as_ref(); - let _ = writeln!( - author, - "You are banned MFS, time left {} seconds ", - (BAN_LIMIT - diff).as_secs_f32() - ); - let _ = author.shutdown(std::net::Shutdown::Both); - } else { - clients.insert( - author_address, - Client { - conn: author.clone(), - last_message: now, - strike_count: 0, - }, - ); + if should_ban { + banned_until.insert(address.ip(), now + BAN_DURATION); + if let Some(client) = clients.get(&address) { + send_to( + client, + "ERR too many rapid messages; you are banned for 10 minutes", + ); + let _ = client.outbound.send(Outbound::Shutdown); + } + remove_client(&mut clients, address); + continue; } - } - Messages::ClientDisconnected { author_addr } => { - clients.remove(&author_addr); - } - Messages::NewMessage { author_addr, bytes } => { - let now = SystemTime::now(); - if let Some(author) = clients.get_mut(&author_addr) { - let freq = now.duration_since(author.last_message).expect("TIME STUFF"); - // Ban rules: utf8 String - if let Ok(_text) = from_utf8(&bytes) { - println!("author {} send {:?}", Sensitive(author_addr), bytes); - // Banned Rules: freq - if freq > MES_FREQ { - author.last_message = now; - for (addr, client) in clients.iter() { - if *addr != author_addr { - let _ = client.conn.as_ref().write(&bytes); - } + if too_quick { + continue; + } + + match parse_client_line(&line) { + None => {} + Some(ClientCommand::Message(text)) => { + if text.is_empty() { + if let Some(client) = clients.get(&address) { + send_to(client, "ERR message cannot be empty"); } - } else { - author.strike_count += 1; - if author.strike_count >= BAN_FREQ { - println!("author {author_addr} was banned"); - banned_mfs.insert(author_addr.ip(), now); - let _ = write!(author.conn.as_ref(), "You are banned MFs"); - let _ = author.conn.shutdown(std::net::Shutdown::Both); + continue; + } + let nickname = clients[&address].nickname.clone(); + let message = format!("MSG {nickname} {text}"); + println!("{address} ({nickname}): {text}"); + broadcast(&clients, address, &message); + } + Some(ClientCommand::Nick(requested)) => { + let requested = requested.to_owned(); + let taken = clients.iter().any(|(other_address, other)| { + *other_address != address + && other.nickname.eq_ignore_ascii_case(&requested) + }); + let error = nickname_error(&requested) + .or(taken.then_some("nickname is already in use")); + + if let Some(reason) = error { + if let Some(client) = clients.get(&address) { + send_to(client, &format!("ERR {reason}")); } + continue; } - } else { - author.strike_count += 1; - if author.strike_count >= BAN_FREQ { - println!("author {author_addr} was banned"); - banned_mfs.insert(author_addr.ip(), now); - let _ = write!(author.conn.as_ref(), "You are banned MFs"); - let _ = author.conn.shutdown(std::net::Shutdown::Both); + + let old_nickname = clients[&address].nickname.clone(); + clients.get_mut(&address).unwrap().nickname = requested.clone(); + if let Some(client) = clients.get(&address) { + send_to(client, &format!("YOU {requested}")); + } + broadcast( + &clients, + address, + &format!("NICK {old_nickname} {requested}"), + ); + } + Some(ClientCommand::Unknown(_)) => { + if let Some(client) = clients.get(&address) { + send_to(client, "ERR unknown protocol command"); } } } @@ -138,136 +239,181 @@ fn server(messages_receiver: Receiver) -> Result<()> { } } -fn authorize(stream: &Arc, author_address: &SocketAddr, token: &String) -> Result<()> { - let _ = write!(stream.as_ref(), "token: ").map_err(|err| { - eprintln!("ERROR: Could not passing message to {author_address}: {err}"); - }); - let mut buffer = [0u8; TOKEN_LEN * 2]; - let mut filled = 0; - while filled < buffer.len() { - let n = stream.as_ref().read(&mut buffer[filled..]).map_err(|err| { - eprintln!("ERROR: Could not read message from {author_address}: {err}"); - })?; - if n == 0 { - eprintln!("ERROR: Client {author_address} disconnected during authorization"); - return Err(()); - } - filled += n; +fn handle_connection( + stream: TcpStream, + address: SocketAddr, + token: &str, + events: Sender, +) -> io::Result<()> { + stream.set_read_timeout(Some(Duration::from_secs(30)))?; + send_line(&stream, "SYS authenticating")?; + + let mut reader = BufReader::new(stream.try_clone()?); + let mut supplied_token = String::new(); + reader.read_line(&mut supplied_token)?; + if supplied_token.trim_end_matches(['\r', '\n']) != token { + send_line(&stream, "ERR invalid access token")?; + stream.shutdown(Shutdown::Both)?; + return Ok(()); } - let buffer = from_utf8(&buffer).map_err(|err| { - eprintln!("ERROR: illeagel utf8 token : {err}"); - })?; + stream.set_read_timeout(None)?; + send_line( + &stream, + "SYS connected; type /help to see available commands", + )?; - println!("get buffer {buffer}"); + let (outbound_tx, outbound_rx) = mpsc::channel::(); + events + .send(ServerEvent::Connected { + address, + outbound: outbound_tx.clone(), + }) + .map_err(|_| io::Error::new(io::ErrorKind::BrokenPipe, "server stopped"))?; + + let writer_stream = stream.try_clone()?; + thread::spawn(move || { + for outbound in outbound_rx { + let result = match outbound { + Outbound::Line(line) => send_line(&writer_stream, &line), + Outbound::Shutdown => { + let _ = writer_stream.shutdown(Shutdown::Both); + break; + } + }; + if result.is_err() { + break; + } + } + }); + drop(outbound_tx); - if token != buffer { - eprintln!("ERROR: Token not valid"); - return Err(()); + loop { + let mut line = String::new(); + match reader.read_line(&mut line) { + Ok(0) => break, + Ok(_) if line.len() > MAX_LINE_BYTES => { + send_line(&stream, "ERR message is too long")?; + } + Ok(_) => { + let line = line.trim_end_matches(['\r', '\n']); + if line.is_empty() { + continue; + } + if events + .send(ServerEvent::LineReceived { + address, + line: line.to_owned(), + }) + .is_err() + { + break; + } + } + Err(error) => { + eprintln!("Read error from {address}: {error}"); + break; + } + } } + let _ = events.send(ServerEvent::Disconnected { address }); Ok(()) } -fn client(stream: Arc, message_sender: Sender, token: String) -> Result<()> { - let author_address = stream - .peer_addr() - .map_err(|err| eprintln!("can not get the peer adderess: {err}"))?; - - authorize(&stream, &author_address, &token).map_err(|()| { - let _ = write!(stream.as_ref(), "Invalid Token ! ").map_err(|err| { - eprintln!("ERROR: Invalid Token : {err} "); - }); - let _ = stream.shutdown(std::net::Shutdown::Both).map_err(|err| { - eprintln!("ERROR: Could not shutdown the connnection : {err} "); - }); - })?; - - let _ = writeln!(stream.as_ref(), "Welcome to fight club buudy!!").map_err(|err| { - eprintln!("ERROR: failed to send message to {author_address}: {err} "); - }); +fn generate_token() -> io::Result { + let mut random_bytes = [0; TOKEN_BYTES]; + getrandom::fill(&mut random_bytes).map_err(io::Error::other)?; - message_sender - .send(Messages::ClientConnected { - author: stream.clone(), - }) - .map_err(|err| { - eprintln!("Could not to connected to client : {err}"); - })?; + let mut token = String::with_capacity(TOKEN_BYTES * 2); + for byte in random_bytes { + write!(token, "{byte:02X}").expect("writing to a String cannot fail"); + } + Ok(token) +} - let mut buffer = Vec::new(); - buffer.resize(64, 0); - loop { - let n = stream.as_ref().read(&mut buffer).map_err(|err| { - eprintln!("Disconnected to client : {err}"); - let _ = message_sender.send(Messages::ClientDisconnected { - author_addr: author_address, - }); - })?; - if n > 0 { - message_sender - .send(Messages::NewMessage { - author_addr: author_address, - bytes: buffer[0..n].to_vec(), - }) - .map_err(|err| { - eprintln!("Disconnected to client : {err}"); - })?; - } else { - let _ = message_sender - .send(Messages::ClientDisconnected { - author_addr: author_address, - }) - .map_err(|err| { - eprintln!("Could not send message to client : {err}"); +fn main() -> io::Result<()> { + let token = generate_token()?; + let listener = TcpListener::bind(LISTEN_ADDRESS)?; + println!("Chat server listening on {LISTEN_ADDRESS}"); + println!("Access token: {token}"); + + let (event_sender, event_receiver) = mpsc::channel(); + thread::spawn(move || run_server(event_receiver)); + + for incoming in listener.incoming() { + match incoming { + Ok(stream) => { + let address = stream.peer_addr()?; + let events = event_sender.clone(); + let token = token.clone(); + thread::spawn(move || { + if let Err(error) = handle_connection(stream, address, &token, events) { + eprintln!("Connection error for {address}: {error}"); + } }); - break; + } + Err(error) => eprintln!("Could not accept connection: {error}"), } } + Ok(()) } -fn main() -> Result<()> { - let mut buffer = [0; TOKEN_LEN]; - let _ = getrandom::fill(&mut buffer).map_err(|err| { - eprintln!("ERROR: Generate token failed : {err}"); - }); - - let mut token = String::new(); +#[cfg(test)] +mod tests { + use super::*; - for x in buffer.iter() { - let _ = write!(token, "{x:02X}"); + #[test] + fn parses_message_command() { + assert_eq!( + parse_client_line("MSG hello world"), + Some(ClientCommand::Message("hello world")) + ); + assert_eq!(parse_client_line("MSG"), Some(ClientCommand::Message(""))); + assert_eq!( + parse_client_line("MSG padded "), + Some(ClientCommand::Message("padded")) + ); } - println!("token is {token}"); - - let address = "0.0.0.0:6969"; - let listener = TcpListener::bind(address).map_err(|err| { - eprintln!( - "can not bind to {} : {}", - Sensitive(address), - Sensitive(err) - ) - })?; + #[test] + fn parses_nick_command() { + assert_eq!( + parse_client_line("NICK alice"), + Some(ClientCommand::Nick("alice")) + ); + assert_eq!(parse_client_line("NICK"), Some(ClientCommand::Nick(""))); + } - println!("INFO: Listening in {}", Sensitive(address)); + #[test] + fn blank_lines_are_ignored() { + assert_eq!(parse_client_line(""), None); + assert_eq!(parse_client_line(" "), None); + } - let (message_sender, message_receiver) = channel(); + #[test] + fn unknown_commands_are_rejected() { + assert_eq!( + parse_client_line("PUT file"), + Some(ClientCommand::Unknown("PUT")) + ); + assert_eq!( + parse_client_line("hello"), + Some(ClientCommand::Unknown("hello")) + ); + } - thread::spawn(|| server(message_receiver)); + #[test] + fn validates_nicknames() { + assert!(nickname_error("bob").is_none()); + assert!(nickname_error("a-b_c1").is_none()); + assert!(nickname_error(&"a".repeat(MAX_NICKNAME_CHARS)).is_none()); - for stream in listener.incoming() { - match stream { - Ok(stream) => { - let stream = Arc::new(stream); - let sender = message_sender.clone(); - let token = token.clone(); - thread::spawn(|| client(stream, sender, token)); - } - Err(e) => { - eprintln!("disconnect with user : {}", Sensitive(e)); - } - } + assert!(nickname_error("").is_some()); + assert!(nickname_error(&"a".repeat(MAX_NICKNAME_CHARS + 1)).is_some()); + assert!(nickname_error("bad name").is_some()); + assert!(nickname_error("böb").is_some()); + assert!(nickname_error("with\nnewline").is_some()); } - Ok(()) }