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