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
17pub 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 if let Err(e) = koan_core::db::connection::Database::open_default() {
33 log::error!("cannot open the database: {e}");
34 eprintln!("koan: cannot open the database: {e}");
35 std::process::exit(1);
36 }
37 let db_path = koan_core::config::db_path();
38 let pool = Arc::new(Pool::new(db_path.clone()));
39 koan_core::db::pool::on_outdated(|| OUTDATED.cancel());
40
41 let (state, _timeline, _viz, cmd_tx) = Player::spawn();
42
43 let watched = db_path.clone();
46 koan_core::helpers::spawn_library_watch(db_path, move |running| {
47 if !running {
50 crate::clients::fulfil_from(&watched);
51 if let Ok(db) = koan_core::db::connection::Database::open_existing(&watched) {
52 crate::clients::changed_if_library_moved(&db.conn);
53 }
54 }
55 });
56
57 if let Err(e) = run_api_blocking(ApiServerOpts {
58 state,
59 cmd_tx,
60 pool,
61 port,
62 bind,
63 subsonic_port,
64 playground,
65 viz: None, headless: true,
67 }) {
68 eprintln!("koan: {}", e);
69 std::process::exit(1);
70 }
71}
72
73const REQUEST_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
80
81const MAX_INFLIGHT_QUERIES: usize = 64;
85
86const MAX_QUERY_BODY: usize = 2 << 20;
90
91fn load_perimeter<S>(router: axum::Router<S>) -> axum::Router<S>
96where
97 S: Clone + Send + Sync + 'static,
98{
99 router
100 .layer(tower_http::catch_panic::CatchPanicLayer::new())
103 .layer(tower_http::limit::RequestBodyLimitLayer::new(
104 MAX_QUERY_BODY,
105 ))
106 .layer(tower_http::timeout::TimeoutLayer::with_status_code(
107 axum::http::StatusCode::REQUEST_TIMEOUT,
108 REQUEST_TIMEOUT,
109 ))
110 .layer(
114 tower::ServiceBuilder::new()
115 .layer(axum::error_handling::HandleErrorLayer::new(
116 |err: tower::BoxError| async move {
117 if err.is::<tower::load_shed::error::Overloaded>() {
118 (
119 axum::http::StatusCode::SERVICE_UNAVAILABLE,
120 "server at capacity",
121 )
122 } else {
123 (
124 axum::http::StatusCode::INTERNAL_SERVER_ERROR,
125 "internal error",
126 )
127 }
128 },
129 ))
130 .load_shed()
131 .concurrency_limit(MAX_INFLIGHT_QUERIES),
132 )
133}
134
135pub struct ApiServerOpts {
137 pub state: Arc<SharedPlayerState>,
138 pub cmd_tx: Sender<PlayerCommand>,
139 pub pool: Arc<Pool>,
142 pub port: Option<u16>,
143 pub bind: Option<std::net::IpAddr>,
144 pub subsonic_port: Option<u16>,
145 pub playground: bool,
146 pub viz: Option<Arc<VizSnapshot>>,
147 pub headless: bool,
149}
150
151fn run_api_blocking(opts: ApiServerOpts) -> Result<(), String> {
157 let ApiServerOpts {
158 state,
159 cmd_tx,
160 pool,
161 port,
162 bind,
163 subsonic_port,
164 playground,
165 viz,
166 headless,
167 } = opts;
168 use axum::routing::{get, post};
169
170 let cfg = Config::load().unwrap_or_default();
171 let port = port.unwrap_or(cfg.graphql.port);
172 let bind = bind.unwrap_or(cfg.graphql.bind);
173 let subsonic_port = subsonic_port.or(cfg.subsonic.port);
174 let playground_enabled = playground || cfg.graphql.playground;
175 let auth_enabled = cfg.graphql.auth_enabled;
176
177 let (private_pem, public_pem) = if auth_enabled {
183 let kp = auth::load_or_generate_keypair().map_err(|e| {
184 format!(
185 "auth_enabled = true but the keypair could not be loaded or created: {}",
186 e
187 )
188 })?;
189 if kp.0.is_empty() || kp.1.is_empty() {
192 return Err("auth_enabled = true but the keypair files are empty. \
193 Run `koan auth regenerate-keys`."
194 .into());
195 }
196 kp
197 } else {
198 auth::load_or_generate_keypair().unwrap_or_default()
201 };
202
203 let _ = std::fs::remove_file(auth::keypair_dir().join("subsonic.key"));
206
207 let access_ttl = parse_duration_secs(&cfg.graphql.access_token_ttl).unwrap_or(900);
208 let refresh_ttl = parse_duration_secs(&cfg.graphql.refresh_token_ttl).unwrap_or(2_592_000);
209
210 let introspection_key = if playground_enabled && auth_enabled {
214 Some(Arc::new(auth::random_token().map_err(|e| {
215 format!("failed to generate introspection key: {}", e)
216 })?))
217 } else {
218 None
219 };
220
221 let auth_state = AuthState {
222 public_pem: Arc::new(public_pem.clone()),
223 auth_enabled,
224 introspection_key: introspection_key.clone(),
225 pool: pool.clone(),
226 };
227
228 crate::auth::set_signing_keys(Arc::new(private_pem.clone()), Arc::new(public_pem.clone()));
229
230 let users = Arc::new(crate::auth::password::PasswordVerifier::new(pool.clone()));
232 let auth_route_state = AuthRouteState {
233 users: users.clone(),
234 pool: pool.clone(),
235 private_pem: Arc::new(private_pem),
236 public_pem: Arc::new(public_pem),
237 access_ttl_secs: access_ttl,
238 refresh_ttl_secs: refresh_ttl,
239 cookie_secure: cfg.graphql.cookie_secure,
240 login_limiter: Arc::new(RateLimiter::default()),
241 };
242
243 koan_core::scrobbling::start(pool.path().to_path_buf());
245
246 let shutdown = tokio_util::sync::CancellationToken::new();
247 let mcp_routes = crate::mcp::router(
248 state.clone(),
249 cmd_tx.clone(),
250 auth_state.clone(),
251 cfg.sharing.public_url.clone(),
252 headless,
253 shutdown.clone(),
254 );
255 let schema = build_schema(state, cmd_tx, pool.clone(), viz);
256
257 if auth_enabled {
258 log::info!(
259 "Auth enabled (Ed25519 JWT, access TTL {}s, refresh TTL {}s)",
260 access_ttl,
261 refresh_ttl
262 );
263 let no_accounts = pool
264 .get()
265 .ok()
266 .and_then(|db| koan_core::db::queries::auth::has_users(&db.conn).ok())
267 == Some(false);
268 if no_accounts && cfg.graphql.setup_wizard {
269 log::warn!(
270 "No accounts yet: whoever opens /setup first becomes the admin. \
271 Set graphql.setup_wizard = false to make it with `koan auth setup` instead."
272 );
273 }
274 } else {
275 log::info!("Auth disabled — all requests treated as admin");
276 }
277
278 let browser_policy = Arc::new(BrowserPolicy {
279 origins: cfg.graphql.cors_origins.clone(),
280 hosts: cfg.graphql.allowed_hosts.clone(),
281 });
282
283 if cfg.graphql.cors_origins.is_empty() {
284 log::info!("CORS: no origins configured — browsers get no cross-origin access");
285 }
286
287 let proxy_auth = crate::ui::ProxyAuth::from_config(
288 &cfg.graphql.proxy_auth_header,
289 &cfg.graphql.proxy_auth_from,
290 )?;
291
292 let rt = tokio::runtime::Runtime::new().expect("failed to create tokio runtime");
293 rt.block_on(async {
294 let query_route = load_perimeter(axum::Router::new().route("/graphql", post(graphql_handler)));
299
300 let gql_app = axum::Router::new()
301 .merge(query_route)
302 .route("/graphql/ws", get(graphql_ws_handler))
303 .layer(axum::middleware::from_fn_with_state(
304 auth_state.clone(),
305 auth_middleware,
306 ))
307 .layer(axum::middleware::from_fn_with_state(
310 browser_policy.clone(),
311 browser_guard,
312 ))
313 .with_state(schema);
314
315 let covers = Arc::new(crate::covers::Covers::in_config_dir());
319 let ui_routes = crate::ui::router(
320 pool.clone(),
321 auth_route_state.clone(),
322 auth_enabled,
323 cfg.graphql.setup_wizard,
324 covers.clone(),
325 cfg.sharing.public_url.clone(),
326 cfg.mcp.redirect_hosts.clone(),
327 proxy_auth,
328 );
329
330 let auth_app = auth_router(auth_route_state);
332
333 let origins: Vec<axum::http::HeaderValue> = cfg
336 .graphql
337 .cors_origins
338 .iter()
339 .filter_map(|o| o.parse().ok())
340 .collect();
341 let cors = tower_http::cors::CorsLayer::new()
342 .allow_origin(origins)
343 .allow_methods([
344 axum::http::Method::GET,
345 axum::http::Method::POST,
346 axum::http::Method::OPTIONS,
347 ])
348 .allow_headers([
349 axum::http::header::AUTHORIZATION,
350 axum::http::header::CONTENT_TYPE,
351 axum::http::HeaderName::from_static("x-introspection-key"),
352 ])
353 .allow_credentials(true);
354
355
356 let share_routes = crate::share::router(
360 pool.clone(),
361 cfg.sharing.public_url.clone(),
362 covers.clone(),
363 )
364 .merge(crate::push::router(pool.clone(), covers.clone()));
365 let subsonic_merged = crate::subsonic::subsonic_router(pool, covers, users);
370 let subsonic_on_main = subsonic_merged.is_some();
371 let subsonic_dedicated = subsonic_merged.clone();
372
373 let mut app = auth_app
374 .merge(gql_app)
375 .merge(share_routes)
376 .merge(ui_routes)
377 .merge(mcp_routes);
378 if let Some(sub) = subsonic_merged {
379 app = app.merge(sub);
380 }
381 if playground_enabled {
382 app = app.route(
383 "/graphql",
384 get(graphql_playground).with_state(introspection_key.clone()),
385 );
386 }
387 let app = app.layer(cors).layer(axum::middleware::from_fn_with_state(
391 browser_policy.clone(),
392 host_guard,
393 ));
394
395 let playground_url = if playground_enabled {
397 if let Some(ref key) = introspection_key {
398 format!("http://{}:{}/graphql?introspection-key={}", bind, port, key)
399 } else {
400 format!("http://{}:{}/graphql", bind, port)
401 }
402 } else {
403 format!("http://{}:{}/graphql", bind, port)
404 };
405
406 let gql_addr = std::net::SocketAddr::new(bind, port);
407
408 let gql_listener = match tokio::net::TcpListener::bind(gql_addr).await {
409 Ok(l) => {
410 log::info!("GraphQL API on http://{}:{}/graphql", bind, port);
411 if subsonic_on_main {
412 log::info!("Subsonic REST on http://{}:{}/rest/", bind, port);
413 }
414 if playground_enabled {
415 log::info!("GraphiQL: {}", playground_url);
416 #[cfg(target_os = "macos")]
418 let _ = std::process::Command::new("open").arg(&playground_url).spawn();
419 #[cfg(target_os = "linux")]
420 let _ = std::process::Command::new("xdg-open").arg(&playground_url).spawn();
421 }
422 l
423 }
424 Err(e) => {
425 return Err(format!(
426 "failed to bind GraphQL port {port} — {e} (another instance running?)"
427 ));
428 }
429 };
430 let gql_server = axum::serve(
431 gql_listener,
432 app.into_make_service_with_connect_info::<std::net::SocketAddr>(),
433 )
434 .with_graceful_shutdown(async move {
435 shutdown_signal().await;
436 shutdown.cancel();
437 });
438
439 let extra_sub_port = subsonic_port.filter(|p| *p != port);
442 if let Some(sub_port) = extra_sub_port
443 && let Some(sub_app) = subsonic_dedicated
444 {
445 let sub_addr = std::net::SocketAddr::new(bind, sub_port);
446 match tokio::net::TcpListener::bind(sub_addr).await {
447 Ok(sub_listener) => {
448 log::info!(
449 "Subsonic REST also on http://{}:{}/rest/ (dedicated port)",
450 bind,
451 sub_port,
452 );
453 let sub_server = axum::serve(
456 sub_listener,
457 sub_app.into_make_service_with_connect_info::<std::net::SocketAddr>(),
458 )
459 .with_graceful_shutdown(shutdown_signal());
460
461 tokio::select! {
462 r = gql_server => { if let Err(e) = r { log::error!("GraphQL server error: {e}"); } },
463 r = sub_server => { if let Err(e) = r { log::error!("Subsonic server error: {e}"); } },
464 _ = drain_deadline() => log::info!("shutting down with connections still open"),
465 }
466 return Ok(());
467 }
468 Err(e) => {
469 log::warn!(
470 "Dedicated Subsonic port {} unavailable — {}. Mounted on GraphQL port only.",
471 sub_port,
472 e,
473 );
474 }
475 }
476 }
477
478 tokio::select! {
479 r = gql_server => { if let Err(e) = r { log::error!("GraphQL server error: {e}"); } },
480 _ = drain_deadline() => log::info!("shutting down with connections still open"),
481 }
482 Ok(())
483 })
484}
485
486pub fn start_api_background(
489 state: Arc<SharedPlayerState>,
490 cmd_tx: Sender<PlayerCommand>,
491 db_path: PathBuf,
492 port: Option<u16>,
493 bind: Option<std::net::IpAddr>,
494 subsonic_port: Option<u16>,
495 playground: bool,
496) {
497 if let Err(e) = run_api_blocking(ApiServerOpts {
500 state,
501 cmd_tx,
502 pool: Arc::new(Pool::new(db_path)),
503 port,
504 bind,
505 subsonic_port,
506 playground,
507 viz: None,
508 headless: false,
509 }) {
510 log::error!("API server not started: {}", e);
511 }
512}
513
514pub(crate) struct BrowserPolicy {
523 origins: Vec<String>,
524 hosts: Vec<String>,
525}
526
527impl BrowserPolicy {
528 fn host_allowed(&self, host: &str) -> bool {
529 if self.hosts.iter().any(|h| h.eq_ignore_ascii_case(host)) {
530 return true;
531 }
532 let bare = strip_port(host);
533 if self.hosts.iter().any(|h| h.eq_ignore_ascii_case(bare)) {
534 return true;
535 }
536 bare.eq_ignore_ascii_case("localhost") || bare.parse::<std::net::IpAddr>().is_ok()
539 }
540
541 fn origin_allowed(&self, origin: &str, host: Option<&str>) -> bool {
544 if self.origins.iter().any(|o| o == origin) {
545 return true;
546 }
547 match (origin.split_once("://"), host) {
548 (Some((_, authority)), Some(host)) => authority.eq_ignore_ascii_case(host),
549 _ => false,
550 }
551 }
552}
553
554fn strip_port(host: &str) -> &str {
556 if let Some(rest) = host.strip_prefix('[') {
557 return rest.split(']').next().unwrap_or(rest);
558 }
559 match host.rsplit_once(':') {
560 Some((h, port)) if !port.is_empty() && port.bytes().all(|b| b.is_ascii_digit()) => h,
561 _ => host,
562 }
563}
564
565fn header_str(request: &axum::extract::Request, name: axum::http::HeaderName) -> Option<&str> {
566 request.headers().get(name).and_then(|v| v.to_str().ok())
567}
568
569async fn host_guard(
571 axum::extract::State(policy): axum::extract::State<Arc<BrowserPolicy>>,
572 request: axum::extract::Request,
573 next: axum::middleware::Next,
574) -> axum::response::Response {
575 use axum::response::IntoResponse;
576
577 let host = header_str(&request, axum::http::header::HOST)
580 .map(str::to_owned)
581 .or_else(|| request.uri().host().map(str::to_owned));
582
583 if let Some(ref host) = host
584 && !policy.host_allowed(host)
585 {
586 log::warn!("rejected request for unrecognised Host: {}", host);
587 return (axum::http::StatusCode::FORBIDDEN, "host not allowed").into_response();
588 }
589
590 next.run(request).await
591}
592
593async fn browser_guard(
602 axum::extract::State(policy): axum::extract::State<Arc<BrowserPolicy>>,
603 request: axum::extract::Request,
604 next: axum::middleware::Next,
605) -> axum::response::Response {
606 use axum::response::IntoResponse;
607
608 let host = header_str(&request, axum::http::header::HOST).map(str::to_owned);
609 if let Some(origin) = header_str(&request, axum::http::header::ORIGIN)
611 && !policy.origin_allowed(origin, host.as_deref())
612 {
613 log::warn!(
614 "rejected GraphQL request from disallowed Origin: {}",
615 origin
616 );
617 return (axum::http::StatusCode::FORBIDDEN, "origin not allowed").into_response();
618 }
619
620 if request.method() == axum::http::Method::POST && !is_graphql_content_type(&request) {
621 return (
622 axum::http::StatusCode::UNSUPPORTED_MEDIA_TYPE,
623 "content type must be application/json or application/graphql",
624 )
625 .into_response();
626 }
627
628 next.run(request).await
629}
630
631fn is_graphql_content_type(request: &axum::extract::Request) -> bool {
632 header_str(request, axum::http::header::CONTENT_TYPE).is_some_and(|ct| {
633 let ct = ct.trim().to_ascii_lowercase();
634 ct.starts_with("application/json") || ct.starts_with("application/graphql")
635 })
636}
637
638static OUTDATED: std::sync::LazyLock<tokio_util::sync::CancellationToken> =
642 std::sync::LazyLock::new(tokio_util::sync::CancellationToken::new);
643
644async fn shutdown_signal() {
649 #[cfg(unix)]
650 {
651 use tokio::signal::unix::{SignalKind, signal};
652 let mut terminate = signal(SignalKind::terminate()).expect("failed to listen for SIGTERM");
653 tokio::select! {
654 r = tokio::signal::ctrl_c() => r.expect("failed to listen for ctrl+c"),
655 _ = terminate.recv() => {}
656 _ = OUTDATED.cancelled() => {}
657 }
658 }
659 #[cfg(not(unix))]
660 tokio::select! {
661 r = tokio::signal::ctrl_c() => r.expect("failed to listen for ctrl+c"),
662 _ = OUTDATED.cancelled() => {}
663 }
664}
665
666const DRAIN: std::time::Duration = std::time::Duration::from_secs(10);
668
669async fn drain_deadline() {
673 shutdown_signal().await;
674 tokio::time::sleep(DRAIN).await;
675}
676
677async fn graphql_handler(
678 axum::Extension(user): axum::Extension<AuthUser>,
679 axum::extract::State(schema): axum::extract::State<KoanSchema>,
680 headers: axum::http::HeaderMap,
681 req: async_graphql_axum::GraphQLRequest,
682) -> async_graphql_axum::GraphQLResponse {
683 let mut request = req.into_inner();
684 if let Some(origin) = crate::origin::origin(&headers, None) {
685 request = request.data(super::RequestOrigin(origin));
686 }
687 request = request.data(user);
690 schema.execute(request).await.into()
691}
692
693async fn graphql_ws_handler(
697 axum::Extension(user): axum::Extension<AuthUser>,
698 lease: Option<axum::Extension<crate::auth::Lease>>,
699 axum::extract::State(schema): axum::extract::State<KoanSchema>,
700 protocol: async_graphql_axum::GraphQLProtocol,
701 websocket: axum::extract::WebSocketUpgrade,
702) -> axum::response::Response {
703 use axum::extract::ws::{CloseFrame, Message, close_code};
704 use futures_util::{SinkExt, StreamExt};
705 websocket
706 .protocols(async_graphql::http::ALL_WEBSOCKET_PROTOCOLS)
707 .on_upgrade(move |socket| async move {
708 let (mut sink, stream) = socket.split();
709 let serve = async_graphql_axum::GraphQLWebSocket::new_with_pair(
710 &mut sink, stream, schema, protocol,
711 )
712 .on_connection_init(move |_| async move {
713 let mut data = async_graphql::Data::default();
714 data.insert(user);
715 Ok(data)
716 })
717 .serve();
718 let ended = async {
719 match lease {
720 Some(axum::Extension(lease)) => lease.ended().await,
721 None => std::future::pending().await,
722 }
723 };
724 let ended = tokio::select! {
725 _ = serve => false,
726 _ = ended => true,
727 };
728 if ended {
729 let _ = sink
730 .send(Message::Close(Some(CloseFrame {
731 code: close_code::NORMAL,
732 reason: "sign in again".into(),
733 })))
734 .await;
735 }
736 })
737}
738
739async fn graphql_playground(
740 axum::extract::Query(params): axum::extract::Query<std::collections::HashMap<String, String>>,
741 axum::extract::State(key): axum::extract::State<Option<Arc<String>>>,
742) -> axum::response::Response {
743 use axum::response::IntoResponse;
744
745 if let Some(ref expected) = key {
747 let provided = params.get("introspection-key");
748 if provided.map(|k| k.as_str()) != Some(expected.as_str()) {
749 return (
750 axum::http::StatusCode::FORBIDDEN,
751 "invalid or missing introspection-key",
752 )
753 .into_response();
754 }
755 }
756
757 let mut source = async_graphql::http::GraphiQLSource::build().endpoint("/graphql");
760 if let Some(ref k) = key {
761 source = source.header("X-Introspection-Key", k.as_str());
762 }
763
764 axum::response::Html(source.finish()).into_response()
765}
766
767pub fn cmd_serve_daemon(
769 port: Option<u16>,
770 bind: Option<std::net::IpAddr>,
771 subsonic_port: Option<u16>,
772 playground: bool,
773) {
774 use std::fs;
775 use std::process::Command;
776
777 let cfg = Config::load().unwrap_or_default();
778 let port_val = port.unwrap_or(cfg.graphql.port);
779 let bind_val = bind.unwrap_or(cfg.graphql.bind);
780
781 let exe = std::env::current_exe().expect("failed to get current exe path");
782 let mut cmd = Command::new(exe);
783
784 cmd.arg("--headless");
785 cmd.arg("--port").arg(port_val.to_string());
786 cmd.arg("--bind").arg(bind_val.to_string());
787 if let Some(sp) = subsonic_port {
788 cmd.arg("--subsonic").arg(sp.to_string());
789 }
790 if playground || cfg.graphql.playground {
791 cmd.arg("--playground");
792 }
793
794 cmd.stdin(std::process::Stdio::null());
795 cmd.stdout(std::process::Stdio::null());
796 cmd.stderr(std::process::Stdio::null());
797
798 let mut child = cmd.spawn().expect("failed to spawn daemon process");
799 let pid = child.id();
800
801 let pid_path = koan_core::config::config_dir().join("koan-serve.pid");
802 fs::write(&pid_path, pid.to_string()).ok();
803
804 std::thread::spawn(move || {
805 let _ = child.wait();
806 });
807
808 eprintln!("kōan daemon started (pid {}) on port {}", pid, port_val);
809 if let Some(sp) = subsonic_port {
810 eprintln!(" Subsonic REST on port {}", sp);
811 }
812 eprintln!(" PID file: {}", pid_path.display());
813}
814
815pub async fn execute_in_process(
824 schema: &KoanSchema,
825 query: &str,
826 variables: Option<serde_json::Value>,
827 caller: AuthUser,
828) -> serde_json::Value {
829 let mut request = async_graphql::Request::new(query).data(caller);
830 if let Some(serde_json::Value::Object(map)) = variables {
831 let mut gql_vars = async_graphql::Variables::default();
832 for (k, v) in map {
833 gql_vars.insert(
834 async_graphql::Name::new(&k),
835 async_graphql::Value::from_json(v).unwrap_or(async_graphql::Value::Null),
836 );
837 }
838 request = request.variables(gql_vars);
839 }
840 let response = schema.execute(request).await;
841 serde_json::to_value(&response).unwrap_or(serde_json::Value::Null)
842}
843
844#[cfg(test)]
849mod tests {
850 use super::*;
851 use axum::body::Body;
852 use axum::http::{Request as HttpRequest, StatusCode};
853 use axum::routing::{get, post};
854 use tower::ServiceExt as _;
855
856 fn policy() -> Arc<BrowserPolicy> {
857 Arc::new(BrowserPolicy {
858 origins: vec!["https://music.example.com".into()],
859 hosts: vec!["koan.local".into()],
860 })
861 }
862
863 async fn ok() -> &'static str {
864 "ok"
865 }
866
867 fn routes() -> axum::Router<Arc<BrowserPolicy>> {
868 axum::Router::new()
869 .route("/graphql", post(ok).get(ok))
870 .route("/graphql/ws", get(ok))
871 }
872
873 async fn run_host(req: HttpRequest<Body>) -> StatusCode {
874 let app = routes()
875 .layer(axum::middleware::from_fn_with_state(policy(), host_guard))
876 .with_state(policy());
877 app.oneshot(req).await.unwrap().status()
878 }
879
880 async fn run_browser(req: HttpRequest<Body>) -> StatusCode {
881 let app = routes()
882 .layer(axum::middleware::from_fn_with_state(
883 policy(),
884 browser_guard,
885 ))
886 .with_state(policy());
887 app.oneshot(req).await.unwrap().status()
888 }
889
890 fn json_post(uri: &str) -> axum::http::request::Builder {
891 HttpRequest::post(uri).header(axum::http::header::CONTENT_TYPE, "application/json")
892 }
893
894 #[test]
897 fn host_policy_accepts_loopback_literals_and_configured_names() {
898 let p = policy();
899 assert!(p.host_allowed("localhost:4000"));
900 assert!(p.host_allowed("127.0.0.1:4000"));
901 assert!(p.host_allowed("192.168.1.20:4000"));
902 assert!(p.host_allowed("[::1]:4000"));
903 assert!(p.host_allowed("koan.local"));
904 assert!(p.host_allowed("koan.local:4000"));
905 }
906
907 #[test]
908 fn host_policy_rejects_attacker_controlled_names() {
909 let p = policy();
910 assert!(!p.host_allowed("evil.com"));
911 assert!(!p.host_allowed("rebind.evil.com:4000"));
912 assert!(!p.host_allowed("koan.local.evil.com"));
913 }
914
915 #[tokio::test]
916 async fn host_guard_rejects_foreign_host() {
917 let req = json_post("/graphql")
918 .header(axum::http::header::HOST, "rebind.evil.com")
919 .body(Body::empty())
920 .unwrap();
921 assert_eq!(run_host(req).await, StatusCode::FORBIDDEN);
922 }
923
924 #[tokio::test]
925 async fn host_guard_allows_known_host_and_missing_host() {
926 let req = json_post("/graphql")
927 .header(axum::http::header::HOST, "127.0.0.1:4000")
928 .body(Body::empty())
929 .unwrap();
930 assert_eq!(run_host(req).await, StatusCode::OK);
931
932 let req = json_post("/graphql").body(Body::empty()).unwrap();
933 assert_eq!(run_host(req).await, StatusCode::OK);
934 }
935
936 #[tokio::test]
939 async fn ws_upgrade_from_foreign_origin_is_rejected() {
940 let req = HttpRequest::get("/graphql/ws")
941 .header(axum::http::header::HOST, "127.0.0.1:4000")
942 .header(axum::http::header::ORIGIN, "https://evil.com")
943 .body(Body::empty())
944 .unwrap();
945 assert_eq!(run_browser(req).await, StatusCode::FORBIDDEN);
946 }
947
948 #[tokio::test]
949 async fn ws_upgrade_without_origin_is_allowed() {
950 let req = HttpRequest::get("/graphql/ws")
951 .header(axum::http::header::HOST, "127.0.0.1:4000")
952 .body(Body::empty())
953 .unwrap();
954 assert_eq!(run_browser(req).await, StatusCode::OK);
955 }
956
957 #[tokio::test]
958 async fn configured_and_same_origin_are_allowed() {
959 let req = HttpRequest::get("/graphql/ws")
960 .header(axum::http::header::HOST, "127.0.0.1:4000")
961 .header(axum::http::header::ORIGIN, "https://music.example.com")
962 .body(Body::empty())
963 .unwrap();
964 assert_eq!(run_browser(req).await, StatusCode::OK);
965
966 let req = json_post("/graphql")
968 .header(axum::http::header::HOST, "127.0.0.1:4000")
969 .header(axum::http::header::ORIGIN, "http://127.0.0.1:4000")
970 .body(Body::empty())
971 .unwrap();
972 assert_eq!(run_browser(req).await, StatusCode::OK);
973 }
974
975 #[tokio::test]
978 async fn text_plain_post_is_rejected() {
979 let req = HttpRequest::post("/graphql")
980 .header(axum::http::header::CONTENT_TYPE, "text/plain")
981 .body(Body::from(r#"{"query":"mutation{clearQueue{ok}}"}"#))
982 .unwrap();
983 assert_eq!(run_browser(req).await, StatusCode::UNSUPPORTED_MEDIA_TYPE);
984 }
985
986 #[tokio::test]
987 async fn post_without_content_type_is_rejected() {
988 let req = HttpRequest::post("/graphql").body(Body::empty()).unwrap();
989 assert_eq!(run_browser(req).await, StatusCode::UNSUPPORTED_MEDIA_TYPE);
990 }
991
992 #[tokio::test]
995 async fn load_perimeter_refuses_an_oversized_query_body() {
996 async fn parse(_: async_graphql_axum::GraphQLRequest) -> StatusCode {
997 StatusCode::OK
998 }
999 let app = load_perimeter(axum::Router::new().route("/graphql", post(parse)));
1000 let body = |padding: usize| {
1003 let chunks = [
1004 axum::body::Bytes::from_static(br#"{"query":"{__typename}""#),
1005 axum::body::Bytes::from(vec![b' '; padding]),
1006 axum::body::Bytes::from_static(b"}"),
1007 ];
1008 Body::from_stream(tokio_stream::iter(chunks.map(Ok::<_, std::io::Error>)))
1009 };
1010 let req = json_post("/graphql").body(body(1 << 10)).unwrap();
1011 assert_eq!(
1012 app.clone().oneshot(req).await.unwrap().status(),
1013 StatusCode::OK
1014 );
1015 let req = json_post("/graphql").body(body(3 << 20)).unwrap();
1016 assert_ne!(app.oneshot(req).await.unwrap().status(), StatusCode::OK);
1017 }
1018
1019 #[tokio::test]
1020 async fn load_perimeter_passes_requests_and_turns_panics_into_500s() {
1021 async fn boom() -> &'static str {
1022 panic!("resolver exploded");
1023 }
1024
1025 let app = load_perimeter(
1026 axum::Router::new()
1027 .route("/graphql", post(ok))
1028 .route("/boom", post(boom)),
1029 );
1030
1031 let req = json_post("/graphql").body(Body::empty()).unwrap();
1032 assert_eq!(
1033 app.clone().oneshot(req).await.unwrap().status(),
1034 StatusCode::OK
1035 );
1036
1037 let req = json_post("/boom").body(Body::empty()).unwrap();
1039 assert_eq!(
1040 app.oneshot(req).await.unwrap().status(),
1041 StatusCode::INTERNAL_SERVER_ERROR
1042 );
1043 }
1044
1045 #[tokio::test]
1046 async fn json_post_is_accepted() {
1047 let req = json_post("/graphql").body(Body::empty()).unwrap();
1048 assert_eq!(run_browser(req).await, StatusCode::OK);
1049
1050 let req = HttpRequest::post("/graphql")
1051 .header(
1052 axum::http::header::CONTENT_TYPE,
1053 "application/json; charset=utf-8",
1054 )
1055 .body(Body::empty())
1056 .unwrap();
1057 assert_eq!(run_browser(req).await, StatusCode::OK);
1058 }
1059}