mod rest;
mod ws;
use std::future::Future;
use std::net::SocketAddr;
use std::sync::Arc;
use axum::Router;
use crate::agent::Agent;
use crate::auth::AuthScheme;
use crate::error::{Error, Result};
pub fn router(agent: Arc<Agent>) -> Router {
rest::router(agent, None)
}
pub fn router_with_auth(agent: Arc<Agent>, auth: AuthScheme) -> Router {
rest::router(agent, Some(Arc::new(auth)))
}
async fn bind(addr: SocketAddr) -> Result<tokio::net::TcpListener> {
tokio::net::TcpListener::bind(addr)
.await
.map_err(|e| Error::Server(format!("failed to bind {addr}: {e}")))
}
pub async fn serve(agent: Arc<Agent>, addr: SocketAddr) -> Result<()> {
let listener = bind(addr).await?;
tracing::info!("REST/WebSocket API listening on http://{addr}");
axum::serve(listener, router(agent))
.await
.map_err(|e| Error::Server(e.to_string()))
}
pub async fn serve_with_auth(agent: Arc<Agent>, addr: SocketAddr, auth: AuthScheme) -> Result<()> {
let listener = bind(addr).await?;
tracing::info!("REST/WebSocket API (authenticated) listening on http://{addr}");
axum::serve(listener, router_with_auth(agent, auth))
.await
.map_err(|e| Error::Server(e.to_string()))
}
pub async fn serve_with_shutdown(
agent: Arc<Agent>,
addr: SocketAddr,
signal: impl Future<Output = ()> + Send + 'static,
) -> Result<()> {
let listener = bind(addr).await?;
tracing::info!("REST/WebSocket API listening on http://{addr} (graceful shutdown armed)");
axum::serve(listener, router(agent))
.with_graceful_shutdown(signal)
.await
.map_err(|e| Error::Server(e.to_string()))
}
pub async fn shutdown_signal() {
let ctrl_c = async {
let _ = tokio::signal::ctrl_c().await;
};
#[cfg(unix)]
{
let mut sigterm = tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())
.expect("failed to install SIGTERM handler");
tokio::select! {
_ = ctrl_c => {}
_ = sigterm.recv() => {}
}
}
#[cfg(not(unix))]
ctrl_c.await;
tracing::info!("shutdown signal received");
}
#[cfg(feature = "tls")]
#[cfg_attr(docsrs, doc(cfg(feature = "tls")))]
pub async fn serve_tls(
agent: Arc<Agent>,
addr: SocketAddr,
tls: &crate::tls::TlsConfig,
) -> Result<()> {
let (cert, key) = tls.pem_pair()?;
let config = axum_server::tls_rustls::RustlsConfig::from_pem(cert, key)
.await
.map_err(|e| Error::Server(format!("invalid TLS material: {e}")))?;
tracing::info!("REST/WebSocket API listening on https://{addr}");
axum_server::bind_rustls(addr, config)
.serve(router(agent).into_make_service())
.await
.map_err(|e| Error::Server(e.to_string()))
}
impl Agent {
pub async fn serve(self: Arc<Self>, addr: SocketAddr) -> Result<()> {
serve(self, addr).await
}
}