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