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