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 access_ttl = parse_duration_secs(&cfg.graphql.access_token_ttl).unwrap_or(900);
197 let refresh_ttl = parse_duration_secs(&cfg.graphql.refresh_token_ttl).unwrap_or(2_592_000);
198
199 let introspection_key = if playground_enabled && auth_enabled {
203 Some(Arc::new(auth::random_token().map_err(|e| {
204 format!("failed to generate introspection key: {}", e)
205 })?))
206 } else {
207 None
208 };
209
210 let auth_state = AuthState {
211 public_pem: Arc::new(public_pem.clone()),
212 auth_enabled,
213 introspection_key: introspection_key.clone(),
214 pool: pool.clone(),
215 };
216
217 let auth_route_state = AuthRouteState {
218 pool: pool.clone(),
219 private_pem: Arc::new(private_pem),
220 public_pem: Arc::new(public_pem),
221 access_ttl_secs: access_ttl,
222 refresh_ttl_secs: refresh_ttl,
223 cookie_secure: cfg.graphql.cookie_secure,
224 login_limiter: Arc::new(RateLimiter::default()),
225 };
226
227 let shutdown = tokio_util::sync::CancellationToken::new();
228 let mcp_routes = crate::mcp::router(
229 state.clone(),
230 cmd_tx.clone(),
231 auth_state.clone(),
232 cfg.sharing.public_url.clone(),
233 headless,
234 shutdown.clone(),
235 );
236 let schema = build_schema(state, cmd_tx, pool.clone(), viz);
237
238 if auth_enabled {
239 log::info!(
240 "Auth enabled (Ed25519 JWT, access TTL {}s, refresh TTL {}s)",
241 access_ttl,
242 refresh_ttl
243 );
244 } else {
245 log::info!("Auth disabled — all requests treated as admin");
246 }
247
248 let browser_policy = Arc::new(BrowserPolicy {
249 origins: cfg.graphql.cors_origins.clone(),
250 hosts: cfg.graphql.allowed_hosts.clone(),
251 });
252
253 if cfg.graphql.cors_origins.is_empty() {
254 log::info!("CORS: no origins configured — browsers get no cross-origin access");
255 }
256
257 let rt = tokio::runtime::Runtime::new().expect("failed to create tokio runtime");
258 rt.block_on(async {
259 let query_route = load_perimeter(axum::Router::new().route("/graphql", post(graphql_handler)));
264
265 let gql_app = axum::Router::new()
266 .merge(query_route)
267 .route("/graphql/ws", get(graphql_ws_handler))
268 .layer(axum::middleware::from_fn_with_state(
269 auth_state.clone(),
270 auth_middleware,
271 ))
272 .layer(axum::middleware::from_fn_with_state(
275 browser_policy.clone(),
276 browser_guard,
277 ))
278 .with_state(schema);
279
280 let covers = Arc::new(crate::covers::Covers::in_config_dir());
284 let ui_routes = crate::ui::router(
285 pool.clone(),
286 auth_route_state.clone(),
287 auth_enabled,
288 covers.clone(),
289 cfg.sharing.public_url.clone(),
290 cfg.mcp.redirect_hosts.clone(),
291 );
292
293 let auth_app = auth_router(auth_route_state);
295
296 let origins: Vec<axum::http::HeaderValue> = cfg
299 .graphql
300 .cors_origins
301 .iter()
302 .filter_map(|o| o.parse().ok())
303 .collect();
304 let cors = tower_http::cors::CorsLayer::new()
305 .allow_origin(origins)
306 .allow_methods([
307 axum::http::Method::GET,
308 axum::http::Method::POST,
309 axum::http::Method::OPTIONS,
310 ])
311 .allow_headers([
312 axum::http::header::AUTHORIZATION,
313 axum::http::header::CONTENT_TYPE,
314 axum::http::HeaderName::from_static("x-introspection-key"),
315 ])
316 .allow_credentials(true);
317
318
319 let share_routes = crate::share::router(
323 pool.clone(),
324 cfg.sharing.public_url.clone(),
325 covers.clone(),
326 )
327 .merge(crate::push::router(pool.clone(), covers.clone()));
328 let subsonic_merged = crate::subsonic::subsonic_router(pool, covers);
333 let subsonic_on_main = subsonic_merged.is_some();
334 let subsonic_dedicated = subsonic_merged.clone();
335
336 let mut app = auth_app
337 .merge(gql_app)
338 .merge(share_routes)
339 .merge(ui_routes)
340 .merge(mcp_routes);
341 if let Some(sub) = subsonic_merged {
342 app = app.merge(sub);
343 }
344 if playground_enabled {
345 app = app.route(
346 "/graphql",
347 get(graphql_playground).with_state(introspection_key.clone()),
348 );
349 }
350 let app = app.layer(cors).layer(axum::middleware::from_fn_with_state(
354 browser_policy.clone(),
355 host_guard,
356 ));
357
358 let playground_url = if playground_enabled {
360 if let Some(ref key) = introspection_key {
361 format!("http://{}:{}/graphql?introspection-key={}", bind, port, key)
362 } else {
363 format!("http://{}:{}/graphql", bind, port)
364 }
365 } else {
366 format!("http://{}:{}/graphql", bind, port)
367 };
368
369 let gql_addr = std::net::SocketAddr::new(bind, port);
370
371 let gql_listener = match tokio::net::TcpListener::bind(gql_addr).await {
372 Ok(l) => {
373 log::info!("GraphQL API on http://{}:{}/graphql", bind, port);
374 if subsonic_on_main {
375 log::info!("Subsonic REST on http://{}:{}/rest/", bind, port);
376 }
377 if playground_enabled {
378 log::info!("GraphiQL: {}", playground_url);
379 #[cfg(target_os = "macos")]
381 let _ = std::process::Command::new("open").arg(&playground_url).spawn();
382 #[cfg(target_os = "linux")]
383 let _ = std::process::Command::new("xdg-open").arg(&playground_url).spawn();
384 }
385 l
386 }
387 Err(e) => {
388 return Err(format!(
389 "failed to bind GraphQL port {port} — {e} (another instance running?)"
390 ));
391 }
392 };
393 let gql_server = axum::serve(
394 gql_listener,
395 app.into_make_service_with_connect_info::<std::net::SocketAddr>(),
396 )
397 .with_graceful_shutdown(async move {
398 shutdown_signal().await;
399 shutdown.cancel();
400 });
401
402 let extra_sub_port = subsonic_port.filter(|p| *p != port);
405 if let Some(sub_port) = extra_sub_port
406 && let Some(sub_app) = subsonic_dedicated
407 {
408 let sub_addr = std::net::SocketAddr::new(bind, sub_port);
409 match tokio::net::TcpListener::bind(sub_addr).await {
410 Ok(sub_listener) => {
411 log::info!(
412 "Subsonic REST also on http://{}:{}/rest/ (dedicated port)",
413 bind,
414 sub_port,
415 );
416 let sub_server = axum::serve(
419 sub_listener,
420 sub_app.into_make_service_with_connect_info::<std::net::SocketAddr>(),
421 )
422 .with_graceful_shutdown(shutdown_signal());
423
424 tokio::select! {
425 r = gql_server => { if let Err(e) = r { log::error!("GraphQL server error: {e}"); } },
426 r = sub_server => { if let Err(e) = r { log::error!("Subsonic server error: {e}"); } },
427 }
428 return Ok(());
429 }
430 Err(e) => {
431 log::warn!(
432 "Dedicated Subsonic port {} unavailable — {}. Mounted on GraphQL port only.",
433 sub_port,
434 e,
435 );
436 }
437 }
438 }
439
440 if let Err(e) = gql_server.await {
441 log::error!("GraphQL server error: {e}");
442 }
443 Ok(())
444 })
445}
446
447pub fn start_api_background(
450 state: Arc<SharedPlayerState>,
451 cmd_tx: Sender<PlayerCommand>,
452 db_path: PathBuf,
453 port: Option<u16>,
454 bind: Option<std::net::IpAddr>,
455 subsonic_port: Option<u16>,
456 playground: bool,
457) {
458 if let Err(e) = run_api_blocking(ApiServerOpts {
461 state,
462 cmd_tx,
463 pool: Arc::new(Pool::new(db_path)),
464 port,
465 bind,
466 subsonic_port,
467 playground,
468 viz: None,
469 headless: false,
470 }) {
471 log::error!("API server not started: {}", e);
472 }
473}
474
475pub(crate) struct BrowserPolicy {
484 origins: Vec<String>,
485 hosts: Vec<String>,
486}
487
488impl BrowserPolicy {
489 fn host_allowed(&self, host: &str) -> bool {
490 if self.hosts.iter().any(|h| h.eq_ignore_ascii_case(host)) {
491 return true;
492 }
493 let bare = strip_port(host);
494 if self.hosts.iter().any(|h| h.eq_ignore_ascii_case(bare)) {
495 return true;
496 }
497 bare.eq_ignore_ascii_case("localhost") || bare.parse::<std::net::IpAddr>().is_ok()
500 }
501
502 fn origin_allowed(&self, origin: &str, host: Option<&str>) -> bool {
505 if self.origins.iter().any(|o| o == origin) {
506 return true;
507 }
508 match (origin.split_once("://"), host) {
509 (Some((_, authority)), Some(host)) => authority.eq_ignore_ascii_case(host),
510 _ => false,
511 }
512 }
513}
514
515fn strip_port(host: &str) -> &str {
517 if let Some(rest) = host.strip_prefix('[') {
518 return rest.split(']').next().unwrap_or(rest);
519 }
520 match host.rsplit_once(':') {
521 Some((h, port)) if !port.is_empty() && port.bytes().all(|b| b.is_ascii_digit()) => h,
522 _ => host,
523 }
524}
525
526fn header_str(request: &axum::extract::Request, name: axum::http::HeaderName) -> Option<&str> {
527 request.headers().get(name).and_then(|v| v.to_str().ok())
528}
529
530async fn host_guard(
532 axum::extract::State(policy): axum::extract::State<Arc<BrowserPolicy>>,
533 request: axum::extract::Request,
534 next: axum::middleware::Next,
535) -> axum::response::Response {
536 use axum::response::IntoResponse;
537
538 let host = header_str(&request, axum::http::header::HOST)
541 .map(str::to_owned)
542 .or_else(|| request.uri().host().map(str::to_owned));
543
544 if let Some(ref host) = host
545 && !policy.host_allowed(host)
546 {
547 log::warn!("rejected request for unrecognised Host: {}", host);
548 return (axum::http::StatusCode::FORBIDDEN, "host not allowed").into_response();
549 }
550
551 next.run(request).await
552}
553
554async fn browser_guard(
563 axum::extract::State(policy): axum::extract::State<Arc<BrowserPolicy>>,
564 request: axum::extract::Request,
565 next: axum::middleware::Next,
566) -> axum::response::Response {
567 use axum::response::IntoResponse;
568
569 let host = header_str(&request, axum::http::header::HOST).map(str::to_owned);
570 if let Some(origin) = header_str(&request, axum::http::header::ORIGIN)
572 && !policy.origin_allowed(origin, host.as_deref())
573 {
574 log::warn!(
575 "rejected GraphQL request from disallowed Origin: {}",
576 origin
577 );
578 return (axum::http::StatusCode::FORBIDDEN, "origin not allowed").into_response();
579 }
580
581 if request.method() == axum::http::Method::POST && !is_graphql_content_type(&request) {
582 return (
583 axum::http::StatusCode::UNSUPPORTED_MEDIA_TYPE,
584 "content type must be application/json or application/graphql",
585 )
586 .into_response();
587 }
588
589 next.run(request).await
590}
591
592fn is_graphql_content_type(request: &axum::extract::Request) -> bool {
593 header_str(request, axum::http::header::CONTENT_TYPE).is_some_and(|ct| {
594 let ct = ct.trim().to_ascii_lowercase();
595 ct.starts_with("application/json") || ct.starts_with("application/graphql")
596 })
597}
598
599async fn shutdown_signal() {
600 tokio::signal::ctrl_c()
601 .await
602 .expect("failed to listen for ctrl+c");
603}
604
605async fn graphql_handler(
606 axum::Extension(user): axum::Extension<AuthUser>,
607 axum::extract::State(schema): axum::extract::State<KoanSchema>,
608 headers: axum::http::HeaderMap,
609 req: async_graphql_axum::GraphQLRequest,
610) -> async_graphql_axum::GraphQLResponse {
611 let mut request = req.into_inner();
612 if let Some(origin) = crate::origin::origin(&headers, None) {
613 request = request.data(super::RequestOrigin(origin));
614 }
615 request = request.data(user);
618 schema.execute(request).await.into()
619}
620
621async fn graphql_ws_handler(
622 axum::Extension(user): axum::Extension<AuthUser>,
623 axum::extract::State(schema): axum::extract::State<KoanSchema>,
624 protocol: async_graphql_axum::GraphQLProtocol,
625 websocket: axum::extract::WebSocketUpgrade,
626) -> axum::response::Response {
627 websocket
628 .protocols(async_graphql::http::ALL_WEBSOCKET_PROTOCOLS)
629 .on_upgrade(move |stream| {
630 let stream = async_graphql_axum::GraphQLWebSocket::new(stream, schema, protocol)
631 .on_connection_init(move |_| async move {
632 let mut data = async_graphql::Data::default();
633 data.insert(user);
634 Ok(data)
635 });
636 async move {
637 stream.serve().await;
638 }
639 })
640}
641
642async fn graphql_playground(
643 axum::extract::Query(params): axum::extract::Query<std::collections::HashMap<String, String>>,
644 axum::extract::State(key): axum::extract::State<Option<Arc<String>>>,
645) -> axum::response::Response {
646 use axum::response::IntoResponse;
647
648 if let Some(ref expected) = key {
650 let provided = params.get("introspection-key");
651 if provided.map(|k| k.as_str()) != Some(expected.as_str()) {
652 return (
653 axum::http::StatusCode::FORBIDDEN,
654 "invalid or missing introspection-key",
655 )
656 .into_response();
657 }
658 }
659
660 let mut source = async_graphql::http::GraphiQLSource::build().endpoint("/graphql");
663 if let Some(ref k) = key {
664 source = source.header("X-Introspection-Key", k.as_str());
665 }
666
667 axum::response::Html(source.finish()).into_response()
668}
669
670pub fn cmd_serve_daemon(
672 port: Option<u16>,
673 bind: Option<std::net::IpAddr>,
674 subsonic_port: Option<u16>,
675 playground: bool,
676) {
677 use std::fs;
678 use std::process::Command;
679
680 let cfg = Config::load().unwrap_or_default();
681 let port_val = port.unwrap_or(cfg.graphql.port);
682 let bind_val = bind.unwrap_or(cfg.graphql.bind);
683
684 let exe = std::env::current_exe().expect("failed to get current exe path");
685 let mut cmd = Command::new(exe);
686
687 cmd.arg("--headless");
688 cmd.arg("--port").arg(port_val.to_string());
689 cmd.arg("--bind").arg(bind_val.to_string());
690 if let Some(sp) = subsonic_port {
691 cmd.arg("--subsonic").arg(sp.to_string());
692 }
693 if playground || cfg.graphql.playground {
694 cmd.arg("--playground");
695 }
696
697 cmd.stdin(std::process::Stdio::null());
698 cmd.stdout(std::process::Stdio::null());
699 cmd.stderr(std::process::Stdio::null());
700
701 let mut child = cmd.spawn().expect("failed to spawn daemon process");
702 let pid = child.id();
703
704 let pid_path = koan_core::config::config_dir().join("koan-serve.pid");
705 fs::write(&pid_path, pid.to_string()).ok();
706
707 std::thread::spawn(move || {
708 let _ = child.wait();
709 });
710
711 eprintln!("kōan daemon started (pid {}) on port {}", pid, port_val);
712 if let Some(sp) = subsonic_port {
713 eprintln!(" Subsonic REST on port {}", sp);
714 }
715 eprintln!(" PID file: {}", pid_path.display());
716}
717
718pub async fn execute_in_process(
727 schema: &KoanSchema,
728 query: &str,
729 variables: Option<serde_json::Value>,
730 caller: AuthUser,
731) -> serde_json::Value {
732 let mut request = async_graphql::Request::new(query).data(caller);
733 if let Some(serde_json::Value::Object(map)) = variables {
734 let mut gql_vars = async_graphql::Variables::default();
735 for (k, v) in map {
736 gql_vars.insert(
737 async_graphql::Name::new(&k),
738 async_graphql::Value::from_json(v).unwrap_or(async_graphql::Value::Null),
739 );
740 }
741 request = request.variables(gql_vars);
742 }
743 let response = schema.execute(request).await;
744 serde_json::to_value(&response).unwrap_or(serde_json::Value::Null)
745}
746
747#[cfg(test)]
752mod tests {
753 use super::*;
754 use axum::body::Body;
755 use axum::http::{Request as HttpRequest, StatusCode};
756 use axum::routing::{get, post};
757 use tower::ServiceExt as _;
758
759 fn policy() -> Arc<BrowserPolicy> {
760 Arc::new(BrowserPolicy {
761 origins: vec!["https://music.example.com".into()],
762 hosts: vec!["koan.local".into()],
763 })
764 }
765
766 async fn ok() -> &'static str {
767 "ok"
768 }
769
770 fn routes() -> axum::Router<Arc<BrowserPolicy>> {
771 axum::Router::new()
772 .route("/graphql", post(ok).get(ok))
773 .route("/graphql/ws", get(ok))
774 }
775
776 async fn run_host(req: HttpRequest<Body>) -> StatusCode {
777 let app = routes()
778 .layer(axum::middleware::from_fn_with_state(policy(), host_guard))
779 .with_state(policy());
780 app.oneshot(req).await.unwrap().status()
781 }
782
783 async fn run_browser(req: HttpRequest<Body>) -> StatusCode {
784 let app = routes()
785 .layer(axum::middleware::from_fn_with_state(
786 policy(),
787 browser_guard,
788 ))
789 .with_state(policy());
790 app.oneshot(req).await.unwrap().status()
791 }
792
793 fn json_post(uri: &str) -> axum::http::request::Builder {
794 HttpRequest::post(uri).header(axum::http::header::CONTENT_TYPE, "application/json")
795 }
796
797 #[test]
800 fn host_policy_accepts_loopback_literals_and_configured_names() {
801 let p = policy();
802 assert!(p.host_allowed("localhost:4000"));
803 assert!(p.host_allowed("127.0.0.1:4000"));
804 assert!(p.host_allowed("192.168.1.20:4000"));
805 assert!(p.host_allowed("[::1]:4000"));
806 assert!(p.host_allowed("koan.local"));
807 assert!(p.host_allowed("koan.local:4000"));
808 }
809
810 #[test]
811 fn host_policy_rejects_attacker_controlled_names() {
812 let p = policy();
813 assert!(!p.host_allowed("evil.com"));
814 assert!(!p.host_allowed("rebind.evil.com:4000"));
815 assert!(!p.host_allowed("koan.local.evil.com"));
816 }
817
818 #[tokio::test]
819 async fn host_guard_rejects_foreign_host() {
820 let req = json_post("/graphql")
821 .header(axum::http::header::HOST, "rebind.evil.com")
822 .body(Body::empty())
823 .unwrap();
824 assert_eq!(run_host(req).await, StatusCode::FORBIDDEN);
825 }
826
827 #[tokio::test]
828 async fn host_guard_allows_known_host_and_missing_host() {
829 let req = json_post("/graphql")
830 .header(axum::http::header::HOST, "127.0.0.1:4000")
831 .body(Body::empty())
832 .unwrap();
833 assert_eq!(run_host(req).await, StatusCode::OK);
834
835 let req = json_post("/graphql").body(Body::empty()).unwrap();
836 assert_eq!(run_host(req).await, StatusCode::OK);
837 }
838
839 #[tokio::test]
842 async fn ws_upgrade_from_foreign_origin_is_rejected() {
843 let req = HttpRequest::get("/graphql/ws")
844 .header(axum::http::header::HOST, "127.0.0.1:4000")
845 .header(axum::http::header::ORIGIN, "https://evil.com")
846 .body(Body::empty())
847 .unwrap();
848 assert_eq!(run_browser(req).await, StatusCode::FORBIDDEN);
849 }
850
851 #[tokio::test]
852 async fn ws_upgrade_without_origin_is_allowed() {
853 let req = HttpRequest::get("/graphql/ws")
854 .header(axum::http::header::HOST, "127.0.0.1:4000")
855 .body(Body::empty())
856 .unwrap();
857 assert_eq!(run_browser(req).await, StatusCode::OK);
858 }
859
860 #[tokio::test]
861 async fn configured_and_same_origin_are_allowed() {
862 let req = HttpRequest::get("/graphql/ws")
863 .header(axum::http::header::HOST, "127.0.0.1:4000")
864 .header(axum::http::header::ORIGIN, "https://music.example.com")
865 .body(Body::empty())
866 .unwrap();
867 assert_eq!(run_browser(req).await, StatusCode::OK);
868
869 let req = json_post("/graphql")
871 .header(axum::http::header::HOST, "127.0.0.1:4000")
872 .header(axum::http::header::ORIGIN, "http://127.0.0.1:4000")
873 .body(Body::empty())
874 .unwrap();
875 assert_eq!(run_browser(req).await, StatusCode::OK);
876 }
877
878 #[tokio::test]
881 async fn text_plain_post_is_rejected() {
882 let req = HttpRequest::post("/graphql")
883 .header(axum::http::header::CONTENT_TYPE, "text/plain")
884 .body(Body::from(r#"{"query":"mutation{clearQueue{ok}}"}"#))
885 .unwrap();
886 assert_eq!(run_browser(req).await, StatusCode::UNSUPPORTED_MEDIA_TYPE);
887 }
888
889 #[tokio::test]
890 async fn post_without_content_type_is_rejected() {
891 let req = HttpRequest::post("/graphql").body(Body::empty()).unwrap();
892 assert_eq!(run_browser(req).await, StatusCode::UNSUPPORTED_MEDIA_TYPE);
893 }
894
895 #[tokio::test]
898 async fn load_perimeter_refuses_an_oversized_query_body() {
899 async fn parse(_: async_graphql_axum::GraphQLRequest) -> StatusCode {
900 StatusCode::OK
901 }
902 let app = load_perimeter(axum::Router::new().route("/graphql", post(parse)));
903 let body = |padding: usize| {
906 let chunks = [
907 axum::body::Bytes::from_static(br#"{"query":"{__typename}""#),
908 axum::body::Bytes::from(vec![b' '; padding]),
909 axum::body::Bytes::from_static(b"}"),
910 ];
911 Body::from_stream(tokio_stream::iter(chunks.map(Ok::<_, std::io::Error>)))
912 };
913 let req = json_post("/graphql").body(body(1 << 10)).unwrap();
914 assert_eq!(
915 app.clone().oneshot(req).await.unwrap().status(),
916 StatusCode::OK
917 );
918 let req = json_post("/graphql").body(body(3 << 20)).unwrap();
919 assert_ne!(app.oneshot(req).await.unwrap().status(), StatusCode::OK);
920 }
921
922 #[tokio::test]
923 async fn load_perimeter_passes_requests_and_turns_panics_into_500s() {
924 async fn boom() -> &'static str {
925 panic!("resolver exploded");
926 }
927
928 let app = load_perimeter(
929 axum::Router::new()
930 .route("/graphql", post(ok))
931 .route("/boom", post(boom)),
932 );
933
934 let req = json_post("/graphql").body(Body::empty()).unwrap();
935 assert_eq!(
936 app.clone().oneshot(req).await.unwrap().status(),
937 StatusCode::OK
938 );
939
940 let req = json_post("/boom").body(Body::empty()).unwrap();
942 assert_eq!(
943 app.oneshot(req).await.unwrap().status(),
944 StatusCode::INTERNAL_SERVER_ERROR
945 );
946 }
947
948 #[tokio::test]
949 async fn json_post_is_accepted() {
950 let req = json_post("/graphql").body(Body::empty()).unwrap();
951 assert_eq!(run_browser(req).await, StatusCode::OK);
952
953 let req = HttpRequest::post("/graphql")
954 .header(
955 axum::http::header::CONTENT_TYPE,
956 "application/json; charset=utf-8",
957 )
958 .body(Body::empty())
959 .unwrap();
960 assert_eq!(run_browser(req).await, StatusCode::OK);
961 }
962}