Skip to main content

rustlavel_http/
server.rs

1//! The HTTP/1.1 server: accept loop, request parsing, keep-alive, and a
2//! graceful shutdown that lets in-flight requests finish.
3
4use 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/// Guard rails applied to every connection.
22#[derive(Debug, Clone)]
23pub struct Limits {
24    pub max_header_bytes: usize,
25    pub max_body_bytes: usize,
26    /// How long an idle keep-alive connection is held open.
27    pub keep_alive_timeout: Duration,
28    /// How long the headers of a single request may take to arrive.
29    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    /// Bind and serve until Ctrl-C, then drain in-flight requests.
65    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                // A single failed accept (fd exhaustion, a dropped SYN) should
88                // not bring down the listener.
89                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        // Small responses should leave for the client immediately.
117        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                // Client hung up between requests: a clean end, not an error.
126                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            // A handler that answered 101 wants the socket. Write the
144            // handshake, then stop speaking HTTP on this connection.
145            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                    // Anything already read past the request belongs to the new
153                    // protocol; dropping it would lose its first frame.
154                    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    /// Run the router, converting a panic into the error page instead of
173    /// letting it kill the connection task.
174    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        // Panics are caught, and the request event is dispatched, inside the
180        // router — so both behave identically under the test client.
181        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    /// Read until the end of the header block, returning the raw head.
196    async fn read_head(
197        &self,
198        reader: &mut tokio::net::tcp::OwnedReadHalf,
199        buffer: &mut Vec<u8>,
200    ) -> Result<Option<Vec<u8>>> {
201        // A connection waiting for its first byte gets the longer keep-alive
202        // budget; once bytes arrive the head must complete promptly.
203        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        // An absolute-form target (`GET http://host/path`) is legal for proxies.
260        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            // Each chunk starts with its size in hex on its own line.
326            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                // The final chunk may be followed by trailer lines; both end at
342                // a blank line.
343                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
388/// Byte offset just past the blank line that ends the header block.
389fn 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}