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, LoginRateLimiter, 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 mcp_bind: Option<std::net::SocketAddr>,
27) {
28 use koan_core::player::Player;
29
30 let _db = koan_core::db::connection::Database::open_default().expect("failed to open database");
32 let db_path = koan_core::config::db_path();
33 let pool = Arc::new(Pool::new(db_path.clone()));
34
35 let (state, _timeline, _viz, cmd_tx) = Player::spawn();
36
37 let watched = db_path.clone();
40 koan_core::helpers::spawn_library_watch(db_path, move |running| {
41 if !running {
44 crate::clients::fulfil_from(&watched);
45 if let Ok(db) = koan_core::db::connection::Database::open(&watched) {
46 crate::clients::changed_if_library_moved(&db.conn);
47 }
48 }
49 });
50
51 if let Some(addr) = mcp_bind {
52 match crate::mcp::spawn_http(addr, state.clone(), cmd_tx.clone(), pool.clone()) {
53 Ok(_) => log::info!("MCP over HTTP at http://{addr}/mcp"),
54 Err(e) => {
55 eprintln!("koan: MCP listener on {addr}: {e}");
56 std::process::exit(1);
57 }
58 }
59 }
60
61 if let Err(e) = run_api_blocking(ApiServerOpts {
62 state,
63 cmd_tx,
64 pool,
65 port,
66 bind,
67 subsonic_port,
68 playground,
69 viz: None, }) {
71 eprintln!("koan: {}", e);
72 std::process::exit(1);
73 }
74}
75
76const REQUEST_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
83
84const MAX_INFLIGHT_QUERIES: usize = 64;
88
89const MAX_QUERY_BODY: usize = 2 << 20;
93
94fn load_perimeter<S>(router: axum::Router<S>) -> axum::Router<S>
99where
100 S: Clone + Send + Sync + 'static,
101{
102 router
103 .layer(tower_http::catch_panic::CatchPanicLayer::new())
106 .layer(tower_http::limit::RequestBodyLimitLayer::new(
107 MAX_QUERY_BODY,
108 ))
109 .layer(tower_http::timeout::TimeoutLayer::with_status_code(
110 axum::http::StatusCode::REQUEST_TIMEOUT,
111 REQUEST_TIMEOUT,
112 ))
113 .layer(
117 tower::ServiceBuilder::new()
118 .layer(axum::error_handling::HandleErrorLayer::new(
119 |err: tower::BoxError| async move {
120 if err.is::<tower::load_shed::error::Overloaded>() {
121 (
122 axum::http::StatusCode::SERVICE_UNAVAILABLE,
123 "server at capacity",
124 )
125 } else {
126 (
127 axum::http::StatusCode::INTERNAL_SERVER_ERROR,
128 "internal error",
129 )
130 }
131 },
132 ))
133 .load_shed()
134 .concurrency_limit(MAX_INFLIGHT_QUERIES),
135 )
136}
137
138pub struct ApiServerOpts {
140 pub state: Arc<SharedPlayerState>,
141 pub cmd_tx: Sender<PlayerCommand>,
142 pub pool: Arc<Pool>,
145 pub port: Option<u16>,
146 pub bind: Option<std::net::IpAddr>,
147 pub subsonic_port: Option<u16>,
148 pub playground: bool,
149 pub viz: Option<Arc<VizSnapshot>>,
150}
151
152fn run_api_blocking(opts: ApiServerOpts) -> Result<(), String> {
158 let ApiServerOpts {
159 state,
160 cmd_tx,
161 pool,
162 port,
163 bind,
164 subsonic_port,
165 playground,
166 viz,
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 access_ttl = parse_duration_secs(&cfg.graphql.access_token_ttl).unwrap_or(900);
204 let refresh_ttl = parse_duration_secs(&cfg.graphql.refresh_token_ttl).unwrap_or(2_592_000);
205
206 let introspection_key = if playground_enabled && auth_enabled {
210 Some(Arc::new(auth::random_token().map_err(|e| {
211 format!("failed to generate introspection key: {}", e)
212 })?))
213 } else {
214 None
215 };
216
217 let auth_state = AuthState {
218 public_pem: Arc::new(public_pem.clone()),
219 auth_enabled,
220 introspection_key: introspection_key.clone(),
221 pool: pool.clone(),
222 };
223
224 let auth_route_state = AuthRouteState {
225 pool: pool.clone(),
226 private_pem: Arc::new(private_pem),
227 public_pem: Arc::new(public_pem),
228 access_ttl_secs: access_ttl,
229 refresh_ttl_secs: refresh_ttl,
230 cookie_secure: cfg.graphql.cookie_secure,
231 login_limiter: Arc::new(LoginRateLimiter::default()),
232 };
233
234 let schema = build_schema(state, cmd_tx, pool.path().to_path_buf(), viz);
235
236 if auth_enabled {
237 log::info!(
238 "Auth enabled (Ed25519 JWT, access TTL {}s, refresh TTL {}s)",
239 access_ttl,
240 refresh_ttl
241 );
242 } else {
243 log::info!("Auth disabled — all requests treated as admin");
244 }
245
246 let browser_policy = Arc::new(BrowserPolicy {
247 origins: cfg.graphql.cors_origins.clone(),
248 hosts: cfg.graphql.allowed_hosts.clone(),
249 });
250
251 if cfg.graphql.cors_origins.is_empty() {
252 log::info!("CORS: no origins configured — browsers get no cross-origin access");
253 }
254
255 let rt = tokio::runtime::Runtime::new().expect("failed to create tokio runtime");
256 rt.block_on(async {
257 let query_route = load_perimeter(axum::Router::new().route("/graphql", post(graphql_handler)));
262
263 let gql_app = axum::Router::new()
264 .merge(query_route)
265 .route("/graphql/ws", get(graphql_ws_handler))
266 .layer(axum::middleware::from_fn_with_state(
267 auth_state.clone(),
268 auth_middleware,
269 ))
270 .layer(axum::middleware::from_fn_with_state(
273 browser_policy.clone(),
274 browser_guard,
275 ))
276 .with_state(schema);
277
278 let covers = Arc::new(crate::covers::Covers::in_config_dir());
282 let ui_routes = crate::ui::router(
283 pool.clone(),
284 auth_route_state.clone(),
285 auth_enabled,
286 covers.clone(),
287 cfg.sharing.public_url.clone(),
288 );
289
290 let auth_app = auth_router(auth_route_state);
292
293 let origins: Vec<axum::http::HeaderValue> = cfg
296 .graphql
297 .cors_origins
298 .iter()
299 .filter_map(|o| o.parse().ok())
300 .collect();
301 let cors = tower_http::cors::CorsLayer::new()
302 .allow_origin(origins)
303 .allow_methods([
304 axum::http::Method::GET,
305 axum::http::Method::POST,
306 axum::http::Method::OPTIONS,
307 ])
308 .allow_headers([
309 axum::http::header::AUTHORIZATION,
310 axum::http::header::CONTENT_TYPE,
311 axum::http::HeaderName::from_static("x-introspection-key"),
312 ])
313 .allow_credentials(true);
314
315
316 let share_routes = crate::share::router(
320 pool.clone(),
321 cfg.sharing.public_url.clone(),
322 covers.clone(),
323 )
324 .merge(crate::push::router(pool.clone(), covers.clone()));
325 let subsonic_merged = crate::subsonic::subsonic_router(pool, covers);
330 let subsonic_on_main = subsonic_merged.is_some();
331 let subsonic_dedicated = subsonic_merged.clone();
332
333 let mut app = auth_app.merge(gql_app).merge(share_routes).merge(ui_routes);
334 if let Some(sub) = subsonic_merged {
335 app = app.merge(sub);
336 }
337 if playground_enabled {
338 app = app.route(
339 "/graphql",
340 get(graphql_playground).with_state(introspection_key.clone()),
341 );
342 }
343 let app = app.layer(cors).layer(axum::middleware::from_fn_with_state(
347 browser_policy.clone(),
348 host_guard,
349 ));
350
351 let playground_url = if playground_enabled {
353 if let Some(ref key) = introspection_key {
354 format!("http://{}:{}/graphql?introspection-key={}", bind, port, key)
355 } else {
356 format!("http://{}:{}/graphql", bind, port)
357 }
358 } else {
359 format!("http://{}:{}/graphql", bind, port)
360 };
361
362 let gql_addr = std::net::SocketAddr::new(bind, port);
363
364 let gql_listener = match tokio::net::TcpListener::bind(gql_addr).await {
365 Ok(l) => {
366 log::info!("GraphQL API on http://{}:{}/graphql", bind, port);
367 if subsonic_on_main {
368 log::info!("Subsonic REST on http://{}:{}/rest/", bind, port);
369 }
370 if playground_enabled {
371 log::info!("GraphiQL: {}", playground_url);
372 #[cfg(target_os = "macos")]
374 let _ = std::process::Command::new("open").arg(&playground_url).spawn();
375 #[cfg(target_os = "linux")]
376 let _ = std::process::Command::new("xdg-open").arg(&playground_url).spawn();
377 }
378 l
379 }
380 Err(e) => {
381 return Err(format!(
382 "failed to bind GraphQL port {port} — {e} (another instance running?)"
383 ));
384 }
385 };
386 let gql_server = axum::serve(
387 gql_listener,
388 app.into_make_service_with_connect_info::<std::net::SocketAddr>(),
389 )
390 .with_graceful_shutdown(shutdown_signal());
391
392 let extra_sub_port = subsonic_port.filter(|p| *p != port);
395 if let Some(sub_port) = extra_sub_port
396 && let Some(sub_app) = subsonic_dedicated
397 {
398 let sub_addr = std::net::SocketAddr::new(bind, sub_port);
399 match tokio::net::TcpListener::bind(sub_addr).await {
400 Ok(sub_listener) => {
401 log::info!(
402 "Subsonic REST also on http://{}:{}/rest/ (dedicated port)",
403 bind,
404 sub_port,
405 );
406 let sub_server = axum::serve(
409 sub_listener,
410 sub_app.into_make_service_with_connect_info::<std::net::SocketAddr>(),
411 )
412 .with_graceful_shutdown(shutdown_signal());
413
414 tokio::select! {
415 r = gql_server => { if let Err(e) = r { log::error!("GraphQL server error: {e}"); } },
416 r = sub_server => { if let Err(e) = r { log::error!("Subsonic server error: {e}"); } },
417 }
418 return Ok(());
419 }
420 Err(e) => {
421 log::warn!(
422 "Dedicated Subsonic port {} unavailable — {}. Mounted on GraphQL port only.",
423 sub_port,
424 e,
425 );
426 }
427 }
428 }
429
430 if let Err(e) = gql_server.await {
431 log::error!("GraphQL server error: {e}");
432 }
433 Ok(())
434 })
435}
436
437pub fn start_api_background(
440 state: Arc<SharedPlayerState>,
441 cmd_tx: Sender<PlayerCommand>,
442 db_path: PathBuf,
443 port: Option<u16>,
444 bind: Option<std::net::IpAddr>,
445 subsonic_port: Option<u16>,
446 playground: bool,
447) {
448 if let Err(e) = run_api_blocking(ApiServerOpts {
451 state,
452 cmd_tx,
453 pool: Arc::new(Pool::new(db_path)),
454 port,
455 bind,
456 subsonic_port,
457 playground,
458 viz: None,
459 }) {
460 log::error!("API server not started: {}", e);
461 }
462}
463
464pub(crate) struct BrowserPolicy {
473 origins: Vec<String>,
474 hosts: Vec<String>,
475}
476
477impl BrowserPolicy {
478 fn host_allowed(&self, host: &str) -> bool {
479 if self.hosts.iter().any(|h| h.eq_ignore_ascii_case(host)) {
480 return true;
481 }
482 let bare = strip_port(host);
483 if self.hosts.iter().any(|h| h.eq_ignore_ascii_case(bare)) {
484 return true;
485 }
486 bare.eq_ignore_ascii_case("localhost") || bare.parse::<std::net::IpAddr>().is_ok()
489 }
490
491 fn origin_allowed(&self, origin: &str, host: Option<&str>) -> bool {
494 if self.origins.iter().any(|o| o == origin) {
495 return true;
496 }
497 match (origin.split_once("://"), host) {
498 (Some((_, authority)), Some(host)) => authority.eq_ignore_ascii_case(host),
499 _ => false,
500 }
501 }
502}
503
504fn strip_port(host: &str) -> &str {
506 if let Some(rest) = host.strip_prefix('[') {
507 return rest.split(']').next().unwrap_or(rest);
508 }
509 match host.rsplit_once(':') {
510 Some((h, port)) if !port.is_empty() && port.bytes().all(|b| b.is_ascii_digit()) => h,
511 _ => host,
512 }
513}
514
515fn header_str(request: &axum::extract::Request, name: axum::http::HeaderName) -> Option<&str> {
516 request.headers().get(name).and_then(|v| v.to_str().ok())
517}
518
519async fn host_guard(
521 axum::extract::State(policy): axum::extract::State<Arc<BrowserPolicy>>,
522 request: axum::extract::Request,
523 next: axum::middleware::Next,
524) -> axum::response::Response {
525 use axum::response::IntoResponse;
526
527 let host = header_str(&request, axum::http::header::HOST)
530 .map(str::to_owned)
531 .or_else(|| request.uri().host().map(str::to_owned));
532
533 if let Some(ref host) = host
534 && !policy.host_allowed(host)
535 {
536 log::warn!("rejected request for unrecognised Host: {}", host);
537 return (axum::http::StatusCode::FORBIDDEN, "host not allowed").into_response();
538 }
539
540 next.run(request).await
541}
542
543async fn browser_guard(
552 axum::extract::State(policy): axum::extract::State<Arc<BrowserPolicy>>,
553 request: axum::extract::Request,
554 next: axum::middleware::Next,
555) -> axum::response::Response {
556 use axum::response::IntoResponse;
557
558 let host = header_str(&request, axum::http::header::HOST).map(str::to_owned);
559 if let Some(origin) = header_str(&request, axum::http::header::ORIGIN)
561 && !policy.origin_allowed(origin, host.as_deref())
562 {
563 log::warn!(
564 "rejected GraphQL request from disallowed Origin: {}",
565 origin
566 );
567 return (axum::http::StatusCode::FORBIDDEN, "origin not allowed").into_response();
568 }
569
570 if request.method() == axum::http::Method::POST && !is_graphql_content_type(&request) {
571 return (
572 axum::http::StatusCode::UNSUPPORTED_MEDIA_TYPE,
573 "content type must be application/json or application/graphql",
574 )
575 .into_response();
576 }
577
578 next.run(request).await
579}
580
581fn is_graphql_content_type(request: &axum::extract::Request) -> bool {
582 header_str(request, axum::http::header::CONTENT_TYPE).is_some_and(|ct| {
583 let ct = ct.trim().to_ascii_lowercase();
584 ct.starts_with("application/json") || ct.starts_with("application/graphql")
585 })
586}
587
588async fn shutdown_signal() {
589 tokio::signal::ctrl_c()
590 .await
591 .expect("failed to listen for ctrl+c");
592}
593
594async fn graphql_handler(
595 axum::Extension(user): axum::Extension<AuthUser>,
596 axum::extract::State(schema): axum::extract::State<KoanSchema>,
597 headers: axum::http::HeaderMap,
598 req: async_graphql_axum::GraphQLRequest,
599) -> async_graphql_axum::GraphQLResponse {
600 let mut request = req.into_inner();
601 if let Some(origin) = crate::origin::origin(&headers, None) {
602 request = request.data(super::RequestOrigin(origin));
603 }
604 request = request.data(user);
607 schema.execute(request).await.into()
608}
609
610async fn graphql_ws_handler(
611 axum::Extension(user): axum::Extension<AuthUser>,
612 axum::extract::State(schema): axum::extract::State<KoanSchema>,
613 protocol: async_graphql_axum::GraphQLProtocol,
614 websocket: axum::extract::WebSocketUpgrade,
615) -> axum::response::Response {
616 websocket
617 .protocols(async_graphql::http::ALL_WEBSOCKET_PROTOCOLS)
618 .on_upgrade(move |stream| {
619 let stream = async_graphql_axum::GraphQLWebSocket::new(stream, schema, protocol)
620 .on_connection_init(move |_| async move {
621 let mut data = async_graphql::Data::default();
622 data.insert(user);
623 Ok(data)
624 });
625 async move {
626 stream.serve().await;
627 }
628 })
629}
630
631async fn graphql_playground(
632 axum::extract::Query(params): axum::extract::Query<std::collections::HashMap<String, String>>,
633 axum::extract::State(key): axum::extract::State<Option<Arc<String>>>,
634) -> axum::response::Response {
635 use axum::response::IntoResponse;
636
637 if let Some(ref expected) = key {
639 let provided = params.get("introspection-key");
640 if provided.map(|k| k.as_str()) != Some(expected.as_str()) {
641 return (
642 axum::http::StatusCode::FORBIDDEN,
643 "invalid or missing introspection-key",
644 )
645 .into_response();
646 }
647 }
648
649 let mut source = async_graphql::http::GraphiQLSource::build().endpoint("/graphql");
652 if let Some(ref k) = key {
653 source = source.header("X-Introspection-Key", k.as_str());
654 }
655
656 axum::response::Html(source.finish()).into_response()
657}
658
659pub fn cmd_serve_daemon(
661 port: Option<u16>,
662 bind: Option<std::net::IpAddr>,
663 subsonic_port: Option<u16>,
664 playground: bool,
665 mcp_bind: Option<std::net::SocketAddr>,
666) {
667 use std::fs;
668 use std::process::Command;
669
670 let cfg = Config::load().unwrap_or_default();
671 let port_val = port.unwrap_or(cfg.graphql.port);
672 let bind_val = bind.unwrap_or(cfg.graphql.bind);
673
674 let exe = std::env::current_exe().expect("failed to get current exe path");
675 let mut cmd = Command::new(exe);
676
677 cmd.arg("--headless");
678 cmd.arg("--port").arg(port_val.to_string());
679 cmd.arg("--bind").arg(bind_val.to_string());
680 if let Some(sp) = subsonic_port {
681 cmd.arg("--subsonic").arg(sp.to_string());
682 }
683 if let Some(addr) = mcp_bind {
684 cmd.arg("--mcp-bind").arg(addr.to_string());
685 }
686 if playground || cfg.graphql.playground {
687 cmd.arg("--playground");
688 }
689
690 cmd.stdin(std::process::Stdio::null());
691 cmd.stdout(std::process::Stdio::null());
692 cmd.stderr(std::process::Stdio::null());
693
694 let mut child = cmd.spawn().expect("failed to spawn daemon process");
695 let pid = child.id();
696
697 let pid_path = koan_core::config::config_dir().join("koan-serve.pid");
698 fs::write(&pid_path, pid.to_string()).ok();
699
700 std::thread::spawn(move || {
701 let _ = child.wait();
702 });
703
704 eprintln!("kōan daemon started (pid {}) on port {}", pid, port_val);
705 if let Some(sp) = subsonic_port {
706 eprintln!(" Subsonic REST on port {}", sp);
707 }
708 eprintln!(" PID file: {}", pid_path.display());
709}
710
711pub async fn execute_in_process(
720 schema: &KoanSchema,
721 query: &str,
722 variables: Option<serde_json::Value>,
723 caller: AuthUser,
724) -> serde_json::Value {
725 let mut request = async_graphql::Request::new(query).data(caller);
726 if let Some(serde_json::Value::Object(map)) = variables {
727 let mut gql_vars = async_graphql::Variables::default();
728 for (k, v) in map {
729 gql_vars.insert(
730 async_graphql::Name::new(&k),
731 async_graphql::Value::from_json(v).unwrap_or(async_graphql::Value::Null),
732 );
733 }
734 request = request.variables(gql_vars);
735 }
736 let response = schema.execute(request).await;
737 serde_json::to_value(&response).unwrap_or(serde_json::Value::Null)
738}
739
740#[cfg(test)]
745mod tests {
746 use super::*;
747 use axum::body::Body;
748 use axum::http::{Request as HttpRequest, StatusCode};
749 use axum::routing::{get, post};
750 use tower::ServiceExt as _;
751
752 fn policy() -> Arc<BrowserPolicy> {
753 Arc::new(BrowserPolicy {
754 origins: vec!["https://music.example.com".into()],
755 hosts: vec!["koan.local".into()],
756 })
757 }
758
759 async fn ok() -> &'static str {
760 "ok"
761 }
762
763 fn routes() -> axum::Router<Arc<BrowserPolicy>> {
764 axum::Router::new()
765 .route("/graphql", post(ok).get(ok))
766 .route("/graphql/ws", get(ok))
767 }
768
769 async fn run_host(req: HttpRequest<Body>) -> StatusCode {
770 let app = routes()
771 .layer(axum::middleware::from_fn_with_state(policy(), host_guard))
772 .with_state(policy());
773 app.oneshot(req).await.unwrap().status()
774 }
775
776 async fn run_browser(req: HttpRequest<Body>) -> StatusCode {
777 let app = routes()
778 .layer(axum::middleware::from_fn_with_state(
779 policy(),
780 browser_guard,
781 ))
782 .with_state(policy());
783 app.oneshot(req).await.unwrap().status()
784 }
785
786 fn json_post(uri: &str) -> axum::http::request::Builder {
787 HttpRequest::post(uri).header(axum::http::header::CONTENT_TYPE, "application/json")
788 }
789
790 #[test]
793 fn host_policy_accepts_loopback_literals_and_configured_names() {
794 let p = policy();
795 assert!(p.host_allowed("localhost:4000"));
796 assert!(p.host_allowed("127.0.0.1:4000"));
797 assert!(p.host_allowed("192.168.1.20:4000"));
798 assert!(p.host_allowed("[::1]:4000"));
799 assert!(p.host_allowed("koan.local"));
800 assert!(p.host_allowed("koan.local:4000"));
801 }
802
803 #[test]
804 fn host_policy_rejects_attacker_controlled_names() {
805 let p = policy();
806 assert!(!p.host_allowed("evil.com"));
807 assert!(!p.host_allowed("rebind.evil.com:4000"));
808 assert!(!p.host_allowed("koan.local.evil.com"));
809 }
810
811 #[tokio::test]
812 async fn host_guard_rejects_foreign_host() {
813 let req = json_post("/graphql")
814 .header(axum::http::header::HOST, "rebind.evil.com")
815 .body(Body::empty())
816 .unwrap();
817 assert_eq!(run_host(req).await, StatusCode::FORBIDDEN);
818 }
819
820 #[tokio::test]
821 async fn host_guard_allows_known_host_and_missing_host() {
822 let req = json_post("/graphql")
823 .header(axum::http::header::HOST, "127.0.0.1:4000")
824 .body(Body::empty())
825 .unwrap();
826 assert_eq!(run_host(req).await, StatusCode::OK);
827
828 let req = json_post("/graphql").body(Body::empty()).unwrap();
829 assert_eq!(run_host(req).await, StatusCode::OK);
830 }
831
832 #[tokio::test]
835 async fn ws_upgrade_from_foreign_origin_is_rejected() {
836 let req = HttpRequest::get("/graphql/ws")
837 .header(axum::http::header::HOST, "127.0.0.1:4000")
838 .header(axum::http::header::ORIGIN, "https://evil.com")
839 .body(Body::empty())
840 .unwrap();
841 assert_eq!(run_browser(req).await, StatusCode::FORBIDDEN);
842 }
843
844 #[tokio::test]
845 async fn ws_upgrade_without_origin_is_allowed() {
846 let req = HttpRequest::get("/graphql/ws")
847 .header(axum::http::header::HOST, "127.0.0.1:4000")
848 .body(Body::empty())
849 .unwrap();
850 assert_eq!(run_browser(req).await, StatusCode::OK);
851 }
852
853 #[tokio::test]
854 async fn configured_and_same_origin_are_allowed() {
855 let req = HttpRequest::get("/graphql/ws")
856 .header(axum::http::header::HOST, "127.0.0.1:4000")
857 .header(axum::http::header::ORIGIN, "https://music.example.com")
858 .body(Body::empty())
859 .unwrap();
860 assert_eq!(run_browser(req).await, StatusCode::OK);
861
862 let req = json_post("/graphql")
864 .header(axum::http::header::HOST, "127.0.0.1:4000")
865 .header(axum::http::header::ORIGIN, "http://127.0.0.1:4000")
866 .body(Body::empty())
867 .unwrap();
868 assert_eq!(run_browser(req).await, StatusCode::OK);
869 }
870
871 #[tokio::test]
874 async fn text_plain_post_is_rejected() {
875 let req = HttpRequest::post("/graphql")
876 .header(axum::http::header::CONTENT_TYPE, "text/plain")
877 .body(Body::from(r#"{"query":"mutation{clearQueue{ok}}"}"#))
878 .unwrap();
879 assert_eq!(run_browser(req).await, StatusCode::UNSUPPORTED_MEDIA_TYPE);
880 }
881
882 #[tokio::test]
883 async fn post_without_content_type_is_rejected() {
884 let req = HttpRequest::post("/graphql").body(Body::empty()).unwrap();
885 assert_eq!(run_browser(req).await, StatusCode::UNSUPPORTED_MEDIA_TYPE);
886 }
887
888 #[tokio::test]
891 async fn load_perimeter_refuses_an_oversized_query_body() {
892 async fn parse(_: async_graphql_axum::GraphQLRequest) -> StatusCode {
893 StatusCode::OK
894 }
895 let app = load_perimeter(axum::Router::new().route("/graphql", post(parse)));
896 let body = |padding: usize| {
899 let chunks = [
900 axum::body::Bytes::from_static(br#"{"query":"{__typename}""#),
901 axum::body::Bytes::from(vec![b' '; padding]),
902 axum::body::Bytes::from_static(b"}"),
903 ];
904 Body::from_stream(tokio_stream::iter(chunks.map(Ok::<_, std::io::Error>)))
905 };
906 let req = json_post("/graphql").body(body(1 << 10)).unwrap();
907 assert_eq!(
908 app.clone().oneshot(req).await.unwrap().status(),
909 StatusCode::OK
910 );
911 let req = json_post("/graphql").body(body(3 << 20)).unwrap();
912 assert_ne!(app.oneshot(req).await.unwrap().status(), StatusCode::OK);
913 }
914
915 #[tokio::test]
916 async fn load_perimeter_passes_requests_and_turns_panics_into_500s() {
917 async fn boom() -> &'static str {
918 panic!("resolver exploded");
919 }
920
921 let app = load_perimeter(
922 axum::Router::new()
923 .route("/graphql", post(ok))
924 .route("/boom", post(boom)),
925 );
926
927 let req = json_post("/graphql").body(Body::empty()).unwrap();
928 assert_eq!(
929 app.clone().oneshot(req).await.unwrap().status(),
930 StatusCode::OK
931 );
932
933 let req = json_post("/boom").body(Body::empty()).unwrap();
935 assert_eq!(
936 app.oneshot(req).await.unwrap().status(),
937 StatusCode::INTERNAL_SERVER_ERROR
938 );
939 }
940
941 #[tokio::test]
942 async fn json_post_is_accepted() {
943 let req = json_post("/graphql").body(Body::empty()).unwrap();
944 assert_eq!(run_browser(req).await, StatusCode::OK);
945
946 let req = HttpRequest::post("/graphql")
947 .header(
948 axum::http::header::CONTENT_TYPE,
949 "application/json; charset=utf-8",
950 )
951 .body(Body::empty())
952 .unwrap();
953 assert_eq!(run_browser(req).await, StatusCode::OK);
954 }
955}