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