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
297 .graphql
298 .cors_origins
299 .iter()
300 .filter_map(|o| o.parse().ok())
301 .collect();
302 let cors = tower_http::cors::CorsLayer::new()
303 .allow_origin(origins)
304 .allow_methods([
305 axum::http::Method::GET,
306 axum::http::Method::POST,
307 axum::http::Method::OPTIONS,
308 ])
309 .allow_headers([
310 axum::http::header::AUTHORIZATION,
311 axum::http::header::CONTENT_TYPE,
312 axum::http::HeaderName::from_static("x-introspection-key"),
313 ])
314 .allow_credentials(true);
315
316 let share_routes = crate::share::router(
326 pool.clone(),
327 cfg.sharing.public_url.clone(),
328 covers.clone(),
329 )
330 .merge(crate::push::router(pool.clone(), covers.clone()));
331 let subsonic_merged = crate::subsonic::subsonic_router(pool, covers);
332 let subsonic_on_main = subsonic_merged.is_some();
333 let subsonic_dedicated = subsonic_merged.clone();
334
335 let mut app = auth_app.merge(gql_app).merge(share_routes).merge(ui_routes);
336 if let Some(sub) = subsonic_merged {
337 app = app.merge(sub);
338 }
339 if playground_enabled {
340 app = app.route(
341 "/graphql",
342 get(graphql_playground).with_state(introspection_key.clone()),
343 );
344 }
345 let app = app.layer(cors).layer(axum::middleware::from_fn_with_state(
349 browser_policy.clone(),
350 host_guard,
351 ));
352
353 let playground_url = if playground_enabled {
355 if let Some(ref key) = introspection_key {
356 format!("http://{}:{}/graphql?introspection-key={}", bind, port, key)
357 } else {
358 format!("http://{}:{}/graphql", bind, port)
359 }
360 } else {
361 format!("http://{}:{}/graphql", bind, port)
362 };
363
364 let gql_addr = std::net::SocketAddr::new(bind, port);
365
366 let gql_listener = match tokio::net::TcpListener::bind(gql_addr).await {
367 Ok(l) => {
368 log::info!("GraphQL API on http://{}:{}/graphql", bind, port);
369 if subsonic_on_main {
370 log::info!("Subsonic REST on http://{}:{}/rest/", bind, port);
371 }
372 if playground_enabled {
373 log::info!("GraphiQL: {}", playground_url);
374 #[cfg(target_os = "macos")]
376 let _ = std::process::Command::new("open").arg(&playground_url).spawn();
377 #[cfg(target_os = "linux")]
378 let _ = std::process::Command::new("xdg-open").arg(&playground_url).spawn();
379 }
380 l
381 }
382 Err(e) => {
383 return Err(format!(
384 "failed to bind GraphQL port {port} — {e} (another instance running?)"
385 ));
386 }
387 };
388 let gql_server = axum::serve(
389 gql_listener,
390 app.into_make_service_with_connect_info::<std::net::SocketAddr>(),
391 )
392 .with_graceful_shutdown(shutdown_signal());
393
394 let extra_sub_port = subsonic_port.filter(|p| *p != port);
398 if let Some(sub_port) = extra_sub_port
399 && let Some(sub_app) = subsonic_dedicated
400 {
401 let sub_addr = std::net::SocketAddr::new(bind, sub_port);
402 match tokio::net::TcpListener::bind(sub_addr).await {
403 Ok(sub_listener) => {
404 log::info!(
405 "Subsonic REST also on http://{}:{}/rest/ (dedicated port)",
406 bind,
407 sub_port,
408 );
409 let sub_server = axum::serve(
412 sub_listener,
413 sub_app.into_make_service_with_connect_info::<std::net::SocketAddr>(),
414 )
415 .with_graceful_shutdown(shutdown_signal());
416
417 tokio::select! {
418 r = gql_server => { if let Err(e) = r { log::error!("GraphQL server error: {e}"); } },
419 r = sub_server => { if let Err(e) = r { log::error!("Subsonic server error: {e}"); } },
420 }
421 return Ok(());
422 }
423 Err(e) => {
424 log::warn!(
425 "Dedicated Subsonic port {} unavailable — {}. Mounted on GraphQL port only.",
426 sub_port,
427 e,
428 );
429 }
430 }
431 }
432
433 if let Err(e) = gql_server.await {
434 log::error!("GraphQL server error: {e}");
435 }
436 Ok(())
437 })
438}
439
440pub fn start_api_background(
446 state: Arc<SharedPlayerState>,
447 cmd_tx: Sender<PlayerCommand>,
448 db_path: PathBuf,
449 port: Option<u16>,
450 bind: Option<std::net::IpAddr>,
451 subsonic_port: Option<u16>,
452 playground: bool,
453) {
454 if let Err(e) = run_api_blocking(ApiServerOpts {
457 state,
458 cmd_tx,
459 pool: Arc::new(Pool::new(db_path)),
460 port,
461 bind,
462 subsonic_port,
463 playground,
464 viz: None,
465 }) {
466 log::error!("API server not started: {}", e);
467 }
468}
469
470pub(crate) struct BrowserPolicy {
479 origins: Vec<String>,
480 hosts: Vec<String>,
481}
482
483impl BrowserPolicy {
484 fn host_allowed(&self, host: &str) -> bool {
485 if self.hosts.iter().any(|h| h.eq_ignore_ascii_case(host)) {
486 return true;
487 }
488 let bare = strip_port(host);
489 if self.hosts.iter().any(|h| h.eq_ignore_ascii_case(bare)) {
490 return true;
491 }
492 bare.eq_ignore_ascii_case("localhost") || bare.parse::<std::net::IpAddr>().is_ok()
495 }
496
497 fn origin_allowed(&self, origin: &str, host: Option<&str>) -> bool {
500 if self.origins.iter().any(|o| o == origin) {
501 return true;
502 }
503 match (origin.split_once("://"), host) {
504 (Some((_, authority)), Some(host)) => authority.eq_ignore_ascii_case(host),
505 _ => false,
506 }
507 }
508}
509
510fn strip_port(host: &str) -> &str {
512 if let Some(rest) = host.strip_prefix('[') {
513 return rest.split(']').next().unwrap_or(rest);
514 }
515 match host.rsplit_once(':') {
516 Some((h, port)) if !port.is_empty() && port.bytes().all(|b| b.is_ascii_digit()) => h,
517 _ => host,
518 }
519}
520
521fn header_str(request: &axum::extract::Request, name: axum::http::HeaderName) -> Option<&str> {
522 request.headers().get(name).and_then(|v| v.to_str().ok())
523}
524
525async fn host_guard(
527 axum::extract::State(policy): axum::extract::State<Arc<BrowserPolicy>>,
528 request: axum::extract::Request,
529 next: axum::middleware::Next,
530) -> axum::response::Response {
531 use axum::response::IntoResponse;
532
533 let host = header_str(&request, axum::http::header::HOST)
536 .map(str::to_owned)
537 .or_else(|| request.uri().host().map(str::to_owned));
538
539 if let Some(ref host) = host
540 && !policy.host_allowed(host)
541 {
542 log::warn!("rejected request for unrecognised Host: {}", host);
543 return (axum::http::StatusCode::FORBIDDEN, "host not allowed").into_response();
544 }
545
546 next.run(request).await
547}
548
549async fn browser_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).map(str::to_owned);
565 if let Some(origin) = header_str(&request, axum::http::header::ORIGIN)
567 && !policy.origin_allowed(origin, host.as_deref())
568 {
569 log::warn!(
570 "rejected GraphQL request from disallowed Origin: {}",
571 origin
572 );
573 return (axum::http::StatusCode::FORBIDDEN, "origin not allowed").into_response();
574 }
575
576 if request.method() == axum::http::Method::POST && !is_graphql_content_type(&request) {
577 return (
578 axum::http::StatusCode::UNSUPPORTED_MEDIA_TYPE,
579 "content type must be application/json or application/graphql",
580 )
581 .into_response();
582 }
583
584 next.run(request).await
585}
586
587fn is_graphql_content_type(request: &axum::extract::Request) -> bool {
588 header_str(request, axum::http::header::CONTENT_TYPE).is_some_and(|ct| {
589 let ct = ct.trim().to_ascii_lowercase();
590 ct.starts_with("application/json") || ct.starts_with("application/graphql")
591 })
592}
593
594async fn shutdown_signal() {
595 tokio::signal::ctrl_c()
596 .await
597 .expect("failed to listen for ctrl+c");
598}
599
600async fn graphql_handler(
601 axum::Extension(user): axum::Extension<AuthUser>,
602 axum::extract::State(schema): axum::extract::State<KoanSchema>,
603 headers: axum::http::HeaderMap,
604 req: async_graphql_axum::GraphQLRequest,
605) -> async_graphql_axum::GraphQLResponse {
606 let mut request = req.into_inner();
607 if let Some(origin) = crate::origin::origin(&headers, None) {
608 request = request.data(super::RequestOrigin(origin));
609 }
610 request = request.data(user);
613 schema.execute(request).await.into()
614}
615
616async fn graphql_ws_handler(
617 axum::Extension(user): axum::Extension<AuthUser>,
618 axum::extract::State(schema): axum::extract::State<KoanSchema>,
619 protocol: async_graphql_axum::GraphQLProtocol,
620 websocket: axum::extract::WebSocketUpgrade,
621) -> axum::response::Response {
622 websocket
623 .protocols(async_graphql::http::ALL_WEBSOCKET_PROTOCOLS)
624 .on_upgrade(move |stream| {
625 let stream = async_graphql_axum::GraphQLWebSocket::new(stream, schema, protocol)
626 .on_connection_init(move |_| async move {
627 let mut data = async_graphql::Data::default();
628 data.insert(user);
629 Ok(data)
630 });
631 async move {
632 stream.serve().await;
633 }
634 })
635}
636
637async fn graphql_playground(
638 axum::extract::Query(params): axum::extract::Query<std::collections::HashMap<String, String>>,
639 axum::extract::State(key): axum::extract::State<Option<Arc<String>>>,
640) -> axum::response::Response {
641 use axum::response::IntoResponse;
642
643 if let Some(ref expected) = key {
645 let provided = params.get("introspection-key");
646 if provided.map(|k| k.as_str()) != Some(expected.as_str()) {
647 return (
648 axum::http::StatusCode::FORBIDDEN,
649 "invalid or missing introspection-key",
650 )
651 .into_response();
652 }
653 }
654
655 let mut source = async_graphql::http::GraphiQLSource::build().endpoint("/graphql");
658 if let Some(ref k) = key {
659 source = source.header("X-Introspection-Key", k.as_str());
660 }
661
662 axum::response::Html(source.finish()).into_response()
663}
664
665pub fn cmd_serve_daemon(
667 port: Option<u16>,
668 bind: Option<std::net::IpAddr>,
669 subsonic_port: Option<u16>,
670 playground: bool,
671 mcp_bind: Option<std::net::SocketAddr>,
672) {
673 use std::fs;
674 use std::process::Command;
675
676 let cfg = Config::load().unwrap_or_default();
677 let port_val = port.unwrap_or(cfg.graphql.port);
678 let bind_val = bind.unwrap_or(cfg.graphql.bind);
679
680 let exe = std::env::current_exe().expect("failed to get current exe path");
681 let mut cmd = Command::new(exe);
682 cmd.arg("--headless");
684 cmd.arg("--port").arg(port_val.to_string());
685 cmd.arg("--bind").arg(bind_val.to_string());
686 if let Some(sp) = subsonic_port {
687 cmd.arg("--subsonic").arg(sp.to_string());
688 }
689 if let Some(addr) = mcp_bind {
690 cmd.arg("--mcp-bind").arg(addr.to_string());
691 }
692 if playground || cfg.graphql.playground {
693 cmd.arg("--playground");
694 }
695
696 cmd.stdin(std::process::Stdio::null());
697 cmd.stdout(std::process::Stdio::null());
698 cmd.stderr(std::process::Stdio::null());
699
700 let mut child = cmd.spawn().expect("failed to spawn daemon process");
701 let pid = child.id();
702
703 let pid_path = koan_core::config::config_dir().join("koan-serve.pid");
704 fs::write(&pid_path, pid.to_string()).ok();
705
706 std::thread::spawn(move || {
707 let _ = child.wait();
708 });
709
710 eprintln!("kōan daemon started (pid {}) on port {}", pid, port_val);
711 if let Some(sp) = subsonic_port {
712 eprintln!(" Subsonic REST on port {}", sp);
713 }
714 eprintln!(" PID file: {}", pid_path.display());
715}
716
717pub async fn execute_in_process(
726 schema: &KoanSchema,
727 query: &str,
728 variables: Option<serde_json::Value>,
729 caller: AuthUser,
730) -> serde_json::Value {
731 let mut request = async_graphql::Request::new(query).data(caller);
732 if let Some(serde_json::Value::Object(map)) = variables {
733 let mut gql_vars = async_graphql::Variables::default();
734 for (k, v) in map {
735 gql_vars.insert(
736 async_graphql::Name::new(&k),
737 async_graphql::Value::from_json(v).unwrap_or(async_graphql::Value::Null),
738 );
739 }
740 request = request.variables(gql_vars);
741 }
742 let response = schema.execute(request).await;
743 serde_json::to_value(&response).unwrap_or(serde_json::Value::Null)
744}
745
746#[cfg(test)]
751mod tests {
752 use super::*;
753 use axum::body::Body;
754 use axum::http::{Request as HttpRequest, StatusCode};
755 use axum::routing::{get, post};
756 use tower::ServiceExt as _;
757
758 fn policy() -> Arc<BrowserPolicy> {
759 Arc::new(BrowserPolicy {
760 origins: vec!["https://music.example.com".into()],
761 hosts: vec!["koan.local".into()],
762 })
763 }
764
765 async fn ok() -> &'static str {
766 "ok"
767 }
768
769 fn routes() -> axum::Router<Arc<BrowserPolicy>> {
770 axum::Router::new()
771 .route("/graphql", post(ok).get(ok))
772 .route("/graphql/ws", get(ok))
773 }
774
775 async fn run_host(req: HttpRequest<Body>) -> StatusCode {
776 let app = routes()
777 .layer(axum::middleware::from_fn_with_state(policy(), host_guard))
778 .with_state(policy());
779 app.oneshot(req).await.unwrap().status()
780 }
781
782 async fn run_browser(req: HttpRequest<Body>) -> StatusCode {
783 let app = routes()
784 .layer(axum::middleware::from_fn_with_state(
785 policy(),
786 browser_guard,
787 ))
788 .with_state(policy());
789 app.oneshot(req).await.unwrap().status()
790 }
791
792 fn json_post(uri: &str) -> axum::http::request::Builder {
793 HttpRequest::post(uri).header(axum::http::header::CONTENT_TYPE, "application/json")
794 }
795
796 #[test]
799 fn host_policy_accepts_loopback_literals_and_configured_names() {
800 let p = policy();
801 assert!(p.host_allowed("localhost:4000"));
802 assert!(p.host_allowed("127.0.0.1:4000"));
803 assert!(p.host_allowed("192.168.1.20:4000"));
804 assert!(p.host_allowed("[::1]:4000"));
805 assert!(p.host_allowed("koan.local"));
806 assert!(p.host_allowed("koan.local:4000"));
807 }
808
809 #[test]
810 fn host_policy_rejects_attacker_controlled_names() {
811 let p = policy();
812 assert!(!p.host_allowed("evil.com"));
813 assert!(!p.host_allowed("rebind.evil.com:4000"));
814 assert!(!p.host_allowed("koan.local.evil.com"));
815 }
816
817 #[tokio::test]
818 async fn host_guard_rejects_foreign_host() {
819 let req = json_post("/graphql")
820 .header(axum::http::header::HOST, "rebind.evil.com")
821 .body(Body::empty())
822 .unwrap();
823 assert_eq!(run_host(req).await, StatusCode::FORBIDDEN);
824 }
825
826 #[tokio::test]
827 async fn host_guard_allows_known_host_and_missing_host() {
828 let req = json_post("/graphql")
829 .header(axum::http::header::HOST, "127.0.0.1:4000")
830 .body(Body::empty())
831 .unwrap();
832 assert_eq!(run_host(req).await, StatusCode::OK);
833
834 let req = json_post("/graphql").body(Body::empty()).unwrap();
835 assert_eq!(run_host(req).await, StatusCode::OK);
836 }
837
838 #[tokio::test]
841 async fn ws_upgrade_from_foreign_origin_is_rejected() {
842 let req = HttpRequest::get("/graphql/ws")
843 .header(axum::http::header::HOST, "127.0.0.1:4000")
844 .header(axum::http::header::ORIGIN, "https://evil.com")
845 .body(Body::empty())
846 .unwrap();
847 assert_eq!(run_browser(req).await, StatusCode::FORBIDDEN);
848 }
849
850 #[tokio::test]
851 async fn ws_upgrade_without_origin_is_allowed() {
852 let req = HttpRequest::get("/graphql/ws")
853 .header(axum::http::header::HOST, "127.0.0.1:4000")
854 .body(Body::empty())
855 .unwrap();
856 assert_eq!(run_browser(req).await, StatusCode::OK);
857 }
858
859 #[tokio::test]
860 async fn configured_and_same_origin_are_allowed() {
861 let req = HttpRequest::get("/graphql/ws")
862 .header(axum::http::header::HOST, "127.0.0.1:4000")
863 .header(axum::http::header::ORIGIN, "https://music.example.com")
864 .body(Body::empty())
865 .unwrap();
866 assert_eq!(run_browser(req).await, StatusCode::OK);
867
868 let req = json_post("/graphql")
870 .header(axum::http::header::HOST, "127.0.0.1:4000")
871 .header(axum::http::header::ORIGIN, "http://127.0.0.1:4000")
872 .body(Body::empty())
873 .unwrap();
874 assert_eq!(run_browser(req).await, StatusCode::OK);
875 }
876
877 #[tokio::test]
880 async fn text_plain_post_is_rejected() {
881 let req = HttpRequest::post("/graphql")
882 .header(axum::http::header::CONTENT_TYPE, "text/plain")
883 .body(Body::from(r#"{"query":"mutation{clearQueue{ok}}"}"#))
884 .unwrap();
885 assert_eq!(run_browser(req).await, StatusCode::UNSUPPORTED_MEDIA_TYPE);
886 }
887
888 #[tokio::test]
889 async fn post_without_content_type_is_rejected() {
890 let req = HttpRequest::post("/graphql").body(Body::empty()).unwrap();
891 assert_eq!(run_browser(req).await, StatusCode::UNSUPPORTED_MEDIA_TYPE);
892 }
893
894 #[tokio::test]
897 async fn load_perimeter_refuses_an_oversized_query_body() {
898 async fn parse(_: async_graphql_axum::GraphQLRequest) -> StatusCode {
899 StatusCode::OK
900 }
901 let app = load_perimeter(axum::Router::new().route("/graphql", post(parse)));
902 let body = |padding: usize| {
905 let chunks = [
906 axum::body::Bytes::from_static(br#"{"query":"{__typename}""#),
907 axum::body::Bytes::from(vec![b' '; padding]),
908 axum::body::Bytes::from_static(b"}"),
909 ];
910 Body::from_stream(tokio_stream::iter(chunks.map(Ok::<_, std::io::Error>)))
911 };
912 let req = json_post("/graphql").body(body(1 << 10)).unwrap();
913 assert_eq!(
914 app.clone().oneshot(req).await.unwrap().status(),
915 StatusCode::OK
916 );
917 let req = json_post("/graphql").body(body(3 << 20)).unwrap();
918 assert_ne!(app.oneshot(req).await.unwrap().status(), StatusCode::OK);
919 }
920
921 #[tokio::test]
922 async fn load_perimeter_passes_requests_and_turns_panics_into_500s() {
923 async fn boom() -> &'static str {
924 panic!("resolver exploded");
925 }
926
927 let app = load_perimeter(
928 axum::Router::new()
929 .route("/graphql", post(ok))
930 .route("/boom", post(boom)),
931 );
932
933 let req = json_post("/graphql").body(Body::empty()).unwrap();
934 assert_eq!(
935 app.clone().oneshot(req).await.unwrap().status(),
936 StatusCode::OK
937 );
938
939 let req = json_post("/boom").body(Body::empty()).unwrap();
941 assert_eq!(
942 app.oneshot(req).await.unwrap().status(),
943 StatusCode::INTERNAL_SERVER_ERROR
944 );
945 }
946
947 #[tokio::test]
948 async fn json_post_is_accepted() {
949 let req = json_post("/graphql").body(Body::empty()).unwrap();
950 assert_eq!(run_browser(req).await, StatusCode::OK);
951
952 let req = HttpRequest::post("/graphql")
953 .header(
954 axum::http::header::CONTENT_TYPE,
955 "application/json; charset=utf-8",
956 )
957 .body(Body::empty())
958 .unwrap();
959 assert_eq!(run_browser(req).await, StatusCode::OK);
960 }
961}