1use crate::error_page;
5use crate::headers::Headers;
6use crate::method::Method;
7use crate::panic;
8use crate::request::Request;
9use crate::response::Response;
10use crate::router::Router;
11use crate::status::Status;
12use crate::url;
13use rustlavel_core::{Context, Error, Result};
14use std::net::SocketAddr;
15use std::sync::Arc;
16use std::sync::atomic::{AtomicUsize, Ordering};
17use std::time::{Duration, Instant};
18use tokio::io::{AsyncReadExt, AsyncWriteExt, BufWriter};
19use tokio::net::{TcpListener, TcpStream};
20
21#[derive(Debug, Clone)]
23pub struct Limits {
24 pub max_header_bytes: usize,
25 pub max_body_bytes: usize,
26 pub keep_alive_timeout: Duration,
28 pub header_timeout: Duration,
30}
31
32impl Default for Limits {
33 fn default() -> Self {
34 Limits {
35 max_header_bytes: 64 * 1024,
36 max_body_bytes: 10 * 1024 * 1024,
37 keep_alive_timeout: Duration::from_secs(15),
38 header_timeout: Duration::from_secs(10),
39 }
40 }
41}
42
43pub struct Server {
44 router: Arc<Router>,
45 context: Context,
46 limits: Limits,
47}
48
49impl Server {
50 pub fn new(mut router: Router, context: Context) -> Self {
51 router.finalize();
52 let limits = Limits {
53 max_body_bytes: context.config().int("server.max_body_bytes", 10 * 1024 * 1024) as usize,
54 ..Limits::default()
55 };
56 Server { router: Arc::new(router), context, limits }
57 }
58
59 pub fn limits(mut self, limits: Limits) -> Self {
60 self.limits = limits;
61 self
62 }
63
64 pub async fn listen(self, addr: impl Into<String>) -> Result<()> {
66 let addr = addr.into();
67 let listener = TcpListener::bind(&addr).await.map_err(Error::Io)?;
68 let local = listener.local_addr().map_err(Error::Io)?;
69
70 panic::install_hook();
71 error_page::set_debug(self.context.config().debug());
72
73 rustlavel_core::info!("Rustlavel serving on http://{local}");
74 rustlavel_core::info!("Press Ctrl-C to stop");
75
76 let in_flight = Arc::new(AtomicUsize::new(0));
77 let shared = Arc::new(self);
78
79 loop {
80 let accepted = tokio::select! {
81 result = listener.accept() => result,
82 _ = tokio::signal::ctrl_c() => break,
83 };
84
85 let (stream, peer) = match accepted {
86 Ok(pair) => pair,
87 Err(e) => {
90 rustlavel_core::warn!("accept failed: {e}");
91 continue;
92 }
93 };
94
95 let server = Arc::clone(&shared);
96 let counter = Arc::clone(&in_flight);
97 counter.fetch_add(1, Ordering::SeqCst);
98 tokio::spawn(async move {
99 if let Err(e) = server.serve_connection(stream, peer).await {
100 rustlavel_core::debug!("connection closed: {e}");
101 }
102 counter.fetch_sub(1, Ordering::SeqCst);
103 });
104 }
105
106 rustlavel_core::info!("Shutting down, waiting for in-flight requests…");
107 let deadline = Instant::now() + Duration::from_secs(10);
108 while in_flight.load(Ordering::SeqCst) > 0 && Instant::now() < deadline {
109 tokio::time::sleep(Duration::from_millis(25)).await;
110 }
111 rustlavel_core::info!("Goodbye.");
112 Ok(())
113 }
114
115 async fn serve_connection(&self, stream: TcpStream, peer: SocketAddr) -> Result<()> {
116 let _ = stream.set_nodelay(true);
118 let (mut reader, writer) = stream.into_split();
119 let mut writer = BufWriter::new(writer);
120 let mut buffer: Vec<u8> = Vec::with_capacity(2048);
121
122 loop {
123 let head = match self.read_head(&mut reader, &mut buffer).await? {
124 Some(head) => head,
125 None => return Ok(()),
127 };
128
129 let (mut request, keep_alive) = match self.parse(&head, &mut reader, &mut buffer, peer).await {
130 Ok(parsed) => parsed,
131 Err(error) => {
132 let response = Response::new(Status::BAD_REQUEST).with_text(error.to_string());
133 writer.write_all(&response.to_bytes(true)).await.map_err(Error::Io)?;
134 writer.flush().await.map_err(Error::Io)?;
135 return Ok(());
136 }
137 };
138
139 request.context = self.context.clone();
140 let is_head = request.method() == Method::Head;
141 let mut response = self.dispatch(request).await;
142
143 if let Some(upgrade) = response.take_upgrade() {
146 writer.write_all(&response.to_bytes(false)).await.map_err(Error::Io)?;
147 writer.flush().await.map_err(Error::Io)?;
148
149 let upgraded = crate::upgrade::Upgraded {
150 reader: Box::new(reader),
151 writer: Box::new(writer),
152 buffered: std::mem::take(&mut buffer),
155 };
156 upgrade.run(upgraded).await;
157 return Ok(());
158 }
159
160 if !keep_alive {
161 response.headers.set("connection", "close");
162 }
163 writer.write_all(&response.to_bytes(!is_head)).await.map_err(Error::Io)?;
164 writer.flush().await.map_err(Error::Io)?;
165
166 if !keep_alive {
167 return Ok(());
168 }
169 }
170 }
171
172 async fn dispatch(&self, request: Request) -> Response {
175 let started = Instant::now();
176 let method = request.method();
177 let path = request.path().to_string();
178
179 let response = self.router.dispatch(request).await;
182 let elapsed = started.elapsed();
183
184 if rustlavel_core::log::enabled(rustlavel_core::log::Level::Debug) {
185 rustlavel_core::debug!(
186 "{method} {path} → {} ({:.1}ms)",
187 response.status.code(),
188 elapsed.as_secs_f64() * 1000.0
189 );
190 }
191
192 response
193 }
194
195 async fn read_head(
197 &self,
198 reader: &mut tokio::net::tcp::OwnedReadHalf,
199 buffer: &mut Vec<u8>,
200 ) -> Result<Option<Vec<u8>>> {
201 let mut timeout = self.limits.keep_alive_timeout;
204
205 loop {
206 if let Some(end) = find_head_end(buffer) {
207 let head = buffer[..end].to_vec();
208 buffer.drain(..end);
209 return Ok(Some(head));
210 }
211 if buffer.len() > self.limits.max_header_bytes {
212 return Err(Error::Protocol("request headers are too large".into()));
213 }
214
215 let mut chunk = [0u8; 4096];
216 let read = match tokio::time::timeout(timeout, reader.read(&mut chunk)).await {
217 Ok(Ok(0)) if buffer.is_empty() => return Ok(None),
218 Ok(Ok(0)) => return Err(Error::Protocol("connection closed mid-request".into())),
219 Ok(Ok(n)) => n,
220 Ok(Err(e)) => return Err(Error::Io(e)),
221 Err(_) if buffer.is_empty() => return Ok(None),
222 Err(_) => return Err(Error::Protocol("timed out reading request headers".into())),
223 };
224 buffer.extend_from_slice(&chunk[..read]);
225 timeout = self.limits.header_timeout;
226 }
227 }
228
229 async fn parse(
230 &self,
231 head: &[u8],
232 reader: &mut tokio::net::tcp::OwnedReadHalf,
233 buffer: &mut Vec<u8>,
234 peer: SocketAddr,
235 ) -> Result<(Request, bool)> {
236 let text = std::str::from_utf8(head).map_err(|_| Error::Protocol("headers are not UTF-8".into()))?;
237 let mut lines = text.split("\r\n");
238
239 let request_line = lines.next().ok_or_else(|| Error::Protocol("empty request".into()))?;
240 let mut parts = request_line.split(' ');
241 let method = parts
242 .next()
243 .and_then(Method::parse)
244 .ok_or_else(|| Error::Protocol("unsupported method".into()))?;
245 let target = parts.next().ok_or_else(|| Error::Protocol("missing request target".into()))?;
246 let version = parts.next().unwrap_or("HTTP/1.1");
247
248 let mut headers = Headers::new();
249 for line in lines {
250 if line.is_empty() {
251 continue;
252 }
253 let (name, value) = line
254 .split_once(':')
255 .ok_or_else(|| Error::Protocol(format!("malformed header line: {line}")))?;
256 headers.append(name.trim(), value.trim());
257 }
258
259 let target = match target.find("://") {
261 Some(scheme_end) => match target[scheme_end + 3..].find('/') {
262 Some(path_start) => &target[scheme_end + 3 + path_start..],
263 None => "/",
264 },
265 None => target,
266 };
267
268 let body = self.read_body(&headers, reader, buffer).await?;
269
270 let keep_alive = match headers.get("connection") {
271 Some(value) if value.eq_ignore_ascii_case("close") => false,
272 Some(value) if value.eq_ignore_ascii_case("keep-alive") => true,
273 _ => version != "HTTP/1.0",
274 };
275
276 let (path, query) = url::split_target(target);
277 let mut request = Request::new(method, target);
278 request.path = url::decode(path);
279 request.query = url::parse_query(query);
280 request.headers = headers;
281 request.peer = Some(peer);
282 Ok((request.with_body(body), keep_alive))
283 }
284
285 async fn read_body(
286 &self,
287 headers: &Headers,
288 reader: &mut tokio::net::tcp::OwnedReadHalf,
289 buffer: &mut Vec<u8>,
290 ) -> Result<Vec<u8>> {
291 if headers.get("transfer-encoding").is_some_and(|te| te.contains("chunked")) {
292 return self.read_chunked_body(reader, buffer).await;
293 }
294
295 let Some(length) = headers.content_length() else {
296 return Ok(Vec::new());
297 };
298 if length > self.limits.max_body_bytes {
299 return Err(Error::Protocol("request body is too large".into()));
300 }
301
302 while buffer.len() < length {
303 let mut chunk = vec![0u8; (length - buffer.len()).min(64 * 1024)];
304 let read = tokio::time::timeout(self.limits.header_timeout, reader.read(&mut chunk))
305 .await
306 .map_err(|_| Error::Protocol("timed out reading request body".into()))?
307 .map_err(Error::Io)?;
308 if read == 0 {
309 return Err(Error::Protocol("request body ended early".into()));
310 }
311 buffer.extend_from_slice(&chunk[..read]);
312 }
313
314 Ok(buffer.drain(..length).collect())
315 }
316
317 async fn read_chunked_body(
318 &self,
319 reader: &mut tokio::net::tcp::OwnedReadHalf,
320 buffer: &mut Vec<u8>,
321 ) -> Result<Vec<u8>> {
322 let mut body = Vec::new();
323
324 loop {
325 let line_end = loop {
327 if let Some(at) = find_crlf(buffer) {
328 break at;
329 }
330 if !fill(reader, buffer, self.limits.header_timeout).await? {
331 return Err(Error::Protocol("chunked body ended early".into()));
332 }
333 };
334
335 let header: Vec<u8> = buffer.drain(..line_end + 2).collect();
336 let size_text = String::from_utf8_lossy(&header[..line_end]);
337 let size = usize::from_str_radix(size_text.split(';').next().unwrap_or("").trim(), 16)
338 .map_err(|_| Error::Protocol("invalid chunk size".into()))?;
339
340 if size == 0 {
341 loop {
344 let end = loop {
345 if let Some(at) = find_crlf(buffer) {
346 break at;
347 }
348 if !fill(reader, buffer, self.limits.header_timeout).await? {
349 return Ok(body);
350 }
351 };
352 buffer.drain(..end + 2);
353 if end == 0 {
354 return Ok(body);
355 }
356 }
357 }
358
359 if body.len() + size > self.limits.max_body_bytes {
360 return Err(Error::Protocol("request body is too large".into()));
361 }
362
363 while buffer.len() < size + 2 {
364 if !fill(reader, buffer, self.limits.header_timeout).await? {
365 return Err(Error::Protocol("chunked body ended early".into()));
366 }
367 }
368 body.extend(buffer.drain(..size));
369 buffer.drain(..2);
370 }
371 }
372}
373
374async fn fill(
375 reader: &mut tokio::net::tcp::OwnedReadHalf,
376 buffer: &mut Vec<u8>,
377 timeout: Duration,
378) -> Result<bool> {
379 let mut chunk = [0u8; 4096];
380 let read = tokio::time::timeout(timeout, reader.read(&mut chunk))
381 .await
382 .map_err(|_| Error::Protocol("timed out reading request body".into()))?
383 .map_err(Error::Io)?;
384 buffer.extend_from_slice(&chunk[..read]);
385 Ok(read > 0)
386}
387
388fn find_head_end(buffer: &[u8]) -> Option<usize> {
390 buffer.windows(4).position(|w| w == b"\r\n\r\n").map(|at| at + 4)
391}
392
393fn find_crlf(buffer: &[u8]) -> Option<usize> {
394 buffer.windows(2).position(|w| w == b"\r\n")
395}
396
397#[cfg(test)]
398mod tests {
399 use super::*;
400
401 #[test]
402 fn finds_the_end_of_a_header_block() {
403 assert_eq!(find_head_end(b"GET / HTTP/1.1\r\n\r\nbody"), Some(18));
404 assert_eq!(find_head_end(b"GET / HTTP/1.1\r\n"), None);
405 }
406
407 #[tokio::test]
408 async fn parses_a_request_with_a_body() {
409 let server = Server::new(Router::new(), Context::default());
410 let head = b"POST /users?page=2 HTTP/1.1\r\nHost: localhost\r\nContent-Type: application/json\r\nContent-Length: 14\r\n\r\n";
411
412 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
413 let addr = listener.local_addr().unwrap();
414 tokio::spawn(async move {
415 let (mut stream, _) = listener.accept().await.unwrap();
416 stream.write_all(br#"{"name":"ada"}"#).await.unwrap();
417 });
418 let stream = TcpStream::connect(addr).await.unwrap();
419 let (mut reader, _writer) = stream.into_split();
420
421 let mut buffer = Vec::new();
422 let (mut request, keep_alive) =
423 server.parse(head, &mut reader, &mut buffer, addr).await.unwrap();
424
425 assert_eq!(request.method(), Method::Post);
426 assert_eq!(request.path(), "/users");
427 assert_eq!(request.query("page"), Some("2"));
428 assert_eq!(request.header("host"), Some("localhost"));
429 assert_eq!(request.input("name").as_deref(), Some("ada"));
430 assert!(keep_alive);
431 }
432
433 #[tokio::test]
434 async fn http_1_0_closes_by_default() {
435 let server = Server::new(Router::new(), Context::default());
436 let head = b"GET / HTTP/1.0\r\nHost: localhost\r\n\r\n";
437
438 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
439 let addr = listener.local_addr().unwrap();
440 tokio::spawn(async move {
441 let _ = listener.accept().await;
442 });
443 let (mut reader, _w) = TcpStream::connect(addr).await.unwrap().into_split();
444
445 let mut buffer = Vec::new();
446 let (_request, keep_alive) =
447 server.parse(head, &mut reader, &mut buffer, addr).await.unwrap();
448
449 assert!(!keep_alive);
450 }
451
452 #[tokio::test]
453 async fn rejects_a_body_larger_than_the_limit() {
454 let mut server = Server::new(Router::new(), Context::default());
455 server.limits.max_body_bytes = 8;
456 let head = b"POST / HTTP/1.1\r\nContent-Length: 9999\r\n\r\n";
457
458 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
459 let addr = listener.local_addr().unwrap();
460 tokio::spawn(async move {
461 let _ = listener.accept().await;
462 });
463 let (mut reader, _w) = TcpStream::connect(addr).await.unwrap().into_split();
464
465 let mut buffer = Vec::new();
466 let error = server.parse(head, &mut reader, &mut buffer, addr).await.unwrap_err();
467
468 assert!(error.to_string().contains("too large"));
469 }
470}