Skip to main content

pitchfork_cli/web/
server.rs

1use crate::Result;
2use crate::settings::settings;
3use axum::{
4    Router,
5    body::Body,
6    extract::Request,
7    http::StatusCode,
8    middleware::{self, Next},
9    response::{Redirect, Response},
10    routing::{get, post},
11};
12use std::net::SocketAddr;
13
14use super::routes;
15use super::static_files::{set_static_base, set_static_token, static_handler};
16
17/// API token middleware - rejects requests without valid X-Pitchfork-Token header
18/// when the server is bound to a non-loopback address and a token is configured.
19async fn token_auth(
20    request: Request<Body>,
21    next: Next,
22    expected_token: String,
23) -> Result<Response, StatusCode> {
24    if expected_token.is_empty() {
25        return Ok(next.run(request).await);
26    }
27    let token = request
28        .headers()
29        .get("X-Pitchfork-Token")
30        .and_then(|v| v.to_str().ok())
31        .unwrap_or("");
32    if token != expected_token {
33        let addr: std::borrow::Cow<'_, str> = request
34            .extensions()
35            .get::<axum::extract::ConnectInfo<SocketAddr>>()
36            .map(|a| a.0.to_string().into())
37            .unwrap_or_else(|| "unknown".into());
38        warn!(
39            "API request rejected: invalid or missing X-Pitchfork-Token from {} to {}",
40            addr,
41            request.uri()
42        );
43        return Err(StatusCode::UNAUTHORIZED);
44    }
45    Ok(next.run(request).await)
46}
47
48/// Check if an IP address is loopback (127.0.0.1 or ::1).
49fn is_loopback(addr: &str) -> bool {
50    addr.parse::<SocketAddr>()
51        .map(|a| a.ip().is_loopback())
52        .unwrap_or_else(|_| {
53            // Try parsing as just an IP without port
54            addr.parse::<std::net::IpAddr>()
55                .map(|ip| ip.is_loopback())
56                .unwrap_or(false)
57        })
58}
59/// Generate a random 32-byte hex token (64 characters).
60fn generate_token() -> String {
61    let a = uuid::Uuid::new_v4();
62    let b = uuid::Uuid::new_v4();
63    format!("{}{}", a.simple(), b.simple())
64}
65
66/// Build the API router (no CSRF - SPA uses JSON).
67fn api_router(token: String) -> Router {
68    let token_clone = token.clone();
69    Router::new()
70        .route("/api/stats", get(routes::api::stats::stats))
71        .route("/api/daemons", get(routes::api::daemons::list))
72        .route("/api/daemons/{id}", get(routes::api::daemons::show))
73        .route("/api/daemons/{id}/start", post(routes::api::daemons::start))
74        .route("/api/daemons/{id}/stop", post(routes::api::daemons::stop))
75        .route(
76            "/api/daemons/{id}/restart",
77            post(routes::api::daemons::restart),
78        )
79        .route(
80            "/api/daemons/{id}/enable",
81            post(routes::api::daemons::enable),
82        )
83        .route(
84            "/api/daemons/{id}/disable",
85            post(routes::api::daemons::disable),
86        )
87        .route("/api/logs/{id}/tail", get(routes::api::logs::tail))
88        .route("/api/logs/{id}/loggers", get(routes::api::logs::loggers))
89        .route(
90            "/api/logs/{id}/field-keys",
91            get(routes::api::logs::field_keys),
92        )
93        .route("/api/namespaces", get(routes::api::namespaces::list))
94        .route("/api/namespaces", post(routes::api::namespaces::register))
95        .route(
96            "/api/namespaces/{name}",
97            axum::routing::delete(routes::api::namespaces::remove),
98        )
99        .route("/api/proxies", get(routes::api::proxies::list))
100        .route("/api/projects", get(routes::api::projects::list))
101        .route("/api/projects/{project}", get(routes::api::projects::show))
102        .route(
103            "/api/projects/{project}/{worktree}",
104            get(routes::api::projects::stack),
105        )
106        .route(
107            "/api/processes/{id}/tree",
108            get(routes::api::processes::tree),
109        )
110        .route("/logs/{id}/stream", get(routes::logs::stream_sse))
111        .layer(middleware::from_fn(move |req, next| {
112            let t = token_clone.clone();
113            async move { token_auth(req, next, t).await }
114        }))
115}
116
117/// Bind a `TcpListener` and return `(listener, actual_port)`, trying
118/// `port_attempts` ports starting from `port`.
119async fn try_bind(
120    bind_address: &str,
121    port: u16,
122    port_attempts: u16,
123) -> Result<(tokio::net::TcpListener, u16)> {
124    let ip_addr: std::net::IpAddr = bind_address
125        .parse()
126        .map_err(|e| miette::miette!("Invalid bind address '{}': {}", bind_address, e))?;
127
128    let mut last_error = None;
129    for offset in 0..port_attempts {
130        let try_port = port.saturating_add(offset);
131        let addr = SocketAddr::from((ip_addr, try_port));
132
133        match tokio::net::TcpListener::bind(addr).await {
134            Ok(listener) => {
135                let actual_port = listener
136                    .local_addr()
137                    .map_err(|e| miette::miette!("Failed to inspect bound port: {}", e))?;
138                return Ok((listener, actual_port.port()));
139            }
140            Err(e) => {
141                debug!("Port {try_port} unavailable: {e}");
142                last_error = Some(e);
143            }
144        }
145    }
146
147    Err(miette::miette!(
148        "Failed to bind: tried ports {}-{}, all in use. Last error: {}",
149        port,
150        port.saturating_add(port_attempts - 1),
151        last_error.map(|e| e.to_string()).unwrap_or_default()
152    ))
153}
154
155pub async fn serve(port: u16, web_path: Option<String>) -> Result<()> {
156    let base_path = super::normalize_base_path(web_path.as_deref())?;
157    super::BASE_PATH
158        .set(base_path.clone())
159        .expect("BASE_PATH already set; serve() must only be called once per process");
160    let s = settings();
161    let bind_address = &s.web.bind_address;
162    let port_attempts: u16 = u16::try_from(s.web.port_attempts)
163        .unwrap_or_else(|_| {
164            warn!(
165                "web.port_attempts value {} is out of range (1-65535), clamping to 10",
166                s.web.port_attempts
167            );
168            10
169        })
170        .max(1);
171
172    // Determine token: use configured token, or auto-generate one if binding to non-loopback
173    let mut token = s.api.token.clone();
174    if token.is_empty() && !is_loopback(bind_address) {
175        token = generate_token();
176        info!(
177            "Web UI bound to non-loopback address {}. Auto-generated API token: {}",
178            bind_address, token
179        );
180        // Also print to stderr so it's visible even with log level filtering
181        eprintln!("pitchfork API security token (auto-generated): {}", token);
182    }
183
184    set_static_token(token.clone());
185    set_static_base(base_path.clone());
186
187    let inner = api_router(token.clone()).fallback(static_handler);
188
189    let app = if base_path.is_empty() {
190        inner
191    } else {
192        let redirect_target = format!("{base_path}/");
193        Router::new()
194            .route(
195                "/",
196                get(move || async move { Redirect::temporary(&redirect_target) }),
197            )
198            .nest(&base_path, inner)
199    };
200
201    let (listener, actual_port) = try_bind(bind_address, port, port_attempts).await?;
202    let _ = super::WEB_PORT.set(actual_port);
203    let actual_addr = listener.local_addr().unwrap();
204    let url_host = match actual_addr.ip() {
205        ip if ip.is_unspecified() => "localhost".to_string(),
206        std::net::IpAddr::V6(ip) => format!("[{ip}]"),
207        std::net::IpAddr::V4(ip) => ip.to_string(),
208    };
209    // The scheme is hardcoded because the web server only speaks plain HTTP.
210    // If TLS support is ever added, this must pick the scheme accordingly or
211    // the URL reported by `supervisor status` will be wrong.
212    let _ = super::WEB_URL.set(format!("http://{url_host}:{actual_port}{base_path}"));
213
214    info!("Web UI listening on http://{actual_addr}");
215
216    axum::serve(listener, app)
217        .await
218        .map_err(|e| miette::miette!("Web server error: {}", e))
219}
220
221/// Serve the API on a dedicated port, separate from the web UI.
222/// Called by the supervisor when `settings.api.bind_port` is configured.
223pub async fn serve_api(port: u16, _web_path: Option<String>) -> Result<()> {
224    let s = settings();
225    let bind_address = &s.api.bind_address;
226    let port_attempts: u16 = u16::try_from(s.api.port_attempts)
227        .unwrap_or_else(|_| {
228            warn!(
229                "api.port_attempts value {} is out of range (1-65535), clamping to 10",
230                s.api.port_attempts
231            );
232            10
233        })
234        .max(1);
235
236    // Determine token for standalone API server
237    let mut token = s.api.token.clone();
238    if token.is_empty() && !is_loopback(bind_address) {
239        token = generate_token();
240        info!(
241            "API server bound to non-loopback address {}. Auto-generated API token: {}",
242            bind_address, token
243        );
244        eprintln!("pitchfork API security token (auto-generated): {}", token);
245    }
246
247    let app = api_router(token);
248
249    let (listener, _actual_port) = try_bind(bind_address, port, port_attempts).await?;
250    let actual_addr = listener.local_addr().unwrap();
251    info!("API server listening on http://{actual_addr}");
252
253    axum::serve(listener, app)
254        .await
255        .map_err(|e| miette::miette!("API server error: {}", e))
256}