use anyhow::Result;
use axum::extract::Request;
use axum::response::IntoResponse;
use axum::routing::{get, post, Route};
use axum::{middleware, Router};
use std::convert::Infallible;
use std::{collections::HashMap, net::Ipv4Addr, sync::Arc};
use tokio::sync::RwLock;
use tower::{Layer, Service};
use tower_http::trace::{DefaultMakeSpan, DefaultOnRequest, DefaultOnResponse, TraceLayer};
use tracing::{info, Level};
use crate::config::{GitHubConfig, DEFAULT_HOST_ADDR, DEFAULT_PORT};
use crate::core::{Context, EventHandlerFn};
use crate::github::{
middlewares::{github_event_middleware, verify_hmac_middleware, HmacConfig},
GitHubAuth, GitHubClient,
};
use super::handlers;
pub type WebhookEventKind = String;
#[derive(Clone, Default)]
pub struct AppState {
pub handlers: Arc<RwLock<HashMap<WebhookEventKind, Vec<EventHandlerFn>>>>,
pub github_client: Option<Arc<GitHubClient>>,
}
pub struct WebhookServer {
state: AppState,
pub host: Ipv4Addr,
pub port: u16,
router: Option<Router>,
}
impl Default for WebhookServer {
fn default() -> Self {
Self::new_default()
}
}
impl WebhookServer {
pub async fn new(
host: Ipv4Addr,
port: u16,
github_config: GitHubConfig,
secret: &str,
hmac_header: &str,
) -> Result<Self> {
let auth = GitHubAuth::from_config(&github_config);
let github_client = Arc::new(GitHubClient::new(auth).await?);
let state = AppState {
handlers: Arc::new(RwLock::new(HashMap::new())),
github_client: Some(github_client),
};
let hmac_config = HmacConfig::new(secret.into(), hmac_header.into());
let router = create_router(state.clone(), hmac_config);
Ok(Self {
state,
host,
port,
router: Some(router),
})
}
pub fn new_default() -> Self {
let state = AppState {
handlers: Arc::new(RwLock::new(HashMap::new())),
github_client: None,
};
let router = create_router(state.clone(), HmacConfig::default());
Self {
state,
host: DEFAULT_HOST_ADDR,
port: DEFAULT_PORT,
router: Some(router),
}
}
pub async fn start(&self) -> Result<()> {
let listener = tokio::net::TcpListener::bind((self.host, self.port)).await?;
info!("Webhook server started on {}:{}", self.host, self.port);
let router = self
.router
.clone()
.ok_or(anyhow::anyhow!("Cannot initialize router"))?;
axum::serve(listener, router).await?;
Ok(())
}
pub async fn on<F, Fut, E>(&mut self, event: impl Into<String>, handler: F, extra: Arc<E>)
where
F: Fn(Context, Arc<E>) -> Fut + Send + Sync + 'static,
Fut: std::future::Future<Output = Result<()>> + Send + 'static,
E: Send + Sync + 'static,
{
let event = event.into();
let boxed_handler: EventHandlerFn = Box::new(move |context| {
let extra = extra.clone();
Box::pin(handler(context, extra))
});
self.state
.handlers
.write()
.await
.entry(event)
.or_default()
.push(boxed_handler);
}
pub fn add_middleware<T>(&mut self, layer: T) -> Result<()>
where
T: Layer<Route> + Clone + Send + Sync + 'static,
T::Service: Service<Request> + Clone + Send + Sync + 'static,
<T::Service as Service<Request>>::Response: IntoResponse + 'static,
<T::Service as Service<Request>>::Error: Into<Infallible> + 'static,
<T::Service as Service<Request>>::Future: Send + 'static,
{
if let Some(router) = self.router.take() {
let new_router = router.layer(layer);
self.router = Some(new_router);
Ok(())
} else {
Err(anyhow::anyhow!("Router not initialized"))
}
}
pub fn github_client(&self) -> Option<&Arc<GitHubClient>> {
self.state.github_client.as_ref()
}
}
fn create_router(state: AppState, hmac_config: HmacConfig) -> Router {
let cors_layer = tower_http::cors::CorsLayer::new()
.allow_origin(tower_http::cors::Any)
.allow_methods(tower_http::cors::Any)
.allow_headers(tower_http::cors::Any);
let trace_layer = TraceLayer::new_for_http()
.make_span_with(DefaultMakeSpan::new().include_headers(true))
.on_request(DefaultOnRequest::new().level(Level::INFO))
.on_response(
DefaultOnResponse::new()
.level(Level::INFO)
.latency_unit(tower_http::LatencyUnit::Micros),
);
Router::new()
.route("/health", get(handlers::handle_health))
.route(
"/webhook",
post(handlers::handle_webhook)
.layer(middleware::from_fn_with_state(
hmac_config,
verify_hmac_middleware,
))
.layer(middleware::from_fn(github_event_middleware)),
)
.layer(trace_layer)
.layer(cors_layer)
.with_state(state)
}