use std::{str::FromStr, thread};
use tiny_http::{Header, HeaderField, Request, Response, Server};
use crate::store::store;
#[derive(Debug)]
pub struct StatsServer {
listen: String,
secret: String,
}
impl StatsServer {
pub fn new<S: Into<String>>(listen: S, secret: S) -> Self {
Self {
listen: listen.into(),
secret: secret.into(),
}
}
pub fn serve(&self) -> anyhow::Result<()> {
let secret = self.secret.to_string();
let server = Server::http(&self.listen)
.map_err(|error| anyhow::anyhow!("start stats server failed: {error}"))?;
thread::spawn(move || {
loop {
if let Ok(req) = server.recv() {
handle(req, secret.clone());
}
}
});
Ok(())
}
}
fn handle(req: Request, secret: String) {
let resp = if header_matches(req.headers(), "Authorization", &secret) {
match req.url() {
"/traffic" => handle_traffic(),
"/dump/streams" => handle_dump_streams(),
_ => handle_404(),
}
} else {
handle_auth_failed()
};
match resp {
Ok(resp) => {
if let Err(e) = req.respond(resp) {
warn!("Error on response {}", e);
}
}
Err(e) => {
warn!("handle failed {}", e);
if let Err(e) =
req.respond(Response::from_string("Internal Server Srror").with_status_code(500))
{
warn!("Error on response {}", e);
}
}
}
}
fn handle_dump_streams() -> anyhow::Result<Response<std::io::Cursor<Vec<u8>>>> {
let conns_all = store().get_conns_all();
let body = serde_json::to_string(&conns_all)?;
Ok(Response::from_data(body))
}
fn handle_traffic() -> anyhow::Result<Response<std::io::Cursor<Vec<u8>>>> {
let traffic_all = store().get_traffic_all();
let body = serde_json::to_string(&traffic_all)?;
Ok(Response::from_data(body))
}
fn handle_404() -> anyhow::Result<Response<std::io::Cursor<Vec<u8>>>> {
Ok(Response::from_data("Not Found").with_status_code(404))
}
fn handle_auth_failed() -> anyhow::Result<Response<std::io::Cursor<Vec<u8>>>> {
Ok(Response::from_data("Unauthorized").with_status_code(401))
}
fn header_matches(headers: &[Header], header: &str, expected: &str) -> bool {
get_header(headers, header).as_deref() == Some(expected)
}
fn get_header(headers: &[Header], header: &str) -> Option<String> {
headers
.iter()
.find(|h| HeaderField::from_str(header).is_ok_and(|field| h.field == field))
.map(|h| h.value.to_string())
}
#[cfg(test)]
mod tests {
use std::io::Read as _;
use std::time::{SystemTime, UNIX_EPOCH};
use tiny_http::{Header, StatusCode};
use super::{
StatsServer, get_header, handle_404, handle_auth_failed, handle_dump_streams,
handle_traffic, header_matches,
};
use crate::store::store;
#[test]
fn test_start_server() {
StatsServer::new("127.0.0.1:8400", "asd");
}
#[test]
fn get_header_finds_matching_value() {
let headers = vec![
Header::from_bytes(&b"Authorization"[..], &b"secret"[..]).unwrap(),
Header::from_bytes(&b"X-Test"[..], &b"value"[..]).unwrap(),
];
assert_eq!(
get_header(&headers, "Authorization").as_deref(),
Some("secret")
);
assert!(header_matches(&headers, "Authorization", "secret"));
assert!(!header_matches(&headers, "Authorization", "other"));
assert_eq!(get_header(&headers, "Missing"), None);
}
#[test]
fn handle_auth_failed_returns_401() {
let response = handle_auth_failed().unwrap();
assert_eq!(response.status_code(), StatusCode(401));
}
#[test]
fn handle_404_returns_404() {
let response = handle_404().unwrap();
assert_eq!(response.status_code(), StatusCode(404));
}
#[test]
fn handle_traffic_and_dump_streams_serialize_store_state() {
let store = store();
let suffix = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos();
let user = format!("alice-{suffix}");
let hash = format!("hash-{suffix}");
let conn_id = format!("conn-{suffix}");
store.insert_conn(
user.clone(),
hash,
conn_id.clone(),
"127.0.0.1:1".to_string(),
"example.com:443".to_string(),
false,
);
store.add_up(&conn_id, 3);
store.add_down(&conn_id, 4);
let traffic = handle_traffic().unwrap();
let dump = handle_dump_streams().unwrap();
assert_eq!(traffic.status_code(), StatusCode(200));
assert_eq!(dump.status_code(), StatusCode(200));
let traffic_text = response_body(traffic);
let dump_text = response_body(dump);
assert!(traffic_text.contains(&format!("\"{user}\"")));
assert!(traffic_text.contains("\"up\":3"));
assert!(traffic_text.contains("\"down\":4"));
assert!(dump_text.contains(&format!("\"{conn_id}\"")));
assert!(dump_text.contains("\"example.com:443\""));
}
fn response_body(response: tiny_http::Response<std::io::Cursor<Vec<u8>>>) -> String {
let mut bytes = Vec::new();
response.into_reader().read_to_end(&mut bytes).unwrap();
String::from_utf8(bytes).unwrap()
}
}