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