use async_trait::async_trait;
use std::{env, io};
#[cfg(feature = "scgi")]
use bytes::BytesMut;
#[cfg(feature = "scgi")]
use std::{collections::HashMap, error::Error, sync::Arc};
#[cfg(feature = "scgi")]
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::{TcpStream, ToSocketAddrs},
};
use crate::{
application::Application,
error::{GemError, ToGemError},
request::Request,
response::Response,
};
#[cfg(feature = "cgi")]
async fn send_cgi_response(response: Response) {
if let Err(err) = response.send_sync(&mut io::stdout()).await {
eprintln!("Could not send response: {err}");
};
}
#[cfg(feature = "cgi")]
fn get_cgi_header(key: &str) -> Result<String, GemError> {
env::var(key).into_gem()
}
#[cfg(feature = "cgi")]
#[async_trait]
pub trait Cgi: Application + Sized + Send + Sync + 'static {
async fn run_cgi(self) {
let request = match Request::parse_request(get_cgi_header) {
Ok(request) => request,
Err(err) => {
eprintln!("Invalid CGI header: {err}");
send_cgi_response(Response::error_cgi("Invalid CGI header")).await;
return;
}
};
let response = match self.handle_request(request).await {
Ok(response) => response,
Err(err) => {
eprintln!("Error while handling request: {err}");
match err.downcast::<GemError>() {
Ok(err) => Response::from(*err),
Err(_) => Response::error_cgi("Internal Server Error"),
}
}
};
send_cgi_response(response).await;
}
}
#[cfg(feature = "cgi")]
impl<A> Cgi for A where A: Application + Send + Sync + 'static {}
#[cfg(feature = "scgi")]
async fn read_scgi_request(conn: &mut TcpStream) -> Result<Request, Box<dyn Error + Send + Sync>> {
let mut buf = Vec::new();
loop {
let chr = conn.read_u8().await?;
if chr == b':' {
break;
}
buf.push(chr);
}
let size: usize = String::from_utf8(buf)?.parse()?;
let mut buffer = BytesMut::zeroed(size);
conn.read_exact(buffer.as_mut()).await?;
let mut headers = HashMap::new();
let mut values = buffer.as_ref().split(|c| *c == b'\0');
loop {
if let Some(key) = values.next() {
if let Some(val) = values.next() {
let key = std::str::from_utf8(key)?;
let val = std::str::from_utf8(val)?;
headers.insert(key, val);
} else {
if !key.is_empty() {
return Err(Box::new(GemError::runtime_error("Missing header value")));
}
break;
}
} else {
break;
}
}
Ok(Request::parse_request(|k| {
headers
.get(k)
.map(|v| (*v).to_owned())
.ok_or(GemError::runtime_error(format!("Missing header {k}")))
})?)
}
#[cfg(feature = "scgi")]
async fn send_scgi_response(mut conn: TcpStream, response: Response) {
if let Err(e) = response.send_async(&mut conn).await {
eprintln!("Could not send body: {e}");
}
if let Err(e) = conn.shutdown().await {
eprintln!("Could not shutdown connection: {e}");
};
}
#[cfg(feature = "scgi")]
#[async_trait]
pub trait Scgi: Application + Sized + Send + Sync + 'static {
async fn run_scgi<A>(self, addr: A) -> io::Result<()>
where
A: ToSocketAddrs + Send + Sync,
{
let listener = tokio::net::TcpListener::bind(addr).await?;
println!("Listening to {:?}", listener.local_addr()?);
let self_arc = Arc::new(self);
loop {
let (mut conn, _) = listener.accept().await?;
let self_ref = self_arc.clone();
tokio::spawn(async move {
let mut path = None;
let response = match read_scgi_request(&mut conn).await {
Ok(request) => {
path = Some(request.path.clone());
match self_ref.handle_request(request).await {
Ok(response) => response,
Err(err) => {
eprintln!("Error while handling request: {err}");
match err.downcast::<GemError>() {
Ok(err) => Response::from(*err),
Err(_) => Response::error_cgi("Internal Server Error"),
}
}
}
}
Err(e) => {
eprintln!("Invalid SCGI header: {e}");
Response::error_cgi("Invalid CGI header")
}
};
println!(
"{}\t{}\t{}",
path.unwrap_or("".into()),
response.code,
response.meta
);
send_scgi_response(conn, response).await;
});
}
}
}
#[cfg(feature = "scgi")]
impl<A> Scgi for A where A: Application + Sized + Send + Sync + 'static {}