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