use bytes::{Bytes, BytesMut};
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
pub const MAX_SUPPORTED_HEADERS: usize = 100;
#[derive(Debug, Clone)]
pub struct BatchConfig {
pub buffer_size: usize,
pub max_requests: usize,
pub max_request_size: usize,
pub max_headers: usize,
}
impl Default for BatchConfig {
fn default() -> Self {
Self {
buffer_size: 65536, max_requests: 32,
max_request_size: 1048576, max_headers: 100,
}
}
}
impl BatchConfig {
pub fn builder() -> BatchConfigBuilder {
BatchConfigBuilder::default()
}
pub fn high_throughput() -> Self {
Self {
buffer_size: 131072, max_requests: 64,
max_request_size: 2097152, max_headers: 100,
}
}
pub fn low_latency() -> Self {
Self {
buffer_size: 16384, max_requests: 8,
max_request_size: 524288, max_headers: 64,
}
}
pub fn memory_efficient() -> Self {
Self {
buffer_size: 32768, max_requests: 16,
max_request_size: 524288, max_headers: 50,
}
}
}
#[derive(Debug, Clone, Default)]
pub struct BatchConfigBuilder {
config: BatchConfig,
}
impl BatchConfigBuilder {
pub fn buffer_size(mut self, size: usize) -> Self {
self.config.buffer_size = size;
self
}
pub fn max_requests(mut self, max: usize) -> Self {
self.config.max_requests = max;
self
}
pub fn max_request_size(mut self, size: usize) -> Self {
self.config.max_request_size = size;
self
}
pub fn max_headers(mut self, max: usize) -> Self {
self.config.max_headers = max.min(MAX_SUPPORTED_HEADERS);
self
}
pub fn build(self) -> BatchConfig {
self.config
}
}
#[derive(Debug)]
pub struct ParsedRequest<'a> {
pub method: &'a str,
pub path: &'a str,
pub version: &'a str,
pub headers: Vec<(&'a str, &'a [u8])>,
pub body: &'a [u8],
pub bytes_consumed: usize,
}
#[derive(Debug, Clone)]
pub struct OwnedParsedRequest {
pub method: String,
pub path: String,
pub version: String,
pub headers: Vec<(String, Vec<u8>)>,
pub body: Bytes,
pub bytes_consumed: usize,
}
impl OwnedParsedRequest {
#[inline]
pub fn header(&self, name: &str) -> Option<&[u8]> {
self.headers
.iter()
.find(|(n, _)| n.eq_ignore_ascii_case(name))
.map(|(_, v)| v.as_slice())
}
#[inline]
pub fn content_length(&self) -> Option<usize> {
self.header("content-length")
.and_then(|v| std::str::from_utf8(v).ok())
.and_then(|s| s.trim().parse().ok())
}
#[inline]
pub fn is_keep_alive(&self) -> bool {
if self.version == "1.1" {
!self
.header("connection")
.map(|v| v.eq_ignore_ascii_case(b"close"))
.unwrap_or(false)
} else {
self.header("connection")
.map(|v| v.eq_ignore_ascii_case(b"keep-alive"))
.unwrap_or(false)
}
}
pub fn to_http_request(&self) -> crate::HttpRequest {
let mut req = crate::HttpRequest::new(self.method.clone(), self.path.clone());
for (name, value) in &self.headers {
if let Ok(v) = std::str::from_utf8(value) {
req.headers.insert(name.clone(), v.to_string());
}
}
if !self.body.is_empty() {
req.set_body_bytes(self.body.clone());
}
req
}
}
#[derive(Debug)]
pub struct ParsedBatch {
pub requests: Vec<OwnedParsedRequest>,
pub bytes_consumed: usize,
pub partial: bool,
pub error: Option<BatchParseError>,
}
impl<'a> ParsedRequest<'a> {
#[inline]
pub fn header(&self, name: &str) -> Option<&[u8]> {
self.headers
.iter()
.find(|(n, _)| n.eq_ignore_ascii_case(name))
.map(|(_, v)| *v)
}
#[inline]
pub fn content_length(&self) -> Option<usize> {
self.header("content-length")
.and_then(|v| std::str::from_utf8(v).ok())
.and_then(|s| s.trim().parse().ok())
}
#[inline]
pub fn is_keep_alive(&self) -> bool {
if self.version == "1.1" {
!self
.header("connection")
.map(|v| v.eq_ignore_ascii_case(b"close"))
.unwrap_or(false)
} else {
self.header("connection")
.map(|v| v.eq_ignore_ascii_case(b"keep-alive"))
.unwrap_or(false)
}
}
pub fn to_http_request(&self) -> crate::HttpRequest {
let mut req = crate::HttpRequest::new(self.method.to_string(), self.path.to_string());
for (name, value) in &self.headers {
if let Ok(v) = std::str::from_utf8(value) {
req.headers.insert(name, v);
}
}
if !self.body.is_empty() {
req.set_body_bytes(Bytes::copy_from_slice(self.body));
}
req
}
}
#[derive(Debug)]
pub struct BatchParseResult<'a> {
pub requests: Vec<ParsedRequest<'a>>,
pub bytes_consumed: usize,
pub partial: bool,
pub error: Option<BatchParseError>,
}
#[derive(Debug, Clone)]
pub enum BatchParseError {
RequestTooLarge(usize),
TooManyHeaders(usize),
InvalidSyntax(String),
BufferOverflow,
}
impl std::fmt::Display for BatchParseError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::RequestTooLarge(size) => write!(f, "Request too large: {} bytes", size),
Self::TooManyHeaders(count) => write!(f, "Too many headers: {}", count),
Self::InvalidSyntax(msg) => write!(f, "Invalid HTTP syntax: {}", msg),
Self::BufferOverflow => write!(f, "Buffer overflow"),
}
}
}
impl std::error::Error for BatchParseError {}
pub struct BatchParser {
config: BatchConfig,
}
impl BatchParser {
pub fn new(config: BatchConfig) -> Self {
Self { config }
}
#[inline]
pub fn parse_batch<'a>(&self, buffer: &'a [u8]) -> BatchParseResult<'a> {
let mut requests = Vec::with_capacity(self.config.max_requests);
let mut offset = 0;
let mut error = None;
while offset < buffer.len() && requests.len() < self.config.max_requests {
match self.parse_single(&buffer[offset..]) {
Ok(Some((req, consumed))) => {
if consumed > self.config.max_request_size {
error = Some(BatchParseError::RequestTooLarge(consumed));
break;
}
if req.headers.len() > self.config.max_headers {
error = Some(BatchParseError::TooManyHeaders(req.headers.len()));
break;
}
requests.push(req);
offset += consumed;
}
Ok(None) => {
break;
}
Err(e) => {
error = Some(e);
break;
}
}
}
BatchParseResult {
requests,
bytes_consumed: offset,
partial: offset < buffer.len() && error.is_none(),
error,
}
}
#[inline]
fn parse_single<'a>(
&self,
buffer: &'a [u8],
) -> Result<Option<(ParsedRequest<'a>, usize)>, BatchParseError> {
if buffer.len() < 16 {
return Ok(None);
}
let mut headers = [httparse::EMPTY_HEADER; MAX_SUPPORTED_HEADERS];
let mut req = httparse::Request::new(&mut headers);
match req.parse(buffer) {
Ok(httparse::Status::Complete(header_len)) => {
let method = req
.method
.ok_or_else(|| BatchParseError::InvalidSyntax("Missing method".to_string()))?;
let path = req
.path
.ok_or_else(|| BatchParseError::InvalidSyntax("Missing path".to_string()))?;
let version = match req.version {
Some(0) => "1.0",
Some(1) => "1.1",
_ => "1.1",
};
let parsed_headers: Vec<(&str, &[u8])> = req
.headers
.iter()
.take_while(|h| !h.name.is_empty())
.map(|h| (h.name, h.value))
.collect();
if parsed_headers
.iter()
.any(|(n, _)| n.eq_ignore_ascii_case("transfer-encoding"))
{
return Err(BatchParseError::InvalidSyntax(
"Transfer-Encoding is not supported by the batch parser".to_string(),
));
}
let content_length = parsed_headers
.iter()
.find(|(n, _)| n.eq_ignore_ascii_case("content-length"))
.and_then(|(_, v)| std::str::from_utf8(v).ok())
.and_then(|s| s.trim().parse::<usize>().ok())
.unwrap_or(0);
let total_len = header_len + content_length;
if buffer.len() < total_len {
if total_len > self.config.buffer_size {
return Err(BatchParseError::BufferOverflow);
}
return Ok(None); }
let body = &buffer[header_len..total_len];
Ok(Some((
ParsedRequest {
method,
path,
version,
headers: parsed_headers,
body,
bytes_consumed: total_len,
},
total_len,
)))
}
Ok(httparse::Status::Partial) => Ok(None),
Err(e) => Err(BatchParseError::InvalidSyntax(e.to_string())),
}
}
pub fn config(&self) -> &BatchConfig {
&self.config
}
}
#[derive(Debug, Default)]
pub struct BatchStats {
batches_processed: AtomicU64,
total_requests: AtomicU64,
total_bytes_read: AtomicU64,
avg_batch_size: AtomicU64,
max_batch_size: AtomicUsize,
parse_errors: AtomicU64,
}
impl BatchStats {
pub fn new() -> Self {
Self::default()
}
#[inline]
pub fn record_batch(&self, request_count: usize, bytes_read: usize) {
self.batches_processed.fetch_add(1, Ordering::Relaxed);
self.total_requests
.fetch_add(request_count as u64, Ordering::Relaxed);
self.total_bytes_read
.fetch_add(bytes_read as u64, Ordering::Relaxed);
let current = self.avg_batch_size.load(Ordering::Relaxed);
let new_avg = (current * 95 + (request_count as u64 * 100) * 5) / 100;
self.avg_batch_size.store(new_avg, Ordering::Relaxed);
self.max_batch_size
.fetch_max(request_count, Ordering::Relaxed);
}
#[inline]
pub fn record_error(&self) {
self.parse_errors.fetch_add(1, Ordering::Relaxed);
}
#[inline]
pub fn batches_processed(&self) -> u64 {
self.batches_processed.load(Ordering::Relaxed)
}
#[inline]
pub fn total_requests(&self) -> u64 {
self.total_requests.load(Ordering::Relaxed)
}
#[inline]
pub fn total_bytes_read(&self) -> u64 {
self.total_bytes_read.load(Ordering::Relaxed)
}
#[inline]
pub fn avg_batch_size(&self) -> f64 {
self.avg_batch_size.load(Ordering::Relaxed) as f64 / 100.0
}
#[inline]
pub fn max_batch_size(&self) -> usize {
self.max_batch_size.load(Ordering::Relaxed)
}
#[inline]
pub fn parse_errors(&self) -> u64 {
self.parse_errors.load(Ordering::Relaxed)
}
}
#[derive(Debug)]
pub struct BatchBuffer {
buffer: BytesMut,
config: BatchConfig,
read_pos: usize,
write_pos: usize,
}
impl BatchBuffer {
pub fn new(config: BatchConfig) -> Self {
Self {
buffer: BytesMut::with_capacity(config.buffer_size),
config,
read_pos: 0,
write_pos: 0,
}
}
#[inline]
pub fn remaining_capacity(&self) -> usize {
self.config.buffer_size - self.write_pos
}
#[inline]
pub fn unprocessed(&self) -> &[u8] {
&self.buffer[self.read_pos..self.write_pos]
}
#[inline]
pub fn write_slice(&mut self) -> &mut [u8] {
if self.buffer.len() < self.config.buffer_size {
self.buffer.resize(self.config.buffer_size, 0);
}
&mut self.buffer[self.write_pos..self.config.buffer_size]
}
#[inline]
pub fn advance_write(&mut self, count: usize) {
self.write_pos += count;
}
#[inline]
pub fn advance_read(&mut self, count: usize) {
self.read_pos += count;
}
#[inline]
pub fn compact(&mut self) {
if self.read_pos > 0 {
let remaining = self.write_pos - self.read_pos;
if remaining > 0 {
self.buffer.copy_within(self.read_pos..self.write_pos, 0);
}
self.read_pos = 0;
self.write_pos = remaining;
}
}
#[inline]
pub fn reset(&mut self) {
self.read_pos = 0;
self.write_pos = 0;
}
#[inline]
pub fn is_empty(&self) -> bool {
self.read_pos >= self.write_pos
}
#[inline]
pub fn len(&self) -> usize {
self.write_pos - self.read_pos
}
}
pub struct BatchReader {
parser: BatchParser,
buffer: BatchBuffer,
stats: BatchStats,
}
impl BatchReader {
pub fn new(config: BatchConfig) -> Self {
Self {
parser: BatchParser::new(config.clone()),
buffer: BatchBuffer::new(config),
stats: BatchStats::new(),
}
}
pub fn config(&self) -> &BatchConfig {
self.parser.config()
}
pub fn stats(&self) -> &BatchStats {
&self.stats
}
#[inline]
pub fn read_buffer(&mut self) -> &mut [u8] {
if self.buffer.remaining_capacity() < 4096 {
self.buffer.compact();
}
self.buffer.write_slice()
}
#[inline]
pub fn data_received(&mut self, count: usize) {
self.buffer.advance_write(count);
}
#[inline]
pub fn parse_available(&mut self) -> ParsedBatch {
let data = self.buffer.unprocessed();
let result = self.parser.parse_batch(data);
let bytes_consumed = result.bytes_consumed;
let request_count = result.requests.len();
let has_error = result.error.is_some();
let partial = result.partial;
let requests: Vec<OwnedParsedRequest> = result
.requests
.into_iter()
.map(|r| OwnedParsedRequest {
method: r.method.to_string(),
path: r.path.to_string(),
version: r.version.to_string(),
headers: r
.headers
.iter()
.map(|(k, v)| (k.to_string(), v.to_vec()))
.collect(),
body: Bytes::copy_from_slice(r.body),
bytes_consumed: r.bytes_consumed,
})
.collect();
let error = result.error;
self.buffer.advance_read(bytes_consumed);
if request_count > 0 {
self.stats.record_batch(request_count, bytes_consumed);
}
if has_error {
self.stats.record_error();
}
ParsedBatch {
requests,
bytes_consumed,
partial,
error,
}
}
#[inline]
pub fn reset(&mut self) {
self.buffer.reset();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_batch_config_builder() {
let config = BatchConfig::builder()
.buffer_size(32768)
.max_requests(16)
.max_request_size(2048)
.build();
assert_eq!(config.buffer_size, 32768);
assert_eq!(config.max_requests, 16);
assert_eq!(config.max_request_size, 2048);
}
#[test]
fn test_parse_single_request() {
let parser = BatchParser::new(BatchConfig::default());
let request = b"GET /api/users HTTP/1.1\r\nHost: example.com\r\nContent-Length: 0\r\n\r\n";
let result = parser.parse_batch(request);
assert_eq!(result.requests.len(), 1);
assert_eq!(result.requests[0].method, "GET");
assert_eq!(result.requests[0].path, "/api/users");
assert_eq!(result.requests[0].version, "1.1");
assert!(!result.partial);
}
#[test]
fn test_parse_multiple_requests() {
let parser = BatchParser::new(BatchConfig::default());
let requests = b"GET / HTTP/1.1\r\nHost: a.com\r\n\r\nPOST /api HTTP/1.1\r\nHost: b.com\r\nContent-Length: 4\r\n\r\ntest";
let result = parser.parse_batch(requests);
assert_eq!(result.requests.len(), 2);
assert_eq!(result.requests[0].method, "GET");
assert_eq!(result.requests[0].path, "/");
assert_eq!(result.requests[1].method, "POST");
assert_eq!(result.requests[1].path, "/api");
assert_eq!(result.requests[1].body, Bytes::from_static(b"test"));
}
#[test]
fn test_parse_partial_request() {
let parser = BatchParser::new(BatchConfig::default());
let partial = b"GET /api HTTP/1.1\r\nHost: exam";
let result = parser.parse_batch(partial);
assert_eq!(result.requests.len(), 0);
assert!(result.partial || result.bytes_consumed == 0);
}
#[test]
fn test_parsed_request_helpers() {
let parser = BatchParser::new(BatchConfig::default());
let request = b"GET /api HTTP/1.1\r\nHost: example.com\r\nContent-Length: 5\r\nConnection: keep-alive\r\n\r\nhello";
let result = parser.parse_batch(request);
assert_eq!(result.requests.len(), 1);
let req = &result.requests[0];
assert_eq!(req.content_length(), Some(5));
assert!(req.is_keep_alive());
assert_eq!(req.header("host"), Some(b"example.com".as_slice()));
}
#[test]
fn test_batch_buffer() {
let config = BatchConfig::default();
let mut buffer = BatchBuffer::new(config);
let data = b"GET / HTTP/1.1\r\n";
buffer.write_slice()[..data.len()].copy_from_slice(data);
buffer.advance_write(data.len());
assert_eq!(buffer.len(), data.len());
assert_eq!(buffer.unprocessed(), data);
buffer.advance_read(4);
assert_eq!(buffer.unprocessed(), b"/ HTTP/1.1\r\n");
buffer.compact();
assert_eq!(buffer.len(), 12);
}
#[test]
fn test_batch_reader() {
let config = BatchConfig::default();
let mut reader = BatchReader::new(config);
let request = b"GET /test HTTP/1.1\r\nHost: localhost\r\n\r\n";
reader.read_buffer()[..request.len()].copy_from_slice(request);
reader.data_received(request.len());
let result = reader.parse_available();
assert_eq!(result.requests.len(), 1);
assert_eq!(result.requests[0].path, "/test");
assert_eq!(reader.stats().total_requests(), 1);
}
#[test]
fn test_transfer_encoding_rejected() {
let parser = BatchParser::new(BatchConfig::default());
let request = b"POST /api HTTP/1.1\r\nHost: a.com\r\nTransfer-Encoding: chunked\r\n\r\n5\r\nhello\r\n0\r\n\r\n";
let result = parser.parse_batch(request);
assert_eq!(result.requests.len(), 0);
assert!(matches!(
result.error,
Some(BatchParseError::InvalidSyntax(_))
));
}
#[test]
fn test_request_exceeding_buffer_capacity_errors() {
let config = BatchConfig::builder().buffer_size(128).build();
let parser = BatchParser::new(config);
let request =
b"POST /api HTTP/1.1\r\nHost: a.com\r\nContent-Length: 100000\r\n\r\npartial body";
let result = parser.parse_batch(request);
assert_eq!(result.requests.len(), 0);
assert!(matches!(
result.error,
Some(BatchParseError::BufferOverflow)
));
}
#[test]
fn test_max_headers_clamped_to_supported() {
let config = BatchConfig::builder().max_headers(500).build();
assert_eq!(config.max_headers, MAX_SUPPORTED_HEADERS);
let config = BatchConfig::builder().max_headers(10).build();
assert_eq!(config.max_headers, 10);
}
#[test]
fn test_batch_stats() {
let stats = BatchStats::new();
stats.record_batch(5, 1024);
stats.record_batch(10, 2048);
assert_eq!(stats.batches_processed(), 2);
assert_eq!(stats.total_requests(), 15);
assert_eq!(stats.total_bytes_read(), 3072);
assert_eq!(stats.max_batch_size(), 10);
}
#[test]
fn test_to_http_request() {
let parser = BatchParser::new(BatchConfig::default());
let request = b"POST /api/data HTTP/1.1\r\nHost: example.com\r\nContent-Type: application/json\r\nContent-Length: 13\r\n\r\n{\"key\":\"val\"}";
let result = parser.parse_batch(request);
assert_eq!(result.requests.len(), 1);
let http_req = result.requests[0].to_http_request();
assert_eq!(http_req.method, "POST");
assert_eq!(http_req.path, "/api/data");
assert_eq!(
http_req.headers.get("Content-Type"),
Some("application/json")
);
assert_eq!(http_req.body_ref(), b"{\"key\":\"val\"}");
}
}