Skip to main content

bestool_alertd/
http_server.rs

1//! HTTP server for alertd daemon control and metrics.
2
3use std::{collections::HashMap, sync::Arc, time::Duration};
4
5use axum::{
6	Router,
7	routing::{get, post},
8};
9use jiff::Timestamp;
10use tower_http::trace::{DefaultMakeSpan, DefaultOnResponse, TraceLayer};
11use tracing::{Level, error, info, warn};
12
13use crate::{
14	context::InternalContext,
15	daemon::DaemonControl,
16	tasks::{BackgroundTask, TaskEndpointHandler},
17};
18
19mod endpoints;
20mod metrics_render;
21mod state;
22#[cfg(test)]
23mod test_utils;
24mod types;
25
26pub use endpoints::*;
27pub use state::ServerState;
28pub use types::*;
29
30#[expect(
31	clippy::too_many_arguments,
32	reason = "daemon wiring; each shared resource is threaded in explicitly"
33)]
34pub async fn start_server(
35	internal_context: Arc<InternalContext>,
36	addrs: Vec<std::net::SocketAddr>,
37	watchdog_timeout: Option<Duration>,
38	background_tasks: &[Arc<dyn BackgroundTask>],
39	control: DaemonControl,
40	backups: Option<Arc<crate::BackupRegistry>>,
41	metrics: Option<crate::doctor::DoctorMetricsHandle>,
42	binary_version: String,
43) {
44	let started_at = Timestamp::now();
45	let pid = std::process::id();
46
47	let task_endpoints = collect_task_endpoints(background_tasks);
48
49	let state = ServerState {
50		started_at,
51		pid,
52		binary_version,
53		internal_context,
54		watchdog_timeout,
55		task_endpoints: Arc::new(task_endpoints),
56		control,
57		backups,
58		metrics,
59	};
60
61	let app = Router::new()
62		.route("/", get(handle_index))
63		.route("/metrics", get(handle_metrics))
64		.route("/status", get(handle_status))
65		.route("/health", get(handle_health))
66		.route("/reload", post(handle_reload))
67		.route("/restart", post(handle_restart))
68		.route("/tasks/{task}/{endpoint}", get(handle_task_endpoint))
69		.layer(
70			TraceLayer::new_for_http()
71				.make_span_with(
72					DefaultMakeSpan::new()
73						.level(Level::INFO)
74						.include_headers(false),
75				)
76				.on_request(|request: &axum::http::Request<_>, _span: &tracing::Span| {
77					info!(
78						method = %request.method(),
79						uri = %request.uri(),
80						"HTTP request"
81					);
82				})
83				.on_response(
84					DefaultOnResponse::new()
85						.level(Level::INFO)
86						.include_headers(false),
87				),
88		)
89		.with_state(Arc::new(state));
90
91	// Use default if no addresses provided
92	let addrs_to_try = if addrs.is_empty() {
93		vec![
94			"[::1]:8271".parse().unwrap(),
95			"127.0.0.1:8271".parse().unwrap(),
96		]
97	} else {
98		addrs
99	};
100
101	let mut listener = None;
102	let mut last_error = None;
103
104	// Try each address in order until one succeeds
105	for addr in &addrs_to_try {
106		match tokio::net::TcpListener::bind(addr).await {
107			Ok(l) => {
108				info!("HTTP server listening on http://{}", addr);
109				listener = Some(l);
110				break;
111			}
112			Err(e) => {
113				warn!("failed to bind HTTP server to {}: {}", addr, e);
114				last_error = Some(e);
115			}
116		}
117	}
118
119	let listener = match listener {
120		Some(l) => l,
121		None => {
122			if let Some(e) = last_error {
123				warn!("failed to bind HTTP server to any address: {}", e);
124			} else {
125				warn!("no addresses provided for HTTP server");
126			}
127			warn!("waiting 10 seconds before continuing without");
128			warn!("use --no-server to disable the HTTP server and this warning");
129
130			tokio::time::sleep(tokio::time::Duration::from_secs(10)).await;
131
132			info!("continuing without HTTP server");
133			return;
134		}
135	};
136
137	if let Err(e) = axum::serve(listener, app).await {
138		error!("HTTP server error: {}", e);
139	}
140}
141
142fn collect_task_endpoints(
143	tasks: &[Arc<dyn BackgroundTask>],
144) -> HashMap<(String, String), TaskEndpointHandler> {
145	let mut map = HashMap::new();
146	for task in tasks {
147		let task_name = task.name();
148		for endpoint in task.http_endpoints() {
149			let key = (task_name.to_string(), endpoint.name.to_string());
150			if map.contains_key(&key) {
151				warn!(
152					task = task_name,
153					endpoint = endpoint.name,
154					"duplicate task endpoint name; later registration wins"
155				);
156			}
157			info!(
158				task = task_name,
159				endpoint = endpoint.name,
160				path = %format!("/tasks/{task_name}/{}", endpoint.name),
161				"mounting task endpoint"
162			);
163			map.insert(key, endpoint.handler);
164		}
165	}
166	map
167}