wowlab_sentinel/http/
mod.rs1mod api_error;
2pub(crate) mod auth;
3mod routes;
4mod services;
5
6use std::sync::Arc;
7
8use axum::{
9 http::{HeaderValue, Method, header},
10 response::{IntoResponse, Response},
11};
12use serde::Serialize;
13use tokio_util::sync::CancellationToken;
14use tower_http::{
15 cors::{AllowOrigin, CorsLayer},
16 trace::TraceLayer,
17};
18
19use crate::state::ServerState;
20
21crate::db_client!(HttpDb, "http", max = 10);
22
23#[derive(Debug)]
24pub(crate) struct PrettyJson<T>(pub(crate) T);
25
26impl<T> IntoResponse for PrettyJson<T>
27where
28 T: Serialize,
29{
30 fn into_response(self) -> Response {
31 let body = serde_json::to_string_pretty(&self.0).expect("serializable response type");
32
33 ([(header::CONTENT_TYPE, "application/json")], body).into_response()
34 }
35}
36
37pub(crate) async fn run(
38 state: Arc<ServerState>,
39 shutdown: CancellationToken,
40) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
41 let port = state.config.http_port;
42 let patterns: Vec<String> = state.config.cors_origins.clone();
43 let cors = CorsLayer::new()
44 .allow_origin(AllowOrigin::predicate(move |origin: &HeaderValue, _| {
45 let Ok(origin) = origin.to_str() else {
46 return false;
47 };
48
49 patterns.iter().any(|p| origin_matches(origin, p))
50 }))
51 .allow_methods([Method::GET, Method::POST, Method::OPTIONS])
52 .allow_headers([
53 header::CONTENT_TYPE,
54 header::HeaderName::from_static("x-node-key"),
55 header::HeaderName::from_static("x-node-sig"),
56 header::HeaderName::from_static("x-node-ts"),
57 ]);
58
59 let app = routes::router(&state)
60 .layer(cors)
61 .layer(TraceLayer::new_for_http())
62 .with_state(state);
63
64 let listener = tokio::net::TcpListener::bind(format!("[::]:{port}")).await?;
65
66 tracing::info!(port, "HTTP server listening");
67
68 axum::serve(listener, app)
69 .with_graceful_shutdown(async move { shutdown.cancelled().await })
70 .await?;
71
72 Ok(())
73}
74
75fn origin_matches(origin: &str, pattern: &str) -> bool {
76 if !pattern.contains('*') {
77 return origin == pattern;
78 }
79
80 let parts: Vec<&str> = pattern.split('*').collect();
81
82 if parts.len() != 2 {
84 return origin == pattern;
85 }
86
87 origin.starts_with(parts[0]) && origin.ends_with(parts[1])
88}