use { std::{collections::HashMap, env, sync::Arc}, tokio::{ io::{AsyncBufReadExt, AsyncWriteExt, BufStream}, net::{TcpListener, TcpStream}, task::spawn_blocking, }, }; enum ResponseCode { BadRequest, Forbidden, NotFound, } /// Remove query parameters and trailing slashes. /// Makes all paths lowercase. fn normalise_path(path: &str) -> String { let mut path = path .split('?') .next() .expect("split operation returns at least one substring"); if path.len() > 1 && let Some('/') = path.chars().next_back() { path = &path[0..path.len() - 1]; } path.to_lowercase() } async fn simple_response(mut socket: BufStream, rc: ResponseCode) { let rc = match rc { ResponseCode::BadRequest => "400 Bad Request", ResponseCode::Forbidden => "403 Forbidden", ResponseCode::NotFound => "404 Not Found", }; let content_length = rc.len(); let response = format!( "HTTP/1.1 {rc}\r\nContent-Type: text/plaintext; charset=utf-8\r\nContent-Length: {content_length}\r\nX-Served-By: urls-txt\r\n\r\n{rc}" ); if let Err(e) = socket.write(response.as_bytes()).await { println!("process_request: failed: error writing response: {e}"); }; if let Err(e) = socket.flush().await { println!("process_request: failed: error flushing socket: {e}"); }; } async fn process_socket(mut socket: BufStream, urlmap: Arc>) { let mut request_line = String::default(); if let Err(e) = socket.read_line(&mut request_line).await { println!("process_request: failed to read line: {e}"); return; } let mut request_line = request_line.split(' '); let (Some(http_method), Some(path)) = (request_line.next(), request_line.next()) else { println!("process_request: malformed request"); simple_response(socket, ResponseCode::BadRequest).await; return; }; let path = normalise_path(path); if http_method != "GET" { println!("process_request: forbidden method: http_method={http_method}"); simple_response(socket, ResponseCode::Forbidden).await; return; } let Some(location) = urlmap.get(&path) else { println!("process_request: not found: location={path}"); simple_response(socket, ResponseCode::NotFound).await; return; }; let content = format!("Found\r\n"); let content_length = content.len(); let response = format!( "HTTP/1.1 302 Found\r\nContent-Type: text/html; charset=utf-8\r\nContent-Length: {content_length}\r\nLocation: {location}\r\nX-Served-By: urls-txt\r\n\r\n{content}" ); if let Err(e) = socket.write(response.as_bytes()).await { println!("process_request: failed: error writing response: {e}"); return; }; if let Err(e) = socket.flush().await { println!("process_request: failed: error flushing socket: {e}"); }; } #[tokio::main] async fn main() -> Result<(), Box> { let args: Vec = env::args().collect(); if args.len() != 2 { println!("USAGE: {} ", args[0]); Err("Incorrect arguments")? } let urls_filename = args[1].clone(); let bind_address = match env::var("URLS_TXT_BIND_ADDRESS") { Ok(addr) => addr, Err(env::VarError::NotPresent) => "127.0.0.1:3000".to_string(), Err(e) => { println!("Failed to read environment variable URLS_TXT_BIND_ADDRESS: {e}"); Err(e)? } }; let config = spawn_blocking(move || std::fs::read_to_string(urls_filename)).await??; let mut urlmap = HashMap::new(); for url_pair in config.split('\n') { if url_pair.is_empty() { continue; } if url_pair .chars() .next() .expect("String at least 1 character") != '/' { continue; } let mut split = url_pair.split_whitespace(); let (Some(slug), Some(redirect)) = (split.next(), split.next()) else { Err(format!("Unable to parse line: '{}'", url_pair))? }; urlmap.insert(normalise_path(slug), redirect.to_string()); } let urlmap = Arc::new(urlmap); println!("Starting TCP Listener..."); let listener = TcpListener::bind(&bind_address).await?; loop { let (socket, _addr) = listener.accept().await?; let urlmap = urlmap.clone(); tokio::spawn(async move { process_socket(BufStream::new(socket), urlmap).await; }); } }