use std::collections::BTreeMap;
use std::panic::AssertUnwindSafe;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use axum::{
Json, Router,
extract::{Request, State},
http::{StatusCode, header::RETRY_AFTER},
middleware::{Next, from_fn, from_fn_with_state},
response::{IntoResponse, Response},
routing::get,
};
use futures_util::FutureExt as _;
use tokio_util::sync::CancellationToken;
use tower::ServiceExt as _;
use url::Url;
use cf_system_sdks::directory::{DirectoryClient, RegisterInstanceInfo, ServiceEndpoint};
use toolkit_canonical_errors::CanonicalError;
use toolkit_http_middleware::{internal_auth_middleware, security_context_middleware};
use toolkit_security::{DynBearerAuthenticator, DynInternalAuthenticator};
use super::readiness::ReadinessState;
use crate::api::canonical_error_middleware;
const DRAIN_RETRY_AFTER_SECONDS: u64 = 5;
const STARTING_RETRY_AFTER_SECONDS: u64 = 1;
#[non_exhaustive]
pub struct OopServeOptions {
pub gear_name: String,
pub instance_id: String,
pub version: Option<String>,
pub advertise_uri: String,
pub listen_addr: std::net::SocketAddr,
pub probe_bind_addr: Option<std::net::SocketAddr>,
pub drain_timeout: Duration,
pub heartbeat_interval: Duration,
pub healthcheck_timeout: Duration,
pub directory: Arc<dyn DirectoryClient>,
pub bearer_authenticator: Option<DynBearerAuthenticator>,
pub internal_authenticator: Option<DynInternalAuthenticator>,
pub labels: BTreeMap<String, String>,
}
impl OopServeOptions {
#[must_use]
pub fn new(
gear_name: String,
instance_id: String,
advertise_uri: String,
listen_addr: std::net::SocketAddr,
directory: Arc<dyn DirectoryClient>,
) -> Self {
Self {
gear_name,
instance_id,
version: None,
advertise_uri,
listen_addr,
probe_bind_addr: None,
drain_timeout: Duration::from_secs(30),
heartbeat_interval: Duration::from_secs(5),
healthcheck_timeout: Duration::from_millis(500),
directory,
bearer_authenticator: None,
internal_authenticator: None,
labels: BTreeMap::new(),
}
}
#[must_use]
pub fn with_version(mut self, version: Option<String>) -> Self {
self.version = version;
self
}
#[must_use]
pub fn with_probe_bind_addr(mut self, addr: Option<std::net::SocketAddr>) -> Self {
self.probe_bind_addr = addr;
self
}
#[must_use]
pub fn with_drain_timeout(mut self, timeout: Duration) -> Self {
self.drain_timeout = timeout;
self
}
#[must_use]
pub fn with_heartbeat_interval(mut self, interval: Duration) -> Self {
self.heartbeat_interval = interval;
self
}
#[must_use]
pub fn with_healthcheck_timeout(mut self, timeout: Duration) -> Self {
self.healthcheck_timeout = timeout;
self
}
#[must_use]
pub fn with_bearer_authenticator(mut self, auth: Option<DynBearerAuthenticator>) -> Self {
self.bearer_authenticator = auth;
self
}
#[must_use]
pub fn with_internal_authenticator(mut self, auth: Option<DynInternalAuthenticator>) -> Self {
self.internal_authenticator = auth;
self
}
#[must_use]
pub fn with_labels(mut self, labels: BTreeMap<String, String>) -> Self {
self.labels = labels;
self
}
}
impl std::fmt::Debug for OopServeOptions {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OopServeOptions")
.field("gear_name", &self.gear_name)
.field("instance_id", &self.instance_id)
.field("version", &self.version)
.field("advertise_uri", &self.advertise_uri)
.field("listen_addr", &self.listen_addr)
.field("probe_bind_addr", &self.probe_bind_addr)
.field("drain_timeout", &self.drain_timeout)
.field("heartbeat_interval", &self.heartbeat_interval)
.field("healthcheck_timeout", &self.healthcheck_timeout)
.field("bearer_authenticator", &self.bearer_authenticator.is_some())
.field(
"internal_authenticator",
&self.internal_authenticator.is_some(),
)
.field("labels", &self.labels)
.finish_non_exhaustive()
}
}
#[derive(Clone, Default)]
struct LateRoutes {
inner: Arc<arc_swap::ArcSwapOption<LateRoutesInner>>,
}
struct LateRoutesInner {
gear: Router,
openapi: Arc<str>,
}
impl LateRoutes {
fn publish(&self, gear: Router, openapi: Arc<str>) {
self.inner
.store(Some(Arc::new(LateRoutesInner { gear, openapi })));
}
fn gear(&self) -> Option<Router> {
self.inner.load_full().map(|i| i.gear.clone())
}
fn openapi(&self) -> Option<Arc<str>> {
self.inner.load_full().map(|i| Arc::clone(&i.openapi))
}
}
fn starting_response() -> Response {
(
StatusCode::SERVICE_UNAVAILABLE,
[(RETRY_AFTER, STARTING_RETRY_AFTER_SECONDS.to_string())],
"starting",
)
.into_response()
}
#[derive(Clone)]
struct ProbeState {
readiness: Arc<ReadinessState>,
late: LateRoutes,
}
fn build_outer_router(readiness: Arc<ReadinessState>, late: LateRoutes) -> Router {
let fallback_late = late.clone();
build_probe_router(readiness, late).fallback_service(tower::util::service_fn(
move |req: Request| {
let late = fallback_late.clone();
async move {
match late.gear() {
Some(router) => Ok::<_, std::convert::Infallible>(
router.oneshot(req).await.unwrap_or_else(|e| match e {}),
),
None => Ok::<_, std::convert::Infallible>(starting_response()),
}
}
},
))
}
fn build_probe_router(readiness: Arc<ReadinessState>, late: LateRoutes) -> Router {
let state = ProbeState { readiness, late };
Router::new()
.route("/healthz", get(healthz))
.route("/readyz", get(readyz))
.route("/health", get(health))
.route("/.well-known/openapi.json", get(openapi))
.route("/openapi.json", get(openapi))
.with_state(state)
}
async fn healthz() -> &'static str {
"ok"
}
async fn readyz(State(state): State<ProbeState>) -> Response {
let report = state.readiness.evaluate().await;
let status = if report.ready {
StatusCode::OK
} else {
StatusCode::SERVICE_UNAVAILABLE
};
(status, Json(report)).into_response()
}
async fn health(State(state): State<ProbeState>) -> impl IntoResponse {
let report = state.readiness.health_report().await;
let status = if report.is_ready() {
StatusCode::OK
} else {
StatusCode::SERVICE_UNAVAILABLE
};
(status, Json(report))
}
async fn openapi(State(state): State<ProbeState>) -> Response {
match state.late.openapi() {
Some(spec) => (
StatusCode::OK,
[(axum::http::header::CONTENT_TYPE, "application/json")],
spec.to_string(),
)
.into_response(),
None => starting_response(),
}
}
#[derive(Clone)]
struct DrainGuard {
readiness: Arc<ReadinessState>,
in_flight: Arc<AtomicUsize>,
}
impl DrainGuard {
fn new(readiness: Arc<ReadinessState>) -> Self {
Self {
readiness,
in_flight: Arc::new(AtomicUsize::new(0)),
}
}
fn in_flight(&self) -> usize {
self.in_flight.load(Ordering::SeqCst)
}
fn begin_drain(&self) {
self.readiness.set_draining(true);
}
}
struct InFlightGuard(Arc<AtomicUsize>);
impl InFlightGuard {
fn new(counter: Arc<AtomicUsize>) -> Self {
counter.fetch_add(1, Ordering::SeqCst);
Self(counter)
}
}
impl Drop for InFlightGuard {
fn drop(&mut self) {
self.0.fetch_sub(1, Ordering::SeqCst);
}
}
async fn drain_guard_middleware(
State(guard): State<DrainGuard>,
request: axum::extract::Request,
next: Next,
) -> Response {
if guard.readiness.is_draining() {
return (
StatusCode::SERVICE_UNAVAILABLE,
[(RETRY_AFTER, DRAIN_RETRY_AFTER_SECONDS.to_string())],
"draining",
)
.into_response();
}
let _in_flight = InFlightGuard::new(guard.in_flight.clone());
let outcome = AssertUnwindSafe(next.run(request)).catch_unwind().await;
match outcome {
Ok(response) => response,
Err(panic) => {
let detail = panic_message(panic.as_ref());
tracing::error!(panic = %detail, "OoP gear handler panicked");
CanonicalError::internal(format!("gear handler panicked: {detail}"))
.create()
.into_response()
}
}
}
fn panic_message(panic: &(dyn std::any::Any + Send)) -> String {
panic
.downcast_ref::<&str>()
.map(|s| (*s).to_owned())
.or_else(|| panic.downcast_ref::<String>().cloned())
.unwrap_or_else(|| "unknown panic".to_owned())
}
fn layer_gear_router(
gear_router: Router,
drain_guard: DrainGuard,
options: &OopServeOptions,
) -> Router {
let mut gear = gear_router;
if let Some(bearer) = options.bearer_authenticator.clone() {
gear = gear.layer(from_fn_with_state(
Arc::new(bearer),
security_context_middleware::<DynBearerAuthenticator>,
));
}
if let Some(internal) = options.internal_authenticator.clone() {
gear = gear.layer(from_fn_with_state(
Arc::new(internal),
internal_auth_middleware::<DynInternalAuthenticator>,
));
}
gear = gear.layer(from_fn_with_state(drain_guard, drain_guard_middleware));
gear.layer(from_fn(canonical_error_middleware))
}
async fn serve_loop(
listener: tokio::net::TcpListener,
router: Router,
drain_guard: DrainGuard,
sidecar: Option<(tokio::net::TcpListener, Router)>,
drain_timeout: Duration,
cancel: CancellationToken,
) -> anyhow::Result<()> {
let sidecar = match sidecar {
Some((probe_listener, probe_router)) => {
let shutdown = {
let cancel = cancel.clone();
async move { cancel.cancelled().await }
};
Some(tokio::spawn(async move {
if let Err(e) = axum::serve(probe_listener, probe_router)
.with_graceful_shutdown(shutdown)
.await
{
tracing::warn!(error = %e, "OoP probe sidecar server error");
}
}))
}
None => None,
};
let shutdown = {
let cancel = cancel.clone();
let guard = drain_guard.clone();
async move {
cancel.cancelled().await;
guard.begin_drain();
tracing::info!("OoP HTTP server draining (graceful shutdown)");
drain_in_flight(&guard, drain_timeout).await;
}
};
let result = axum::serve(
listener,
router.into_make_service_with_connect_info::<std::net::SocketAddr>(),
)
.with_graceful_shutdown(shutdown)
.await
.map_err(anyhow::Error::from);
if let Some(handle) = sidecar {
handle.abort();
if let Err(e) = handle.await
&& !e.is_cancelled()
{
tracing::warn!(error = %e, "OoP probe sidecar task join error");
}
}
result
}
pub(super) struct OopHttpServer {
readiness: Arc<ReadinessState>,
late: LateRoutes,
drain_guard: DrainGuard,
options: OopServeOptions,
cancel: CancellationToken,
registration_task: Option<tokio::task::JoinHandle<()>>,
serve: tokio::task::JoinHandle<anyhow::Result<()>>,
}
impl OopHttpServer {
pub(super) async fn start(
readiness: Arc<ReadinessState>,
options: OopServeOptions,
cancel: CancellationToken,
) -> anyhow::Result<Self> {
let late = LateRoutes::default();
let drain_guard = DrainGuard::new(Arc::clone(&readiness));
let mut options = options;
let listener = tokio::net::TcpListener::bind(options.listen_addr).await?;
let bound = listener.local_addr()?;
let configured_port = options.listen_addr.port();
if bound.port() != configured_port {
if let Some(rewritten) =
rewrite_advertise_uri_port(&options.advertise_uri, bound.port(), configured_port)
{
tracing::info!(
old = %options.advertise_uri,
new = %rewritten,
"advertise_uri rewritten to use actually bound port"
);
options.advertise_uri = rewritten;
} else {
tracing::info!(
advertise_uri = %options.advertise_uri,
bound_port = bound.port(),
configured_port,
"bound port differs from configured; leaving advertise_uri unchanged"
);
}
}
options.listen_addr = bound;
tracing::info!(
addr = %options.listen_addr,
"OoP HTTP server bound (probes live; gear routes attach after start)"
);
let outer = build_outer_router(Arc::clone(&readiness), late.clone());
let sidecar = if let Some(addr) = options.probe_bind_addr {
let probe_listener = tokio::net::TcpListener::bind(addr).await?;
let bound_probe = probe_listener.local_addr()?;
options.probe_bind_addr = Some(bound_probe);
tracing::info!(addr = %bound_probe, "OoP probe sidecar bound");
Some((
probe_listener,
build_probe_router(Arc::clone(&readiness), late.clone()),
))
} else {
None
};
let serve = tokio::spawn(serve_loop(
listener,
outer,
drain_guard.clone(),
sidecar,
options.drain_timeout,
cancel.clone(),
));
Ok(Self {
readiness,
late,
drain_guard,
options,
cancel,
registration_task: None,
serve,
})
}
pub(super) fn options(&self) -> &OopServeOptions {
&self.options
}
pub(super) fn resolve_bearer_authenticator(&mut self, hub: &crate::ClientHub) {
if self.options.bearer_authenticator.is_some() {
return;
}
if let Ok(auth) = hub.get::<DynBearerAuthenticator>() {
self.options.bearer_authenticator = Some((*auth).clone());
tracing::info!(
gear = %self.options.gear_name,
"tenant-plane authenticator installed (security_context_middleware enabled)"
);
} else {
tracing::warn!(
gear = %self.options.gear_name,
"no tenant-plane authenticator registered in ClientHub; tenant plane not installed"
);
}
}
pub(super) fn attach(&mut self, gear_router: Router, openapi_json: String) {
let openapi_arc: Arc<str> = Arc::from(openapi_json);
let layered = layer_gear_router(gear_router, self.drain_guard.clone(), &self.options);
self.late.publish(layered, Arc::clone(&openapi_arc));
self.readiness.mark_startup_complete();
tracing::info!(gear = %self.options.gear_name, "OoP gear routes attached (now serving)");
let mut registration_info = RegisterInstanceInfo::new(
self.options.gear_name.clone(),
self.options.instance_id.clone(),
)
.with_rest_endpoint(ServiceEndpoint::new(self.options.advertise_uri.clone()))
.with_openapi_spec(openapi_arc.to_string())
.with_labels(self.options.labels.clone());
if let Some(version) = self.options.version.clone() {
registration_info = registration_info.with_version(version);
}
self.registration_task = Some(tokio::spawn(super::oop_registration::presence_loop(
Arc::clone(&self.options.directory),
registration_info,
self.options.heartbeat_interval,
self.cancel.clone(),
)));
}
pub(super) async fn join(mut self) -> anyhow::Result<()> {
let serve_result = match self.serve.await {
Ok(r) => r,
Err(e) => Err(anyhow::anyhow!("OoP serve task join error: {e}")),
};
if let Some(task) = self.registration_task.take() {
task.abort();
}
if let Err(e) = self
.options
.directory
.deregister_instance(&self.options.gear_name, &self.options.instance_id)
.await
{
tracing::warn!(
gear = %self.options.gear_name,
error = %e,
"deregistration from DirectoryService failed on shutdown"
);
} else {
tracing::info!(gear = %self.options.gear_name, "deregistered from DirectoryService");
}
serve_result
}
}
fn rewrite_advertise_uri_port(
advertise_uri: &str,
bound_port: u16,
configured_port: u16,
) -> Option<String> {
let mut url = Url::parse(advertise_uri).ok()?;
if url.port_or_known_default() != Some(configured_port) {
return None;
}
url.set_port(Some(bound_port)).ok()?;
Some(url.to_string())
}
async fn drain_in_flight(guard: &DrainGuard, timeout: Duration) {
let deadline = tokio::time::Instant::now() + timeout;
loop {
let in_flight = guard.in_flight();
if in_flight == 0 {
tracing::info!("OoP drain complete: no in-flight requests");
return;
}
if tokio::time::Instant::now() >= deadline {
tracing::warn!(
in_flight,
timeout_secs = timeout.as_secs(),
"OoP drain timed out with in-flight requests remaining"
);
return;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
#[cfg(test)]
#[cfg_attr(coverage_nightly, coverage(off))]
#[path = "oop_serve_tests.rs"]
mod tests;