Skip to main content

koan_server/graphql/
server.rs

1use std::path::PathBuf;
2use std::sync::Arc;
3
4use crossbeam_channel::Sender;
5use koan_core::audio::viz::VizSnapshot;
6use koan_core::auth::{self, parse_duration_secs};
7use koan_core::config::Config;
8use koan_core::db::pool::Pool;
9use koan_core::player::commands::PlayerCommand;
10use koan_core::player::state::SharedPlayerState;
11
12use super::{KoanSchema, build_schema};
13use crate::auth::AuthUser;
14use crate::auth::middleware::{AuthState, auth_middleware};
15use crate::auth::routes::{AuthRouteState, RateLimiter, auth_router};
16
17// ---------------------------------------------------------------------------
18// `koan --headless` entry point (standalone headless server)
19// ---------------------------------------------------------------------------
20
21pub fn cmd_serve(
22    port: Option<u16>,
23    bind: Option<std::net::IpAddr>,
24    subsonic_port: Option<u16>,
25    playground: bool,
26) {
27    use koan_core::player::Player;
28
29    // Open, and migrate, before binding: a database this build cannot read,
30    // such as one a newer koan has migrated, stops it here with the reason
31    // rather than serving errors.
32    if let Err(e) = koan_core::db::connection::Database::open_default() {
33        log::error!("cannot open the database: {e}");
34        eprintln!("koan: cannot open the database: {e}");
35        std::process::exit(1);
36    }
37    let db_path = koan_core::config::db_path();
38    let pool = Arc::new(Pool::new(db_path.clone()));
39    koan_core::db::pool::on_outdated(|| OUTDATED.cancel());
40
41    let (state, _timeline, _viz, cmd_tx) = Player::spawn();
42
43    // A server has no one to press "scan": index at start and whenever the
44    // library folders change, as the macOS app does.
45    let watched = db_path.clone();
46    koan_core::helpers::spawn_library_watch(db_path, move |running| {
47        // A scan that just finished may have brought in an album someone asked
48        // to have queued when it arrived.
49        if !running {
50            crate::clients::fulfil_from(&watched);
51            if let Ok(db) = koan_core::db::connection::Database::open_existing(&watched) {
52                crate::clients::changed_if_library_moved(&db.conn);
53            }
54        }
55    });
56
57    if let Err(e) = run_api_blocking(ApiServerOpts {
58        state,
59        cmd_tx,
60        pool,
61        port,
62        bind,
63        subsonic_port,
64        playground,
65        viz: None, // headless — no viz analyzer
66        headless: true,
67    }) {
68        eprintln!("koan: {}", e);
69        std::process::exit(1);
70    }
71}
72
73// ---------------------------------------------------------------------------
74// Shared API server logic — used by both headless and TUI+API modes
75// ---------------------------------------------------------------------------
76
77/// Ceiling on a single GraphQL query. Anything longer than this —
78/// a library scan, a remote sync — runs as a job instead.
79const REQUEST_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
80
81/// Queries executing at once. Resolvers do their SQLite and HTTP work on the
82/// blocking pool, so this bounds concurrent work rather than protecting the
83/// runtime's workers from it.
84const MAX_INFLIGHT_QUERIES: usize = 64;
85
86/// Largest query body accepted. async-graphql reads the whole body before
87/// parsing it, with no limit of its own. A query is text; even a playlist of
88/// thousands of ids is a fraction of this.
89const MAX_QUERY_BODY: usize = 2 << 20;
90
91/// Timeout, panic catch, body limit and load shed for the query route.
92///
93/// Not applied to `/graphql/ws`: a subscription is meant to outlive any request
94/// timeout.
95fn load_perimeter<S>(router: axum::Router<S>) -> axum::Router<S>
96where
97    S: Clone + Send + Sync + 'static,
98{
99    router
100        // Innermost so it is inside the timeout: a panicking resolver becomes a
101        // 500 rather than a silently dropped connection.
102        .layer(tower_http::catch_panic::CatchPanicLayer::new())
103        .layer(tower_http::limit::RequestBodyLimitLayer::new(
104            MAX_QUERY_BODY,
105        ))
106        .layer(tower_http::timeout::TimeoutLayer::with_status_code(
107            axum::http::StatusCode::REQUEST_TIMEOUT,
108            REQUEST_TIMEOUT,
109        ))
110        // Shed rather than queue. A concurrency limit on its own parks callers
111        // on a semaphore, so an overloaded server answers every client slowly
112        // instead of telling the surplus to come back.
113        .layer(
114            tower::ServiceBuilder::new()
115                .layer(axum::error_handling::HandleErrorLayer::new(
116                    |err: tower::BoxError| async move {
117                        if err.is::<tower::load_shed::error::Overloaded>() {
118                            (
119                                axum::http::StatusCode::SERVICE_UNAVAILABLE,
120                                "server at capacity",
121                            )
122                        } else {
123                            (
124                                axum::http::StatusCode::INTERNAL_SERVER_ERROR,
125                                "internal error",
126                            )
127                        }
128                    },
129                ))
130                .load_shed()
131                .concurrency_limit(MAX_INFLIGHT_QUERIES),
132        )
133}
134
135/// Options for the API server — avoids too-many-arguments.
136pub struct ApiServerOpts {
137    pub state: Arc<SharedPlayerState>,
138    pub cmd_tx: Sender<PlayerCommand>,
139    /// Shared by every route that reads the library outside GraphQL, and by
140    /// the MCP listener when there is one.
141    pub pool: Arc<Pool>,
142    pub port: Option<u16>,
143    pub bind: Option<std::net::IpAddr>,
144    pub subsonic_port: Option<u16>,
145    pub playground: bool,
146    pub viz: Option<Arc<VizSnapshot>>,
147    /// No TUI or app: the player here is heard by nobody.
148    pub headless: bool,
149}
150
151/// Run the GraphQL (+ optional Subsonic) API server, blocking the current thread.
152/// Called from `cmd_serve` (headless) and `start_api_background` (TUI companion).
153///
154/// `Err` means the server refused to start on a misconfiguration the caller has
155/// to surface — never a silent downgrade to an unauthenticated server.
156fn run_api_blocking(opts: ApiServerOpts) -> Result<(), String> {
157    let ApiServerOpts {
158        state,
159        cmd_tx,
160        pool,
161        port,
162        bind,
163        subsonic_port,
164        playground,
165        viz,
166        headless,
167    } = opts;
168    use axum::routing::{get, post};
169
170    let cfg = Config::load().unwrap_or_default();
171    let port = port.unwrap_or(cfg.graphql.port);
172    let bind = bind.unwrap_or(cfg.graphql.bind);
173    let subsonic_port = subsonic_port.or(cfg.subsonic.port);
174    let playground_enabled = playground || cfg.graphql.playground;
175    let auth_enabled = cfg.graphql.auth_enabled;
176
177    // Load or generate Ed25519 keypair for JWT signing. A server's first start
178    // makes its own: it is this server's signing key and nothing else, a fresh
179    // one invalidates no token (there can be none yet), and a server in a
180    // container has no terminal to run `koan auth setup` in before it starts.
181    // Accounts are still created deliberately; until one exists, nothing signs in.
182    let (private_pem, public_pem) = if auth_enabled {
183        let kp = auth::load_or_generate_keypair().map_err(|e| {
184            format!(
185                "auth_enabled = true but the keypair could not be loaded or created: {}",
186                e
187            )
188        })?;
189        // An empty or truncated key file would otherwise leave every request an
190        // unauthenticated admin, which is the opposite of what was asked for.
191        if kp.0.is_empty() || kp.1.is_empty() {
192            return Err("auth_enabled = true but the keypair files are empty. \
193                 Run `koan auth regenerate-keys`."
194                .into());
195        }
196        kp
197    } else {
198        // When auth is disabled, we still need dummy keys for the route state
199        // (routes exist but won't be hit by middleware). Generate if available.
200        auth::load_or_generate_keypair().unwrap_or_default()
201    };
202
203    // What once encrypted account passwords for Subsonic token auth. Nothing
204    // is encrypted with it any more; a copy left on disk is a liability.
205    let _ = std::fs::remove_file(auth::keypair_dir().join("subsonic.key"));
206
207    let access_ttl = parse_duration_secs(&cfg.graphql.access_token_ttl).unwrap_or(900);
208    let refresh_ttl = parse_duration_secs(&cfg.graphql.refresh_token_ttl).unwrap_or(2_592_000);
209
210    // Process-scoped introspection key for playground access. It is a bearer
211    // credential compared verbatim, so it has to be full-entropy random — a
212    // UUID would leak the server start time and cut the guessable space.
213    let introspection_key = if playground_enabled && auth_enabled {
214        Some(Arc::new(auth::random_token().map_err(|e| {
215            format!("failed to generate introspection key: {}", e)
216        })?))
217    } else {
218        None
219    };
220
221    let auth_state = AuthState {
222        public_pem: Arc::new(public_pem.clone()),
223        auth_enabled,
224        introspection_key: introspection_key.clone(),
225        pool: pool.clone(),
226    };
227
228    crate::auth::set_signing_keys(Arc::new(private_pem.clone()), Arc::new(public_pem.clone()));
229
230    // One for every door a password comes through, so its limits are shared.
231    let users = Arc::new(crate::auth::password::PasswordVerifier::new(pool.clone()));
232    let auth_route_state = AuthRouteState {
233        users: users.clone(),
234        pool: pool.clone(),
235        private_pem: Arc::new(private_pem),
236        public_pem: Arc::new(public_pem),
237        access_ttl_secs: access_ttl,
238        refresh_ttl_secs: refresh_ttl,
239        cookie_secure: cfg.graphql.cookie_secure,
240        login_limiter: Arc::new(RateLimiter::default()),
241    };
242
243    // Plays accounts forward to ListenBrainz; it sleeps while none do.
244    koan_core::scrobbling::start(pool.path().to_path_buf());
245
246    let shutdown = tokio_util::sync::CancellationToken::new();
247    let mcp_routes = crate::mcp::router(
248        state.clone(),
249        cmd_tx.clone(),
250        auth_state.clone(),
251        cfg.sharing.public_url.clone(),
252        headless,
253        shutdown.clone(),
254    );
255    let schema = build_schema(state, cmd_tx, pool.clone(), viz);
256
257    if auth_enabled {
258        log::info!(
259            "Auth enabled (Ed25519 JWT, access TTL {}s, refresh TTL {}s)",
260            access_ttl,
261            refresh_ttl
262        );
263        let no_accounts = pool
264            .get()
265            .ok()
266            .and_then(|db| koan_core::db::queries::auth::has_users(&db.conn).ok())
267            == Some(false);
268        if no_accounts && cfg.graphql.setup_wizard {
269            log::warn!(
270                "No accounts yet: whoever opens /setup first becomes the admin. \
271                 Set graphql.setup_wizard = false to make it with `koan auth setup` instead."
272            );
273        }
274    } else {
275        log::info!("Auth disabled — all requests treated as admin");
276    }
277
278    let browser_policy = Arc::new(BrowserPolicy {
279        origins: cfg.graphql.cors_origins.clone(),
280        hosts: cfg.graphql.allowed_hosts.clone(),
281    });
282
283    if cfg.graphql.cors_origins.is_empty() {
284        log::info!("CORS: no origins configured — browsers get no cross-origin access");
285    }
286
287    let proxy_auth = crate::ui::ProxyAuth::from_config(
288        &cfg.graphql.proxy_auth_header,
289        &cfg.graphql.proxy_auth_from,
290    )?;
291
292    let rt = tokio::runtime::Runtime::new().expect("failed to create tokio runtime");
293    rt.block_on(async {
294        // GraphQL routes — protected by auth middleware.
295        //
296        // The query route carries the load perimeter; the WebSocket route does
297        // not, because a subscription is meant to outlive any request timeout.
298        let query_route = load_perimeter(axum::Router::new().route("/graphql", post(graphql_handler)));
299
300        let gql_app = axum::Router::new()
301            .merge(query_route)
302            .route("/graphql/ws", get(graphql_ws_handler))
303            .layer(axum::middleware::from_fn_with_state(
304                auth_state.clone(),
305                auth_middleware,
306            ))
307            // Runs before auth: a rejected request should never reach a
308            // credential check, let alone execute.
309            .layer(axum::middleware::from_fn_with_state(
310                browser_policy.clone(),
311                browser_guard,
312            ))
313            .with_state(schema);
314
315        // The web UI checks its own session (a page load without one is sent
316        // to sign in rather than refused), so it sits outside the GraphQL auth
317        // layer. Built before the auth routes take their state.
318        let covers = Arc::new(crate::covers::Covers::in_config_dir());
319        let ui_routes = crate::ui::router(
320            pool.clone(),
321            auth_route_state.clone(),
322            auth_enabled,
323            cfg.graphql.setup_wizard,
324            covers.clone(),
325            cfg.sharing.public_url.clone(),
326            cfg.mcp.redirect_hosts.clone(),
327            proxy_auth,
328        );
329
330        // Auth routes — always accessible (no auth middleware).
331        let auth_app = auth_router(auth_route_state);
332
333        // CORS. An empty origin list emits no `Access-Control-Allow-Origin` at
334        // all; a wildcard would let any web page read this library.
335        let origins: Vec<axum::http::HeaderValue> = cfg
336            .graphql
337            .cors_origins
338            .iter()
339            .filter_map(|o| o.parse().ok())
340            .collect();
341        let cors = tower_http::cors::CorsLayer::new()
342            .allow_origin(origins)
343            .allow_methods([
344                axum::http::Method::GET,
345                axum::http::Method::POST,
346                axum::http::Method::OPTIONS,
347            ])
348            .allow_headers([
349                axum::http::header::AUTHORIZATION,
350                axum::http::header::CONTENT_TYPE,
351                axum::http::HeaderName::from_static("x-introspection-key"),
352            ])
353            .allow_credentials(true);
354
355
356        // Public by design, so outside the auth layers: each route answers for
357        // one share's own tracks, or the one cover a notification's signed
358        // link names, and nothing else. The Host guard still applies.
359        let share_routes = crate::share::router(
360            pool.clone(),
361            cfg.sharing.public_url.clone(),
362            covers.clone(),
363        )
364        .merge(crate::push::router(pool.clone(), covers.clone()));
365        // Subsonic on the GraphQL port whenever `[subsonic]` is enabled: the
366        // remote TUI bridge builds its stream URL off the GraphQL base. Built
367        // once and cloned for the dedicated listener, since each build reads
368        // the config from disk.
369        let subsonic_merged = crate::subsonic::subsonic_router(pool, covers, users);
370        let subsonic_on_main = subsonic_merged.is_some();
371        let subsonic_dedicated = subsonic_merged.clone();
372
373        let mut app = auth_app
374            .merge(gql_app)
375            .merge(share_routes)
376            .merge(ui_routes)
377            .merge(mcp_routes);
378        if let Some(sub) = subsonic_merged {
379            app = app.merge(sub);
380        }
381        if playground_enabled {
382            app = app.route(
383                "/graphql",
384                get(graphql_playground).with_state(introspection_key.clone()),
385            );
386        }
387        // Outermost: a request whose `Host` we do not recognise is refused
388        // before anything else looks at it. Without this a DNS-rebinding page
389        // reaches the API as same-origin and CORS stops mattering.
390        let app = app.layer(cors).layer(axum::middleware::from_fn_with_state(
391            browser_policy.clone(),
392            host_guard,
393        ));
394
395        // Build playground URL with introspection key.
396        let playground_url = if playground_enabled {
397            if let Some(ref key) = introspection_key {
398                format!("http://{}:{}/graphql?introspection-key={}", bind, port, key)
399            } else {
400                format!("http://{}:{}/graphql", bind, port)
401            }
402        } else {
403            format!("http://{}:{}/graphql", bind, port)
404        };
405
406        let gql_addr = std::net::SocketAddr::new(bind, port);
407
408        let gql_listener = match tokio::net::TcpListener::bind(gql_addr).await {
409            Ok(l) => {
410                log::info!("GraphQL API on http://{}:{}/graphql", bind, port);
411                if subsonic_on_main {
412                    log::info!("Subsonic REST on http://{}:{}/rest/", bind, port);
413                }
414                if playground_enabled {
415                    log::info!("GraphiQL: {}", playground_url);
416                    // Open browser on macOS/Linux.
417                    #[cfg(target_os = "macos")]
418                    let _ = std::process::Command::new("open").arg(&playground_url).spawn();
419                    #[cfg(target_os = "linux")]
420                    let _ = std::process::Command::new("xdg-open").arg(&playground_url).spawn();
421                }
422                l
423            }
424            Err(e) => {
425                return Err(format!(
426                    "failed to bind GraphQL port {port} — {e} (another instance running?)"
427                ));
428            }
429        };
430        let gql_server = axum::serve(
431            gql_listener,
432            app.into_make_service_with_connect_info::<std::net::SocketAddr>(),
433        )
434        .with_graceful_shutdown(async move {
435            shutdown_signal().await;
436            shutdown.cancel();
437        });
438
439        // `--subsonic <port>` set to something other than the GraphQL port adds
440        // a dedicated Subsonic listener.
441        let extra_sub_port = subsonic_port.filter(|p| *p != port);
442        if let Some(sub_port) = extra_sub_port
443            && let Some(sub_app) = subsonic_dedicated
444        {
445            let sub_addr = std::net::SocketAddr::new(bind, sub_port);
446            match tokio::net::TcpListener::bind(sub_addr).await {
447                Ok(sub_listener) => {
448                    log::info!(
449                        "Subsonic REST also on http://{}:{}/rest/ (dedicated port)",
450                        bind,
451                        sub_port,
452                    );
453                    // With connection info, as the main listener: the sign-in
454                    // throttle keys on the client's address.
455                    let sub_server = axum::serve(
456                        sub_listener,
457                        sub_app.into_make_service_with_connect_info::<std::net::SocketAddr>(),
458                    )
459                    .with_graceful_shutdown(shutdown_signal());
460
461                    tokio::select! {
462                        r = gql_server => { if let Err(e) = r { log::error!("GraphQL server error: {e}"); } },
463                        r = sub_server => { if let Err(e) = r { log::error!("Subsonic server error: {e}"); } },
464                        _ = drain_deadline() => log::info!("shutting down with connections still open"),
465                    }
466                    return Ok(());
467                }
468                Err(e) => {
469                    log::warn!(
470                        "Dedicated Subsonic port {} unavailable — {}. Mounted on GraphQL port only.",
471                        sub_port,
472                        e,
473                    );
474                }
475            }
476        }
477
478        tokio::select! {
479            r = gql_server => { if let Err(e) = r { log::error!("GraphQL server error: {e}"); } },
480            _ = drain_deadline() => log::info!("shutting down with connections still open"),
481        }
482        Ok(())
483    })
484}
485
486/// Start the API server on the current thread (blocks forever).
487/// Called from a background thread when TUI mode has API enabled.
488pub fn start_api_background(
489    state: Arc<SharedPlayerState>,
490    cmd_tx: Sender<PlayerCommand>,
491    db_path: PathBuf,
492    port: Option<u16>,
493    bind: Option<std::net::IpAddr>,
494    subsonic_port: Option<u16>,
495    playground: bool,
496) {
497    // Runs on a spawned thread in TUI mode, where a panic would take the API
498    // down with nothing on screen to say so.
499    if let Err(e) = run_api_blocking(ApiServerOpts {
500        state,
501        cmd_tx,
502        pool: Arc::new(Pool::new(db_path)),
503        port,
504        bind,
505        subsonic_port,
506        playground,
507        viz: None,
508        headless: false,
509    }) {
510        log::error!("API server not started: {}", e);
511    }
512}
513
514// ---------------------------------------------------------------------------
515// Browser perimeter
516// ---------------------------------------------------------------------------
517
518/// What this server will answer to when the caller is a browser.
519///
520/// Two separate questions: which `Host` values name this server (DNS rebinding),
521/// and which `Origin` values may talk to it (CSRF, cross-site WebSockets).
522pub(crate) struct BrowserPolicy {
523    origins: Vec<String>,
524    hosts: Vec<String>,
525}
526
527impl BrowserPolicy {
528    fn host_allowed(&self, host: &str) -> bool {
529        if self.hosts.iter().any(|h| h.eq_ignore_ascii_case(host)) {
530            return true;
531        }
532        let bare = strip_port(host);
533        if self.hosts.iter().any(|h| h.eq_ignore_ascii_case(bare)) {
534            return true;
535        }
536        // A rebinding attack needs a name it controls; literals and localhost
537        // resolve to this machine by definition.
538        bare.eq_ignore_ascii_case("localhost") || bare.parse::<std::net::IpAddr>().is_ok()
539    }
540
541    /// An origin is allowed if it is configured, or if it is simply this server
542    /// talking to itself — which is what the bundled playground does.
543    fn origin_allowed(&self, origin: &str, host: Option<&str>) -> bool {
544        if self.origins.iter().any(|o| o == origin) {
545            return true;
546        }
547        match (origin.split_once("://"), host) {
548            (Some((_, authority)), Some(host)) => authority.eq_ignore_ascii_case(host),
549            _ => false,
550        }
551    }
552}
553
554/// `example.com:4000` -> `example.com`, `[::1]:4000` -> `::1`.
555fn strip_port(host: &str) -> &str {
556    if let Some(rest) = host.strip_prefix('[') {
557        return rest.split(']').next().unwrap_or(rest);
558    }
559    match host.rsplit_once(':') {
560        Some((h, port)) if !port.is_empty() && port.bytes().all(|b| b.is_ascii_digit()) => h,
561        _ => host,
562    }
563}
564
565fn header_str(request: &axum::extract::Request, name: axum::http::HeaderName) -> Option<&str> {
566    request.headers().get(name).and_then(|v| v.to_str().ok())
567}
568
569/// Reject requests carrying an unrecognised `Host`.
570async fn host_guard(
571    axum::extract::State(policy): axum::extract::State<Arc<BrowserPolicy>>,
572    request: axum::extract::Request,
573    next: axum::middleware::Next,
574) -> axum::response::Response {
575    use axum::response::IntoResponse;
576
577    // No `Host` at all means no browser: only HTTP/1.0 and raw tooling omit it,
578    // and neither can be steered by an attacker page.
579    let host = header_str(&request, axum::http::header::HOST)
580        .map(str::to_owned)
581        .or_else(|| request.uri().host().map(str::to_owned));
582
583    if let Some(ref host) = host
584        && !policy.host_allowed(host)
585    {
586        log::warn!("rejected request for unrecognised Host: {}", host);
587        return (axum::http::StatusCode::FORBIDDEN, "host not allowed").into_response();
588    }
589
590    next.run(request).await
591}
592
593/// Reject cross-site GraphQL traffic.
594///
595/// Two holes, one guard. A WebSocket handshake is exempt from CORS entirely, so
596/// a foreign page can open `/graphql/ws`, have the browser attach the session
597/// cookie, and read every response. And a POST whose content type is
598/// CORS-safelisted (`text/plain`) is sent without a preflight, yet
599/// async-graphql parses it as JSON regardless — so the mutation lands even
600/// though the reply is unreadable.
601async fn browser_guard(
602    axum::extract::State(policy): axum::extract::State<Arc<BrowserPolicy>>,
603    request: axum::extract::Request,
604    next: axum::middleware::Next,
605) -> axum::response::Response {
606    use axum::response::IntoResponse;
607
608    let host = header_str(&request, axum::http::header::HOST).map(str::to_owned);
609    // No `Origin` means a non-browser client, which CSRF cannot reach.
610    if let Some(origin) = header_str(&request, axum::http::header::ORIGIN)
611        && !policy.origin_allowed(origin, host.as_deref())
612    {
613        log::warn!(
614            "rejected GraphQL request from disallowed Origin: {}",
615            origin
616        );
617        return (axum::http::StatusCode::FORBIDDEN, "origin not allowed").into_response();
618    }
619
620    if request.method() == axum::http::Method::POST && !is_graphql_content_type(&request) {
621        return (
622            axum::http::StatusCode::UNSUPPORTED_MEDIA_TYPE,
623            "content type must be application/json or application/graphql",
624        )
625            .into_response();
626    }
627
628    next.run(request).await
629}
630
631fn is_graphql_content_type(request: &axum::extract::Request) -> bool {
632    header_str(request, axum::http::header::CONTENT_TYPE).is_some_and(|ct| {
633        let ct = ct.trim().to_ascii_lowercase();
634        ct.starts_with("application/json") || ct.starts_with("application/graphql")
635    })
636}
637
638/// Cancelled when the database turns out to have been upgraded by a newer
639/// koan, as the one replacing this server does when it starts: this build
640/// can no longer serve it, so it drains and exits as on SIGTERM.
641static OUTDATED: std::sync::LazyLock<tokio_util::sync::CancellationToken> =
642    std::sync::LazyLock::new(tokio_util::sync::CancellationToken::new);
643
644/// Resolves on SIGINT, SIGTERM where there is one, or the database being
645/// upgraded past this build. SIGTERM is how a service manager stops a
646/// process, and as PID 1 in a container an unhandled one is dropped by the
647/// kernel, leaving the server running until it is killed.
648async fn shutdown_signal() {
649    #[cfg(unix)]
650    {
651        use tokio::signal::unix::{SignalKind, signal};
652        let mut terminate = signal(SignalKind::terminate()).expect("failed to listen for SIGTERM");
653        tokio::select! {
654            r = tokio::signal::ctrl_c() => r.expect("failed to listen for ctrl+c"),
655            _ = terminate.recv() => {}
656            _ = OUTDATED.cancelled() => {}
657        }
658    }
659    #[cfg(not(unix))]
660    tokio::select! {
661        r = tokio::signal::ctrl_c() => r.expect("failed to listen for ctrl+c"),
662        _ = OUTDATED.cancelled() => {}
663    }
664}
665
666/// How long shutdown waits for open connections before giving up on them.
667const DRAIN: std::time::Duration = std::time::Duration::from_secs(10);
668
669/// Resolves `DRAIN` after a shutdown signal. Graceful shutdown waits for every
670/// connection to close, and subscription websockets and audio streams never
671/// close on their own, so without a deadline the wait ends only in a kill.
672async fn drain_deadline() {
673    shutdown_signal().await;
674    tokio::time::sleep(DRAIN).await;
675}
676
677async fn graphql_handler(
678    axum::Extension(user): axum::Extension<AuthUser>,
679    axum::extract::State(schema): axum::extract::State<KoanSchema>,
680    headers: axum::http::HeaderMap,
681    req: async_graphql_axum::GraphQLRequest,
682) -> async_graphql_axum::GraphQLResponse {
683    let mut request = req.into_inner();
684    if let Some(origin) = crate::origin::origin(&headers, None) {
685        request = request.data(super::RequestOrigin(origin));
686    }
687    // The auth middleware always injects AuthUser (anonymous_admin when auth is
688    // disabled, or a real user when auth is enabled). No fallback needed here.
689    request = request.data(user);
690    schema.execute(request).await.into()
691}
692
693/// Subscriptions over a WebSocket, as the account the upgrade request
694/// authenticated. Closed when that account changes or its token lapses (see
695/// `Lease`); without a lease, as when auth is off, it stays open.
696async fn graphql_ws_handler(
697    axum::Extension(user): axum::Extension<AuthUser>,
698    lease: Option<axum::Extension<crate::auth::Lease>>,
699    axum::extract::State(schema): axum::extract::State<KoanSchema>,
700    protocol: async_graphql_axum::GraphQLProtocol,
701    websocket: axum::extract::WebSocketUpgrade,
702) -> axum::response::Response {
703    use axum::extract::ws::{CloseFrame, Message, close_code};
704    use futures_util::{SinkExt, StreamExt};
705    websocket
706        .protocols(async_graphql::http::ALL_WEBSOCKET_PROTOCOLS)
707        .on_upgrade(move |socket| async move {
708            let (mut sink, stream) = socket.split();
709            let serve = async_graphql_axum::GraphQLWebSocket::new_with_pair(
710                &mut sink, stream, schema, protocol,
711            )
712            .on_connection_init(move |_| async move {
713                let mut data = async_graphql::Data::default();
714                data.insert(user);
715                Ok(data)
716            })
717            .serve();
718            let ended = async {
719                match lease {
720                    Some(axum::Extension(lease)) => lease.ended().await,
721                    None => std::future::pending().await,
722                }
723            };
724            let ended = tokio::select! {
725                _ = serve => false,
726                _ = ended => true,
727            };
728            if ended {
729                let _ = sink
730                    .send(Message::Close(Some(CloseFrame {
731                        code: close_code::NORMAL,
732                        reason: "sign in again".into(),
733                    })))
734                    .await;
735            }
736        })
737}
738
739async fn graphql_playground(
740    axum::extract::Query(params): axum::extract::Query<std::collections::HashMap<String, String>>,
741    axum::extract::State(key): axum::extract::State<Option<Arc<String>>>,
742) -> axum::response::Response {
743    use axum::response::IntoResponse;
744
745    // If an introspection key exists, require it in the URL.
746    if let Some(ref expected) = key {
747        let provided = params.get("introspection-key");
748        if provided.map(|k| k.as_str()) != Some(expected.as_str()) {
749            return (
750                axum::http::StatusCode::FORBIDDEN,
751                "invalid or missing introspection-key",
752            )
753                .into_response();
754        }
755    }
756
757    // Use async-graphql's built-in GraphiQL (self-contained, no CDN).
758    // Inject the introspection key as a default header so all queries are authed.
759    let mut source = async_graphql::http::GraphiQLSource::build().endpoint("/graphql");
760    if let Some(ref k) = key {
761        source = source.header("X-Introspection-Key", k.as_str());
762    }
763
764    axum::response::Html(source.finish()).into_response()
765}
766
767/// Run the server as a background daemon (fork + detach).
768pub fn cmd_serve_daemon(
769    port: Option<u16>,
770    bind: Option<std::net::IpAddr>,
771    subsonic_port: Option<u16>,
772    playground: bool,
773) {
774    use std::fs;
775    use std::process::Command;
776
777    let cfg = Config::load().unwrap_or_default();
778    let port_val = port.unwrap_or(cfg.graphql.port);
779    let bind_val = bind.unwrap_or(cfg.graphql.bind);
780
781    let exe = std::env::current_exe().expect("failed to get current exe path");
782    let mut cmd = Command::new(exe);
783
784    cmd.arg("--headless");
785    cmd.arg("--port").arg(port_val.to_string());
786    cmd.arg("--bind").arg(bind_val.to_string());
787    if let Some(sp) = subsonic_port {
788        cmd.arg("--subsonic").arg(sp.to_string());
789    }
790    if playground || cfg.graphql.playground {
791        cmd.arg("--playground");
792    }
793
794    cmd.stdin(std::process::Stdio::null());
795    cmd.stdout(std::process::Stdio::null());
796    cmd.stderr(std::process::Stdio::null());
797
798    let mut child = cmd.spawn().expect("failed to spawn daemon process");
799    let pid = child.id();
800
801    let pid_path = koan_core::config::config_dir().join("koan-serve.pid");
802    fs::write(&pid_path, pid.to_string()).ok();
803
804    std::thread::spawn(move || {
805        let _ = child.wait();
806    });
807
808    eprintln!("kōan daemon started (pid {}) on port {}", pid, port_val);
809    if let Some(sp) = subsonic_port {
810        eprintln!("  Subsonic REST on port {}", sp);
811    }
812    eprintln!("  PID file: {}", pid_path.display());
813}
814
815// ---------------------------------------------------------------------------
816// In-process execution (for MCP `graphql` tool)
817// ---------------------------------------------------------------------------
818
819/// Execute a GraphQL query in-process (no HTTP round-trip).
820///
821/// There is no credential to check, so the caller states who the query runs as
822/// and at what role — see `mcp::mcp_role`.
823pub async fn execute_in_process(
824    schema: &KoanSchema,
825    query: &str,
826    variables: Option<serde_json::Value>,
827    caller: AuthUser,
828) -> serde_json::Value {
829    let mut request = async_graphql::Request::new(query).data(caller);
830    if let Some(serde_json::Value::Object(map)) = variables {
831        let mut gql_vars = async_graphql::Variables::default();
832        for (k, v) in map {
833            gql_vars.insert(
834                async_graphql::Name::new(&k),
835                async_graphql::Value::from_json(v).unwrap_or(async_graphql::Value::Null),
836            );
837        }
838        request = request.variables(gql_vars);
839    }
840    let response = schema.execute(request).await;
841    serde_json::to_value(&response).unwrap_or(serde_json::Value::Null)
842}
843
844// ---------------------------------------------------------------------------
845// Tests
846// ---------------------------------------------------------------------------
847
848#[cfg(test)]
849mod tests {
850    use super::*;
851    use axum::body::Body;
852    use axum::http::{Request as HttpRequest, StatusCode};
853    use axum::routing::{get, post};
854    use tower::ServiceExt as _;
855
856    fn policy() -> Arc<BrowserPolicy> {
857        Arc::new(BrowserPolicy {
858            origins: vec!["https://music.example.com".into()],
859            hosts: vec!["koan.local".into()],
860        })
861    }
862
863    async fn ok() -> &'static str {
864        "ok"
865    }
866
867    fn routes() -> axum::Router<Arc<BrowserPolicy>> {
868        axum::Router::new()
869            .route("/graphql", post(ok).get(ok))
870            .route("/graphql/ws", get(ok))
871    }
872
873    async fn run_host(req: HttpRequest<Body>) -> StatusCode {
874        let app = routes()
875            .layer(axum::middleware::from_fn_with_state(policy(), host_guard))
876            .with_state(policy());
877        app.oneshot(req).await.unwrap().status()
878    }
879
880    async fn run_browser(req: HttpRequest<Body>) -> StatusCode {
881        let app = routes()
882            .layer(axum::middleware::from_fn_with_state(
883                policy(),
884                browser_guard,
885            ))
886            .with_state(policy());
887        app.oneshot(req).await.unwrap().status()
888    }
889
890    fn json_post(uri: &str) -> axum::http::request::Builder {
891        HttpRequest::post(uri).header(axum::http::header::CONTENT_TYPE, "application/json")
892    }
893
894    // -- Host allowlist (DNS rebinding) --
895
896    #[test]
897    fn host_policy_accepts_loopback_literals_and_configured_names() {
898        let p = policy();
899        assert!(p.host_allowed("localhost:4000"));
900        assert!(p.host_allowed("127.0.0.1:4000"));
901        assert!(p.host_allowed("192.168.1.20:4000"));
902        assert!(p.host_allowed("[::1]:4000"));
903        assert!(p.host_allowed("koan.local"));
904        assert!(p.host_allowed("koan.local:4000"));
905    }
906
907    #[test]
908    fn host_policy_rejects_attacker_controlled_names() {
909        let p = policy();
910        assert!(!p.host_allowed("evil.com"));
911        assert!(!p.host_allowed("rebind.evil.com:4000"));
912        assert!(!p.host_allowed("koan.local.evil.com"));
913    }
914
915    #[tokio::test]
916    async fn host_guard_rejects_foreign_host() {
917        let req = json_post("/graphql")
918            .header(axum::http::header::HOST, "rebind.evil.com")
919            .body(Body::empty())
920            .unwrap();
921        assert_eq!(run_host(req).await, StatusCode::FORBIDDEN);
922    }
923
924    #[tokio::test]
925    async fn host_guard_allows_known_host_and_missing_host() {
926        let req = json_post("/graphql")
927            .header(axum::http::header::HOST, "127.0.0.1:4000")
928            .body(Body::empty())
929            .unwrap();
930        assert_eq!(run_host(req).await, StatusCode::OK);
931
932        let req = json_post("/graphql").body(Body::empty()).unwrap();
933        assert_eq!(run_host(req).await, StatusCode::OK);
934    }
935
936    // -- Cross-site WebSocket --
937
938    #[tokio::test]
939    async fn ws_upgrade_from_foreign_origin_is_rejected() {
940        let req = HttpRequest::get("/graphql/ws")
941            .header(axum::http::header::HOST, "127.0.0.1:4000")
942            .header(axum::http::header::ORIGIN, "https://evil.com")
943            .body(Body::empty())
944            .unwrap();
945        assert_eq!(run_browser(req).await, StatusCode::FORBIDDEN);
946    }
947
948    #[tokio::test]
949    async fn ws_upgrade_without_origin_is_allowed() {
950        let req = HttpRequest::get("/graphql/ws")
951            .header(axum::http::header::HOST, "127.0.0.1:4000")
952            .body(Body::empty())
953            .unwrap();
954        assert_eq!(run_browser(req).await, StatusCode::OK);
955    }
956
957    #[tokio::test]
958    async fn configured_and_same_origin_are_allowed() {
959        let req = HttpRequest::get("/graphql/ws")
960            .header(axum::http::header::HOST, "127.0.0.1:4000")
961            .header(axum::http::header::ORIGIN, "https://music.example.com")
962            .body(Body::empty())
963            .unwrap();
964        assert_eq!(run_browser(req).await, StatusCode::OK);
965
966        // The bundled playground posts to the host it was served from.
967        let req = json_post("/graphql")
968            .header(axum::http::header::HOST, "127.0.0.1:4000")
969            .header(axum::http::header::ORIGIN, "http://127.0.0.1:4000")
970            .body(Body::empty())
971            .unwrap();
972        assert_eq!(run_browser(req).await, StatusCode::OK);
973    }
974
975    // -- CSRF via a CORS-safelisted content type --
976
977    #[tokio::test]
978    async fn text_plain_post_is_rejected() {
979        let req = HttpRequest::post("/graphql")
980            .header(axum::http::header::CONTENT_TYPE, "text/plain")
981            .body(Body::from(r#"{"query":"mutation{clearQueue{ok}}"}"#))
982            .unwrap();
983        assert_eq!(run_browser(req).await, StatusCode::UNSUPPORTED_MEDIA_TYPE);
984    }
985
986    #[tokio::test]
987    async fn post_without_content_type_is_rejected() {
988        let req = HttpRequest::post("/graphql").body(Body::empty()).unwrap();
989        assert_eq!(run_browser(req).await, StatusCode::UNSUPPORTED_MEDIA_TYPE);
990    }
991
992    // -- Load perimeter --
993
994    #[tokio::test]
995    async fn load_perimeter_refuses_an_oversized_query_body() {
996        async fn parse(_: async_graphql_axum::GraphQLRequest) -> StatusCode {
997            StatusCode::OK
998        }
999        let app = load_perimeter(axum::Router::new().route("/graphql", post(parse)));
1000        // A valid query padded with whitespace, streamed so no Content-Length
1001        // warns the limit ahead of the bytes.
1002        let body = |padding: usize| {
1003            let chunks = [
1004                axum::body::Bytes::from_static(br#"{"query":"{__typename}""#),
1005                axum::body::Bytes::from(vec![b' '; padding]),
1006                axum::body::Bytes::from_static(b"}"),
1007            ];
1008            Body::from_stream(tokio_stream::iter(chunks.map(Ok::<_, std::io::Error>)))
1009        };
1010        let req = json_post("/graphql").body(body(1 << 10)).unwrap();
1011        assert_eq!(
1012            app.clone().oneshot(req).await.unwrap().status(),
1013            StatusCode::OK
1014        );
1015        let req = json_post("/graphql").body(body(3 << 20)).unwrap();
1016        assert_ne!(app.oneshot(req).await.unwrap().status(), StatusCode::OK);
1017    }
1018
1019    #[tokio::test]
1020    async fn load_perimeter_passes_requests_and_turns_panics_into_500s() {
1021        async fn boom() -> &'static str {
1022            panic!("resolver exploded");
1023        }
1024
1025        let app = load_perimeter(
1026            axum::Router::new()
1027                .route("/graphql", post(ok))
1028                .route("/boom", post(boom)),
1029        );
1030
1031        let req = json_post("/graphql").body(Body::empty()).unwrap();
1032        assert_eq!(
1033            app.clone().oneshot(req).await.unwrap().status(),
1034            StatusCode::OK
1035        );
1036
1037        // Without CatchPanicLayer this drops the connection with nothing logged.
1038        let req = json_post("/boom").body(Body::empty()).unwrap();
1039        assert_eq!(
1040            app.oneshot(req).await.unwrap().status(),
1041            StatusCode::INTERNAL_SERVER_ERROR
1042        );
1043    }
1044
1045    #[tokio::test]
1046    async fn json_post_is_accepted() {
1047        let req = json_post("/graphql").body(Body::empty()).unwrap();
1048        assert_eq!(run_browser(req).await, StatusCode::OK);
1049
1050        let req = HttpRequest::post("/graphql")
1051            .header(
1052                axum::http::header::CONTENT_TYPE,
1053                "application/json; charset=utf-8",
1054            )
1055            .body(Body::empty())
1056            .unwrap();
1057        assert_eq!(run_browser(req).await, StatusCode::OK);
1058    }
1059}