use std::io::{Error, ErrorKind, Write};
use std::net::SocketAddr;
use anyhow::Result;
use bytes::{Buf, BytesMut};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, BufReader};
use crate::request::form::Form;
use crate::request::Request;
use crate::response::Response;
use crate::server::protocol::TcpHandler;
use crate::server::Server;
use crate::server::protocol::tcp::http1::ws::Ws;
use crate::utils::http::Headers;
use crate::utils::mem::Instance;
use crate::utils::url::parse_query;
use crate::utils::Values;
pub mod ws;
const MAX_HEADER_LENGTH: usize = 8192;
pub struct Http1 {
server: Instance<Server>,
addr: SocketAddr,
}
impl TcpHandler for Http1 {
fn new(server: Instance<Server>, addr: SocketAddr) -> Self {
Self { server, addr }
}
async fn handle<RW>(&mut self, mut rw: BufReader<RW>) -> Result<()>
where
RW: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static,
{
let req = match self.deserialize(&mut rw).await {
Ok(req) => req,
Err(_) => return Ok(()),
};
if req.header("upgrade").eq_ignore_ascii_case("websocket") {
return Ws::new(self.server.clone(), self.addr).handle(rw, req).await;
}
let (_, res) = self.server.as_mut().on_http(req, Response::new()).await;
let serialized_res = Self::serialize(&res);
rw.write_all(&serialized_res).await?;
rw.flush().await?;
Ok(())
}
}
impl Http1 {
async fn deserialize<RW>(&mut self, rw: &mut BufReader<RW>) -> Result<Request>
where
RW: AsyncRead + AsyncWrite + Unpin + Send + Sync,
{
let mut buffer = BytesMut::with_capacity(MAX_HEADER_LENGTH);
let header_size = loop {
let n = rw.read_buf(&mut buffer).await?;
if n == 0 {
return Err(Error::new(ErrorKind::UnexpectedEof, "connection closed").into());
}
let mut headers_ptr = [httparse::EMPTY_HEADER; 64];
let mut req = httparse::Request::new(&mut headers_ptr);
match req.parse(&buffer) {
Ok(httparse::Status::Complete(size)) => break size,
Ok(httparse::Status::Partial) => {
if buffer.len() >= MAX_HEADER_LENGTH {
return Err(
Error::new(ErrorKind::InvalidData, "HTTP header limit exceeded").into(),
);
}
}
Err(e) => return Err(Error::new(ErrorKind::InvalidData, e).into()),
}
};
let header_bytes = buffer.split_to(header_size);
let leftover_body = buffer;
let mut headers_ptr = [httparse::EMPTY_HEADER; 64];
let mut parsed_req = httparse::Request::new(&mut headers_ptr);
parsed_req
.parse(&header_bytes)
.map_err(|e| Error::new(ErrorKind::InvalidData, e))?;
let mut headers = Headers::new();
let mut content_length: u64 = 0;
let mut is_chunked = false;
for h in parsed_req.headers.iter().filter(|h| !h.name.is_empty()) {
let val_str = std::str::from_utf8(h.value).unwrap_or("").trim();
let name_lower = h.name.to_ascii_lowercase();
if name_lower == "content-length" {
content_length = val_str.parse().unwrap_or(0);
} else if name_lower == "transfer-encoding" && val_str.contains("chunked") {
is_chunked = true;
}
headers.insert(name_lower, val_str.to_string());
}
let body = if is_chunked {
self.read_chunked_body(rw, leftover_body).await?
} else {
self.read_fixed_body(rw, leftover_body, content_length)
.await?
};
let raw_url = parsed_req.path.unwrap_or("");
let (path, queries) = match raw_url.find('?') {
Some(i) => (&raw_url[..i], parse_query(&raw_url[i + 1..])),
None => (raw_url, Values::new()),
};
let host = headers.get("host").cloned().unwrap_or_default();
Ok(Request {
server: self.server.clone(),
addr: self.addr,
protocol: "HTTP/1.1".to_string(),
method: parsed_req.method.unwrap_or("GET").to_string(),
path: path.to_string(),
queries: queries,
host: host,
headers: headers,
parameters: Values::new(),
cookies: Default::default(),
session: Default::default(),
body: body.into(),
form: Form::default(),
})
}
async fn read_fixed_body<RW>(
&mut self,
rw: &mut BufReader<RW>,
mut leftover: BytesMut,
content_length: u64,
) -> Result<Vec<u8>>
where
RW: AsyncRead + Unpin + Send + Sync,
{
if content_length == 0 {
return Ok(Vec::new());
}
let cl_usize = content_length as usize;
if leftover.len() >= cl_usize {
let body_bytes = leftover.split_to(cl_usize);
return Ok(body_bytes.to_vec());
}
let mut body = Vec::with_capacity(cl_usize);
let leftover_len = leftover.len();
body.extend_from_slice(&leftover);
let remaining = cl_usize - leftover_len;
let mut limited = rw.take(remaining as u64);
limited.read_to_end(&mut body).await?;
Ok(body)
}
async fn read_chunked_body<RW>(
&mut self,
rw: &mut BufReader<RW>,
mut buf: BytesMut,
) -> Result<Vec<u8>>
where
RW: AsyncRead + Unpin + Send + Sync,
{
let mut body = Vec::new();
loop {
let line_end = loop {
if let Some(pos) = buf.windows(2).position(|w| w == b"\r\n") {
break pos;
}
if rw.read_buf(&mut buf).await? == 0 {
return Err(
Error::new(ErrorKind::UnexpectedEof, "Truncated chunk size").into(),
);
}
};
let size_bytes = buf.split_to(line_end);
buf.advance(2);
let size_str = std::str::from_utf8(&size_bytes)
.map_err(|_| Error::new(ErrorKind::InvalidData, "Invalid UTF-8 in chunk size"))?;
let hex_str = size_str.split(';').next().unwrap_or("").trim();
let chunk_size = usize::from_str_radix(hex_str, 16)
.map_err(|_| Error::new(ErrorKind::InvalidData, "Invalid hex chunk size"))?;
if chunk_size == 0 {
while buf.len() < 2 {
if rw.read_buf(&mut buf).await? == 0 {
break;
}
}
break;
}
let total_needed = chunk_size + 2;
while buf.len() < total_needed {
if rw.read_buf(&mut buf).await? == 0 {
return Err(
Error::new(ErrorKind::UnexpectedEof, "Truncated chunk body").into(),
);
}
}
body.extend_from_slice(&buf[..chunk_size]);
buf.advance(total_needed); }
Ok(body)
}
fn serialize(res: &Response) -> Vec<u8> {
let content_length = res.content.len();
let mut serialized = Vec::with_capacity(128 + (res.headers.len() * 32) + content_length);
let status_text = http::StatusCode::from_u16(res.status_code)
.map(|s| s.canonical_reason().unwrap_or("OK"))
.unwrap_or("OK");
let _ = write!(serialized, "HTTP/1.1 {} {}\r\n", res.status_code, status_text);
for (k, v) in &res.headers {
if !k.eq_ignore_ascii_case("content-length") {
let _ = write!(serialized, "{}: {}\r\n", k, v);
}
}
let _ = write!(serialized, "Content-Length: {}\r\n\r\n", content_length);
serialized.extend_from_slice(&res.content);
serialized
}
}