use std::cell::RefCell;
use std::collections::VecDeque;
use std::io::Read;
use std::net::{TcpListener, TcpStream, ToSocketAddrs};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use std::thread;
use crate::config::Config;
use crate::dispatch;
use crate::error::Error;
use crate::hpack_decoder::{Decoder, PathKind};
use crate::http2::*;
use crate::macros::*;
use crate::response_end::ResponseEnd;
use crate::{PajamaxService, Response};
pub fn serve_with_config<A>(
services: Vec<Arc<dyn PajamaxService + Send + Sync + 'static>>,
config: Config,
addr: A,
) -> std::io::Result<()>
where
A: ToSocketAddrs,
{
let concurrent = Arc::new(AtomicUsize::new(0));
let listener = TcpListener::bind(addr)?;
for c in listener.incoming() {
if concurrent.load(Ordering::Relaxed) >= config.max_concurrent_connections {
error!("drop new connection for limit");
continue;
}
concurrent.fetch_add(1, Ordering::Relaxed);
let c = c?;
info!("new connection from {}", c.local_addr().unwrap().ip());
c.set_read_timeout(Some(config.idle_timeout))?;
c.set_write_timeout(Some(config.write_timeout))?;
let concurrent = concurrent.clone();
let services = services.clone();
thread::Builder::new()
.name(String::from("pajamax-w"))
.spawn(move || {
match handle(services, c, config) {
Ok(_) => info!("connection closed"),
Err(err) => error!("connection fail: {:?}", err),
}
concurrent.fetch_sub(1, Ordering::Relaxed);
})
.unwrap();
}
unreachable!();
}
thread_local! {
static RESPONSE_END: RefCell<ResponseEnd> = panic!();
}
struct Stream {
id: u32,
isvc: usize, req_disc: usize,
}
pub fn local_build_response<Reply>(stream_id: u32, response: Response<Reply>) -> Result<(), Error>
where
Reply: prost::Message,
{
RESPONSE_END.with_borrow_mut(|resp_end| Ok(resp_end.build(stream_id, response)?))
}
pub fn handle(
services: Vec<Arc<dyn PajamaxService + Send + Sync + 'static>>,
mut c: TcpStream,
config: Config,
) -> Result<(), Error> {
handshake(&mut c, &config)?;
trace!("handshake done");
let mut input = Vec::new();
input.resize(config.max_frame_size, 0);
let mut streams = VecDeque::new();
let mut hpack_decoder: Decoder = Decoder::new();
let mut route_cache = Vec::new();
let c2 = Arc::new(Mutex::new(c.try_clone()?));
if services.iter().any(|svc| svc.is_dispatch_mode()) {
dispatch::new_response_routine(c2.clone(), &config);
}
RESPONSE_END.set(ResponseEnd::new(c2, &config));
let mut last_end = 0;
while let Ok(len) = c.read(&mut input[last_end..]) {
trace!("receive data {len}");
if len == 0 {
return Ok(());
}
let end = last_end + len;
let mut data_len = 0;
let mut pos = 0;
while let Some(frame) = Frame::parse(&input[pos..end]) {
pos += Frame::HEAD_SIZE + frame.len;
trace!(
"get frame {:?} {:?}, len:{}, stream_id:{}",
frame.kind,
frame.flags,
frame.stream_id,
frame.len
);
match frame.kind {
FrameKind::Headers => {
let headers_buf = frame.process_headers()?;
let (isvc, req_disc) = match hpack_decoder.find_path(headers_buf)? {
PathKind::Cached(cached) => {
trace!("route cache hit: {cached}");
route_cache[cached]
}
PathKind::Plain(path) => {
let len0 = route_cache.len();
for (i, svc) in services.iter().enumerate() {
if let Some(req_disc) = svc.route(&path) {
route_cache.push((i, req_disc));
break;
}
}
if route_cache.len() == len0 {
return Err(Error::UnknownMethod(
String::from_utf8_lossy(&path).into(),
));
}
trace!(
"route cache new ({len0}): {}",
String::from_utf8_lossy(&path)
);
route_cache[len0]
}
};
streams.push_back(Stream {
id: frame.stream_id,
isvc,
req_disc,
});
}
FrameKind::Data => {
let req_buf = frame.process_data()?;
if req_buf.len() == 0 {
continue;
}
if req_buf.len() < 5 {
return Err(Error::InvalidHttp2("DATA frame too short for grpc"));
}
let req_buf = &req_buf[5..];
let Some(i) = streams.iter().position(|s| s.id == frame.stream_id) else {
return Err(Error::InvalidHttp2("DATA frame without HEADER"));
};
let Stream { id, isvc, req_disc } = streams.remove(i).unwrap();
trace!("handle isvc:{isvc}, req_disc:{req_disc}");
services[isvc].handle(req_disc, req_buf, id)?;
data_len += frame.len;
}
_ => (),
}
}
RESPONSE_END.with_borrow_mut(|resp_end| {
resp_end.window_update(data_len);
resp_end.flush()
})?;
if pos == 0 {
return Err(Error::InvalidHttp2("too long frame"));
}
if pos < end {
trace!("left data {}", end - pos);
input.copy_within(pos..end, 0);
last_end = end - pos;
} else {
last_end = 0;
}
}
Ok(())
}