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
89fn load_perimeter<S>(router: axum::Router<S>) -> axum::Router<S>
94where
95 S: Clone + Send + Sync + 'static,
96{
97 router
98 .layer(tower_http::catch_panic::CatchPanicLayer::new())
101 .layer(tower_http::timeout::TimeoutLayer::with_status_code(
102 axum::http::StatusCode::REQUEST_TIMEOUT,
103 REQUEST_TIMEOUT,
104 ))
105 .layer(
109 tower::ServiceBuilder::new()
110 .layer(axum::error_handling::HandleErrorLayer::new(
111 |err: tower::BoxError| async move {
112 if err.is::<tower::load_shed::error::Overloaded>() {
113 (
114 axum::http::StatusCode::SERVICE_UNAVAILABLE,
115 "server at capacity",
116 )
117 } else {
118 (
119 axum::http::StatusCode::INTERNAL_SERVER_ERROR,
120 "internal error",
121 )
122 }
123 },
124 ))
125 .load_shed()
126 .concurrency_limit(MAX_INFLIGHT_QUERIES),
127 )
128}
129
130pub struct ApiServerOpts {
132 pub state: Arc<SharedPlayerState>,
133 pub cmd_tx: Sender<PlayerCommand>,
134 pub pool: Arc<Pool>,
137 pub port: Option<u16>,
138 pub bind: Option<std::net::IpAddr>,
139 pub subsonic_port: Option<u16>,
140 pub playground: bool,
141 pub viz: Option<Arc<VizSnapshot>>,
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 } = opts;
160 use axum::routing::{get, post};
161
162 let cfg = Config::load().unwrap_or_default();
163 let port = port.unwrap_or(cfg.graphql.port);
164 let bind = bind.unwrap_or(cfg.graphql.bind);
165 let subsonic_port = subsonic_port.or(cfg.subsonic.port);
166 let playground_enabled = playground || cfg.graphql.playground;
167 let auth_enabled = cfg.graphql.auth_enabled;
168
169 let (private_pem, public_pem) = if auth_enabled {
175 let kp = auth::load_or_generate_keypair().map_err(|e| {
176 format!(
177 "auth_enabled = true but the keypair could not be loaded or created: {}",
178 e
179 )
180 })?;
181 if kp.0.is_empty() || kp.1.is_empty() {
184 return Err("auth_enabled = true but the keypair files are empty. \
185 Run `koan auth regenerate-keys`."
186 .into());
187 }
188 kp
189 } else {
190 auth::load_or_generate_keypair().unwrap_or_default()
193 };
194
195 let access_ttl = parse_duration_secs(&cfg.graphql.access_token_ttl).unwrap_or(900);
196 let refresh_ttl = parse_duration_secs(&cfg.graphql.refresh_token_ttl).unwrap_or(2_592_000);
197
198 let introspection_key = if playground_enabled && auth_enabled {
202 Some(Arc::new(auth::random_token().map_err(|e| {
203 format!("failed to generate introspection key: {}", e)
204 })?))
205 } else {
206 None
207 };
208
209 let auth_state = AuthState {
210 public_pem: Arc::new(public_pem.clone()),
211 auth_enabled,
212 introspection_key: introspection_key.clone(),
213 };
214
215 let auth_route_state = AuthRouteState {
216 pool: pool.clone(),
217 private_pem: Arc::new(private_pem),
218 public_pem: Arc::new(public_pem),
219 access_ttl_secs: access_ttl,
220 refresh_ttl_secs: refresh_ttl,
221 cookie_secure: cfg.graphql.cookie_secure,
222 login_limiter: Arc::new(LoginRateLimiter::default()),
223 };
224
225 let schema = build_schema(state, cmd_tx, pool.path().to_path_buf(), viz);
226
227 if auth_enabled {
228 log::info!(
229 "Auth enabled (Ed25519 JWT, access TTL {}s, refresh TTL {}s)",
230 access_ttl,
231 refresh_ttl
232 );
233 } else {
234 log::info!("Auth disabled — all requests treated as admin");
235 }
236
237 let browser_policy = Arc::new(BrowserPolicy {
238 origins: cfg.graphql.cors_origins.clone(),
239 hosts: cfg.graphql.allowed_hosts.clone(),
240 });
241
242 if cfg.graphql.cors_origins.is_empty() {
243 log::info!("CORS: no origins configured — browsers get no cross-origin access");
244 }
245
246 let rt = tokio::runtime::Runtime::new().expect("failed to create tokio runtime");
247 rt.block_on(async {
248 let query_route = load_perimeter(axum::Router::new().route("/graphql", post(graphql_handler)));
253
254 let gql_app = axum::Router::new()
255 .merge(query_route)
256 .route("/graphql/ws", get(graphql_ws_handler))
257 .layer(axum::middleware::from_fn_with_state(
258 auth_state.clone(),
259 auth_middleware,
260 ))
261 .layer(axum::middleware::from_fn_with_state(
264 browser_policy.clone(),
265 browser_guard,
266 ))
267 .with_state(schema);
268
269 let covers = Arc::new(crate::covers::Covers::in_config_dir());
273 let ui_routes = crate::ui::router(
274 pool.clone(),
275 auth_route_state.clone(),
276 auth_enabled,
277 covers.clone(),
278 cfg.sharing.public_url.clone(),
279 );
280
281 let auth_app = auth_router(auth_route_state);
283
284 let origins: Vec<axum::http::HeaderValue> = cfg
288 .graphql
289 .cors_origins
290 .iter()
291 .filter_map(|o| o.parse().ok())
292 .collect();
293 let cors = tower_http::cors::CorsLayer::new()
294 .allow_origin(origins)
295 .allow_methods([
296 axum::http::Method::GET,
297 axum::http::Method::POST,
298 axum::http::Method::OPTIONS,
299 ])
300 .allow_headers([
301 axum::http::header::AUTHORIZATION,
302 axum::http::header::CONTENT_TYPE,
303 axum::http::HeaderName::from_static("x-introspection-key"),
304 ])
305 .allow_credentials(true);
306
307 let share_routes = crate::share::router(
317 pool.clone(),
318 cfg.sharing.public_url.clone(),
319 covers.clone(),
320 )
321 .merge(crate::push::router(pool.clone(), covers));
322 let subsonic_merged = crate::subsonic::subsonic_router(pool);
323 let subsonic_on_main = subsonic_merged.is_some();
324 let subsonic_dedicated = subsonic_merged.clone();
325
326 let mut app = auth_app.merge(gql_app).merge(share_routes).merge(ui_routes);
327 if let Some(sub) = subsonic_merged {
328 app = app.merge(sub);
329 }
330 if playground_enabled {
331 app = app.route(
332 "/graphql",
333 get(graphql_playground).with_state(introspection_key.clone()),
334 );
335 }
336 let app = app.layer(cors).layer(axum::middleware::from_fn_with_state(
340 browser_policy.clone(),
341 host_guard,
342 ));
343
344 let playground_url = if playground_enabled {
346 if let Some(ref key) = introspection_key {
347 format!("http://{}:{}/graphql?introspection-key={}", bind, port, key)
348 } else {
349 format!("http://{}:{}/graphql", bind, port)
350 }
351 } else {
352 format!("http://{}:{}/graphql", bind, port)
353 };
354
355 let gql_addr = std::net::SocketAddr::new(bind, port);
356
357 let gql_listener = match tokio::net::TcpListener::bind(gql_addr).await {
358 Ok(l) => {
359 log::info!("GraphQL API on http://{}:{}/graphql", bind, port);
360 if subsonic_on_main {
361 log::info!("Subsonic REST on http://{}:{}/rest/", bind, port);
362 }
363 if playground_enabled {
364 log::info!("GraphiQL: {}", playground_url);
365 #[cfg(target_os = "macos")]
367 let _ = std::process::Command::new("open").arg(&playground_url).spawn();
368 #[cfg(target_os = "linux")]
369 let _ = std::process::Command::new("xdg-open").arg(&playground_url).spawn();
370 }
371 l
372 }
373 Err(e) => {
374 return Err(format!(
375 "failed to bind GraphQL port {port} — {e} (another instance running?)"
376 ));
377 }
378 };
379 let gql_server = axum::serve(
380 gql_listener,
381 app.into_make_service_with_connect_info::<std::net::SocketAddr>(),
382 )
383 .with_graceful_shutdown(shutdown_signal());
384
385 let extra_sub_port = subsonic_port.filter(|p| *p != port);
389 if let Some(sub_port) = extra_sub_port
390 && let Some(sub_app) = subsonic_dedicated
391 {
392 let sub_addr = std::net::SocketAddr::new(bind, sub_port);
393 match tokio::net::TcpListener::bind(sub_addr).await {
394 Ok(sub_listener) => {
395 log::info!(
396 "Subsonic REST also on http://{}:{}/rest/ (dedicated port)",
397 bind,
398 sub_port,
399 );
400 let sub_server = axum::serve(sub_listener, sub_app)
401 .with_graceful_shutdown(shutdown_signal());
402
403 tokio::select! {
404 r = gql_server => { if let Err(e) = r { log::error!("GraphQL server error: {e}"); } },
405 r = sub_server => { if let Err(e) = r { log::error!("Subsonic server error: {e}"); } },
406 }
407 return Ok(());
408 }
409 Err(e) => {
410 log::warn!(
411 "Dedicated Subsonic port {} unavailable — {}. Mounted on GraphQL port only.",
412 sub_port,
413 e,
414 );
415 }
416 }
417 }
418
419 if let Err(e) = gql_server.await {
420 log::error!("GraphQL server error: {e}");
421 }
422 Ok(())
423 })
424}
425
426pub fn start_api_background(
432 state: Arc<SharedPlayerState>,
433 cmd_tx: Sender<PlayerCommand>,
434 db_path: PathBuf,
435 port: Option<u16>,
436 bind: Option<std::net::IpAddr>,
437 subsonic_port: Option<u16>,
438 playground: bool,
439) {
440 if let Err(e) = run_api_blocking(ApiServerOpts {
443 state,
444 cmd_tx,
445 pool: Arc::new(Pool::new(db_path)),
446 port,
447 bind,
448 subsonic_port,
449 playground,
450 viz: None,
451 }) {
452 log::error!("API server not started: {}", e);
453 }
454}
455
456pub(crate) struct BrowserPolicy {
465 origins: Vec<String>,
466 hosts: Vec<String>,
467}
468
469impl BrowserPolicy {
470 fn host_allowed(&self, host: &str) -> bool {
471 if self.hosts.iter().any(|h| h.eq_ignore_ascii_case(host)) {
472 return true;
473 }
474 let bare = strip_port(host);
475 if self.hosts.iter().any(|h| h.eq_ignore_ascii_case(bare)) {
476 return true;
477 }
478 bare.eq_ignore_ascii_case("localhost") || bare.parse::<std::net::IpAddr>().is_ok()
481 }
482
483 fn origin_allowed(&self, origin: &str, host: Option<&str>) -> bool {
486 if self.origins.iter().any(|o| o == origin) {
487 return true;
488 }
489 match (origin.split_once("://"), host) {
490 (Some((_, authority)), Some(host)) => authority.eq_ignore_ascii_case(host),
491 _ => false,
492 }
493 }
494}
495
496fn strip_port(host: &str) -> &str {
498 if let Some(rest) = host.strip_prefix('[') {
499 return rest.split(']').next().unwrap_or(rest);
500 }
501 match host.rsplit_once(':') {
502 Some((h, port)) if !port.is_empty() && port.bytes().all(|b| b.is_ascii_digit()) => h,
503 _ => host,
504 }
505}
506
507fn header_str(request: &axum::extract::Request, name: axum::http::HeaderName) -> Option<&str> {
508 request.headers().get(name).and_then(|v| v.to_str().ok())
509}
510
511async fn host_guard(
513 axum::extract::State(policy): axum::extract::State<Arc<BrowserPolicy>>,
514 request: axum::extract::Request,
515 next: axum::middleware::Next,
516) -> axum::response::Response {
517 use axum::response::IntoResponse;
518
519 let host = header_str(&request, axum::http::header::HOST)
522 .map(str::to_owned)
523 .or_else(|| request.uri().host().map(str::to_owned));
524
525 if let Some(ref host) = host
526 && !policy.host_allowed(host)
527 {
528 log::warn!("rejected request for unrecognised Host: {}", host);
529 return (axum::http::StatusCode::FORBIDDEN, "host not allowed").into_response();
530 }
531
532 next.run(request).await
533}
534
535async fn browser_guard(
544 axum::extract::State(policy): axum::extract::State<Arc<BrowserPolicy>>,
545 request: axum::extract::Request,
546 next: axum::middleware::Next,
547) -> axum::response::Response {
548 use axum::response::IntoResponse;
549
550 let host = header_str(&request, axum::http::header::HOST).map(str::to_owned);
551 if let Some(origin) = header_str(&request, axum::http::header::ORIGIN)
553 && !policy.origin_allowed(origin, host.as_deref())
554 {
555 log::warn!(
556 "rejected GraphQL request from disallowed Origin: {}",
557 origin
558 );
559 return (axum::http::StatusCode::FORBIDDEN, "origin not allowed").into_response();
560 }
561
562 if request.method() == axum::http::Method::POST && !is_graphql_content_type(&request) {
563 return (
564 axum::http::StatusCode::UNSUPPORTED_MEDIA_TYPE,
565 "content type must be application/json or application/graphql",
566 )
567 .into_response();
568 }
569
570 next.run(request).await
571}
572
573fn is_graphql_content_type(request: &axum::extract::Request) -> bool {
574 header_str(request, axum::http::header::CONTENT_TYPE).is_some_and(|ct| {
575 let ct = ct.trim().to_ascii_lowercase();
576 ct.starts_with("application/json") || ct.starts_with("application/graphql")
577 })
578}
579
580async fn shutdown_signal() {
581 tokio::signal::ctrl_c()
582 .await
583 .expect("failed to listen for ctrl+c");
584}
585
586async fn graphql_handler(
587 axum::Extension(user): axum::Extension<AuthUser>,
588 axum::extract::State(schema): axum::extract::State<KoanSchema>,
589 headers: axum::http::HeaderMap,
590 req: async_graphql_axum::GraphQLRequest,
591) -> async_graphql_axum::GraphQLResponse {
592 let mut request = req.into_inner();
593 if let Some(origin) = crate::origin::origin(&headers, None) {
594 request = request.data(super::RequestOrigin(origin));
595 }
596 request = request.data(user);
599 schema.execute(request).await.into()
600}
601
602async fn graphql_ws_handler(
603 axum::Extension(user): axum::Extension<AuthUser>,
604 axum::extract::State(schema): axum::extract::State<KoanSchema>,
605 protocol: async_graphql_axum::GraphQLProtocol,
606 websocket: axum::extract::WebSocketUpgrade,
607) -> axum::response::Response {
608 websocket
609 .protocols(async_graphql::http::ALL_WEBSOCKET_PROTOCOLS)
610 .on_upgrade(move |stream| {
611 let stream = async_graphql_axum::GraphQLWebSocket::new(stream, schema, protocol)
612 .on_connection_init(move |_| async move {
613 let mut data = async_graphql::Data::default();
614 data.insert(user);
615 Ok(data)
616 });
617 async move {
618 stream.serve().await;
619 }
620 })
621}
622
623async fn graphql_playground(
624 axum::extract::Query(params): axum::extract::Query<std::collections::HashMap<String, String>>,
625 axum::extract::State(key): axum::extract::State<Option<Arc<String>>>,
626) -> axum::response::Response {
627 use axum::response::IntoResponse;
628
629 if let Some(ref expected) = key {
631 let provided = params.get("introspection-key");
632 if provided.map(|k| k.as_str()) != Some(expected.as_str()) {
633 return (
634 axum::http::StatusCode::FORBIDDEN,
635 "invalid or missing introspection-key",
636 )
637 .into_response();
638 }
639 }
640
641 let mut source = async_graphql::http::GraphiQLSource::build().endpoint("/graphql");
644 if let Some(ref k) = key {
645 source = source.header("X-Introspection-Key", k.as_str());
646 }
647
648 axum::response::Html(source.finish()).into_response()
649}
650
651pub fn cmd_serve_daemon(
653 port: Option<u16>,
654 bind: Option<std::net::IpAddr>,
655 subsonic_port: Option<u16>,
656 playground: bool,
657 mcp_bind: Option<std::net::SocketAddr>,
658) {
659 use std::fs;
660 use std::process::Command;
661
662 let cfg = Config::load().unwrap_or_default();
663 let port_val = port.unwrap_or(cfg.graphql.port);
664 let bind_val = bind.unwrap_or(cfg.graphql.bind);
665
666 let exe = std::env::current_exe().expect("failed to get current exe path");
667 let mut cmd = Command::new(exe);
668 cmd.arg("--headless");
670 cmd.arg("--port").arg(port_val.to_string());
671 cmd.arg("--bind").arg(bind_val.to_string());
672 if let Some(sp) = subsonic_port {
673 cmd.arg("--subsonic").arg(sp.to_string());
674 }
675 if let Some(addr) = mcp_bind {
676 cmd.arg("--mcp-bind").arg(addr.to_string());
677 }
678 if playground || cfg.graphql.playground {
679 cmd.arg("--playground");
680 }
681
682 cmd.stdin(std::process::Stdio::null());
683 cmd.stdout(std::process::Stdio::null());
684 cmd.stderr(std::process::Stdio::null());
685
686 let mut child = cmd.spawn().expect("failed to spawn daemon process");
687 let pid = child.id();
688
689 let pid_path = koan_core::config::config_dir().join("koan-serve.pid");
690 fs::write(&pid_path, pid.to_string()).ok();
691
692 std::thread::spawn(move || {
693 let _ = child.wait();
694 });
695
696 eprintln!("kōan daemon started (pid {}) on port {}", pid, port_val);
697 if let Some(sp) = subsonic_port {
698 eprintln!(" Subsonic REST on port {}", sp);
699 }
700 eprintln!(" PID file: {}", pid_path.display());
701}
702
703pub async fn execute_in_process(
712 schema: &KoanSchema,
713 query: &str,
714 variables: Option<serde_json::Value>,
715 user_id: i64,
716 role: koan_core::auth::Role,
717) -> serde_json::Value {
718 let mut request = async_graphql::Request::new(query);
719 request = request.data(AuthUser {
720 user_id,
721 role,
722 ..AuthUser::anonymous_admin()
723 });
724 if let Some(serde_json::Value::Object(map)) = variables {
725 let mut gql_vars = async_graphql::Variables::default();
726 for (k, v) in map {
727 gql_vars.insert(
728 async_graphql::Name::new(&k),
729 async_graphql::Value::from_json(v).unwrap_or(async_graphql::Value::Null),
730 );
731 }
732 request = request.variables(gql_vars);
733 }
734 let response = schema.execute(request).await;
735 serde_json::to_value(&response).unwrap_or(serde_json::Value::Null)
736}
737
738#[cfg(test)]
743mod tests {
744 use super::*;
745 use axum::body::Body;
746 use axum::http::{Request as HttpRequest, StatusCode};
747 use axum::routing::{get, post};
748 use tower::ServiceExt as _;
749
750 fn policy() -> Arc<BrowserPolicy> {
751 Arc::new(BrowserPolicy {
752 origins: vec!["https://music.example.com".into()],
753 hosts: vec!["koan.local".into()],
754 })
755 }
756
757 async fn ok() -> &'static str {
758 "ok"
759 }
760
761 fn routes() -> axum::Router<Arc<BrowserPolicy>> {
762 axum::Router::new()
763 .route("/graphql", post(ok).get(ok))
764 .route("/graphql/ws", get(ok))
765 }
766
767 async fn run_host(req: HttpRequest<Body>) -> StatusCode {
768 let app = routes()
769 .layer(axum::middleware::from_fn_with_state(policy(), host_guard))
770 .with_state(policy());
771 app.oneshot(req).await.unwrap().status()
772 }
773
774 async fn run_browser(req: HttpRequest<Body>) -> StatusCode {
775 let app = routes()
776 .layer(axum::middleware::from_fn_with_state(
777 policy(),
778 browser_guard,
779 ))
780 .with_state(policy());
781 app.oneshot(req).await.unwrap().status()
782 }
783
784 fn json_post(uri: &str) -> axum::http::request::Builder {
785 HttpRequest::post(uri).header(axum::http::header::CONTENT_TYPE, "application/json")
786 }
787
788 #[test]
791 fn host_policy_accepts_loopback_literals_and_configured_names() {
792 let p = policy();
793 assert!(p.host_allowed("localhost:4000"));
794 assert!(p.host_allowed("127.0.0.1:4000"));
795 assert!(p.host_allowed("192.168.1.20:4000"));
796 assert!(p.host_allowed("[::1]:4000"));
797 assert!(p.host_allowed("koan.local"));
798 assert!(p.host_allowed("koan.local:4000"));
799 }
800
801 #[test]
802 fn host_policy_rejects_attacker_controlled_names() {
803 let p = policy();
804 assert!(!p.host_allowed("evil.com"));
805 assert!(!p.host_allowed("rebind.evil.com:4000"));
806 assert!(!p.host_allowed("koan.local.evil.com"));
807 }
808
809 #[tokio::test]
810 async fn host_guard_rejects_foreign_host() {
811 let req = json_post("/graphql")
812 .header(axum::http::header::HOST, "rebind.evil.com")
813 .body(Body::empty())
814 .unwrap();
815 assert_eq!(run_host(req).await, StatusCode::FORBIDDEN);
816 }
817
818 #[tokio::test]
819 async fn host_guard_allows_known_host_and_missing_host() {
820 let req = json_post("/graphql")
821 .header(axum::http::header::HOST, "127.0.0.1:4000")
822 .body(Body::empty())
823 .unwrap();
824 assert_eq!(run_host(req).await, StatusCode::OK);
825
826 let req = json_post("/graphql").body(Body::empty()).unwrap();
827 assert_eq!(run_host(req).await, StatusCode::OK);
828 }
829
830 #[tokio::test]
833 async fn ws_upgrade_from_foreign_origin_is_rejected() {
834 let req = HttpRequest::get("/graphql/ws")
835 .header(axum::http::header::HOST, "127.0.0.1:4000")
836 .header(axum::http::header::ORIGIN, "https://evil.com")
837 .body(Body::empty())
838 .unwrap();
839 assert_eq!(run_browser(req).await, StatusCode::FORBIDDEN);
840 }
841
842 #[tokio::test]
843 async fn ws_upgrade_without_origin_is_allowed() {
844 let req = HttpRequest::get("/graphql/ws")
845 .header(axum::http::header::HOST, "127.0.0.1:4000")
846 .body(Body::empty())
847 .unwrap();
848 assert_eq!(run_browser(req).await, StatusCode::OK);
849 }
850
851 #[tokio::test]
852 async fn configured_and_same_origin_are_allowed() {
853 let req = HttpRequest::get("/graphql/ws")
854 .header(axum::http::header::HOST, "127.0.0.1:4000")
855 .header(axum::http::header::ORIGIN, "https://music.example.com")
856 .body(Body::empty())
857 .unwrap();
858 assert_eq!(run_browser(req).await, StatusCode::OK);
859
860 let req = json_post("/graphql")
862 .header(axum::http::header::HOST, "127.0.0.1:4000")
863 .header(axum::http::header::ORIGIN, "http://127.0.0.1:4000")
864 .body(Body::empty())
865 .unwrap();
866 assert_eq!(run_browser(req).await, StatusCode::OK);
867 }
868
869 #[tokio::test]
872 async fn text_plain_post_is_rejected() {
873 let req = HttpRequest::post("/graphql")
874 .header(axum::http::header::CONTENT_TYPE, "text/plain")
875 .body(Body::from(r#"{"query":"mutation{clearQueue{ok}}"}"#))
876 .unwrap();
877 assert_eq!(run_browser(req).await, StatusCode::UNSUPPORTED_MEDIA_TYPE);
878 }
879
880 #[tokio::test]
881 async fn post_without_content_type_is_rejected() {
882 let req = HttpRequest::post("/graphql").body(Body::empty()).unwrap();
883 assert_eq!(run_browser(req).await, StatusCode::UNSUPPORTED_MEDIA_TYPE);
884 }
885
886 #[tokio::test]
889 async fn load_perimeter_passes_requests_and_turns_panics_into_500s() {
890 async fn boom() -> &'static str {
891 panic!("resolver exploded");
892 }
893
894 let app = load_perimeter(
895 axum::Router::new()
896 .route("/graphql", post(ok))
897 .route("/boom", post(boom)),
898 );
899
900 let req = json_post("/graphql").body(Body::empty()).unwrap();
901 assert_eq!(
902 app.clone().oneshot(req).await.unwrap().status(),
903 StatusCode::OK
904 );
905
906 let req = json_post("/boom").body(Body::empty()).unwrap();
908 assert_eq!(
909 app.oneshot(req).await.unwrap().status(),
910 StatusCode::INTERNAL_SERVER_ERROR
911 );
912 }
913
914 #[tokio::test]
915 async fn json_post_is_accepted() {
916 let req = json_post("/graphql").body(Body::empty()).unwrap();
917 assert_eq!(run_browser(req).await, StatusCode::OK);
918
919 let req = HttpRequest::post("/graphql")
920 .header(
921 axum::http::header::CONTENT_TYPE,
922 "application/json; charset=utf-8",
923 )
924 .body(Body::empty())
925 .unwrap();
926 assert_eq!(run_browser(req).await, StatusCode::OK);
927 }
928}