use std::sync::Arc;
use std::sync::atomic::AtomicBool;
use crate::config::AppConfig;
use crate::errors::OrionError;
pub fn create_tcp_listener(addr: &str) -> Result<tokio::net::TcpListener, OrionError> {
let socket_addr = addr
.parse::<std::net::SocketAddr>()
.map_err(|e| OrionError::internal(format!("Invalid address '{addr}': {e}")))?;
let domain = if socket_addr.is_ipv4() {
socket2::Domain::IPV4
} else {
socket2::Domain::IPV6
};
let map_err = |stage: &str, e: std::io::Error| OrionError::Internal {
context: format!("Failed to {stage} for {addr}"),
source: Some(Box::new(e)),
};
let socket = socket2::Socket::new(domain, socket2::Type::STREAM, Some(socket2::Protocol::TCP))
.map_err(|e| map_err("create socket", e))?;
socket.set_tcp_nodelay(true).ok();
socket.set_reuse_address(true).ok();
socket
.bind(&socket_addr.into())
.map_err(|e| map_err("bind", e))?;
socket.listen(1024).map_err(|e| map_err("listen", e))?;
socket.set_nonblocking(true).ok();
tokio::net::TcpListener::from_std(socket.into())
.map_err(|e| map_err("create async listener", e))
}
pub async fn serve_tls(
config: Arc<AppConfig>,
ready: Arc<AtomicBool>,
router: axum::Router,
handle: axum_server::Handle<std::net::SocketAddr>,
shutdown: impl std::future::Future<Output = ()> + Send + 'static,
) -> Result<(), OrionError> {
let rustls_config =
super::tls::load_rustls_config(&config.server.tls.cert_path, &config.server.tls.key_path)
.await?;
let addr = format!("{}:{}", config.server.host, config.server.port);
let bind_addr: std::net::SocketAddr = addr
.parse()
.map_err(|e| OrionError::internal(format!("Invalid address '{addr}': {e}")))?;
let shutdown_handle = handle.clone();
let drain_secs = config.server.shutdown_drain_secs;
let force_timeout_secs = config.server.shutdown_force_timeout_secs;
let ready_for_drain = ready.clone();
tokio::spawn(async move {
super::drain::drain_gate(
shutdown,
ready_for_drain,
std::time::Duration::from_secs(drain_secs),
)
.await;
let force =
(force_timeout_secs > 0).then(|| std::time::Duration::from_secs(force_timeout_secs));
shutdown_handle.graceful_shutdown(force);
});
tracing::info!(
address = %addr,
storage = %crate::connector::redact_url_secrets_or_raw(&config.storage.url),
tls = true,
"Orion is ready (HTTPS)"
);
axum_server::bind_rustls(bind_addr, rustls_config)
.handle(handle)
.serve(router.into_make_service_with_connect_info::<std::net::SocketAddr>())
.await
.map_err(|e| OrionError::Internal {
context: format!("HTTPS server error on {addr}"),
source: Some(Box::new(e)),
})
}
pub async fn serve_metrics(
listener: tokio::net::TcpListener,
config: Arc<AppConfig>,
router: axum::Router,
shutdown: impl std::future::Future<Output = ()> + Send + 'static,
) -> Result<(), OrionError> {
let drain = std::time::Duration::from_secs(config.server.shutdown_drain_secs);
let addr = listener
.local_addr()
.map(|a| a.to_string())
.unwrap_or_default();
tracing::info!(
address = %addr,
"Metrics listener ready (GET /metrics, unauthenticated)"
);
axum::serve(listener, router)
.with_graceful_shutdown(async move {
shutdown.await;
tokio::time::sleep(drain).await;
})
.await
.map_err(|e| OrionError::Internal {
context: format!("Metrics server error on {addr}"),
source: Some(Box::new(e)),
})
}
pub async fn serve_plain_http(
listener: tokio::net::TcpListener,
config: Arc<AppConfig>,
ready: Arc<AtomicBool>,
router: axum::Router,
shutdown: impl std::future::Future<Output = ()> + Send + 'static,
) -> Result<(), OrionError> {
let drain_secs = config.server.shutdown_drain_secs;
let force_timeout_secs = config.server.shutdown_force_timeout_secs;
let addr = listener
.local_addr()
.map(|a| a.to_string())
.unwrap_or_default();
tracing::info!(
address = %addr,
storage = %crate::connector::redact_url_secrets_or_raw(&config.storage.url),
tcp_nodelay = true,
"Orion is ready"
);
let (drained_tx, drained_rx) = tokio::sync::oneshot::channel::<()>();
let gate = async move {
super::drain::drain_gate(shutdown, ready, std::time::Duration::from_secs(drain_secs)).await;
let _ = drained_tx.send(());
};
let serve = axum::serve(
listener,
router.into_make_service_with_connect_info::<std::net::SocketAddr>(),
)
.with_graceful_shutdown(gate);
let map_err = |e: std::io::Error| OrionError::Internal {
context: format!("HTTP server error on {addr}"),
source: Some(Box::new(e)),
};
if force_timeout_secs == 0 {
serve.await.map_err(map_err)?;
} else {
tokio::select! {
result = serve => result.map_err(map_err)?,
_ = async {
match drained_rx.await {
Ok(()) => {
tokio::time::sleep(std::time::Duration::from_secs(force_timeout_secs)).await;
tracing::warn!(
force_timeout_secs,
"Shutdown force timeout elapsed; aborting remaining in-flight connections"
);
}
Err(_) => std::future::pending::<()>().await,
}
} => {}
}
}
Ok(())
}