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 );
279
280 let auth_app = auth_router(auth_route_state);
282
283 let origins: Vec<axum::http::HeaderValue> = cfg
287 .graphql
288 .cors_origins
289 .iter()
290 .filter_map(|o| o.parse().ok())
291 .collect();
292 let cors = tower_http::cors::CorsLayer::new()
293 .allow_origin(origins)
294 .allow_methods([
295 axum::http::Method::GET,
296 axum::http::Method::POST,
297 axum::http::Method::OPTIONS,
298 ])
299 .allow_headers([
300 axum::http::header::AUTHORIZATION,
301 axum::http::header::CONTENT_TYPE,
302 axum::http::HeaderName::from_static("x-introspection-key"),
303 ])
304 .allow_credentials(true);
305
306 let share_routes = crate::share::router(
316 pool.clone(),
317 cfg.sharing.public_url.clone(),
318 covers.clone(),
319 )
320 .merge(crate::push::router(pool.clone(), covers));
321 let subsonic_merged = crate::subsonic::subsonic_router(pool);
322 let subsonic_on_main = subsonic_merged.is_some();
323 let subsonic_dedicated = subsonic_merged.clone();
324
325 let mut app = auth_app.merge(gql_app).merge(share_routes).merge(ui_routes);
326 if let Some(sub) = subsonic_merged {
327 app = app.merge(sub);
328 }
329 if playground_enabled {
330 app = app.route(
331 "/graphql",
332 get(graphql_playground).with_state(introspection_key.clone()),
333 );
334 }
335 let app = app.layer(cors).layer(axum::middleware::from_fn_with_state(
339 browser_policy.clone(),
340 host_guard,
341 ));
342
343 let playground_url = if playground_enabled {
345 if let Some(ref key) = introspection_key {
346 format!("http://{}:{}/graphql?introspection-key={}", bind, port, key)
347 } else {
348 format!("http://{}:{}/graphql", bind, port)
349 }
350 } else {
351 format!("http://{}:{}/graphql", bind, port)
352 };
353
354 let gql_addr = std::net::SocketAddr::new(bind, port);
355
356 let gql_listener = match tokio::net::TcpListener::bind(gql_addr).await {
357 Ok(l) => {
358 log::info!("GraphQL API on http://{}:{}/graphql", bind, port);
359 if subsonic_on_main {
360 log::info!("Subsonic REST on http://{}:{}/rest/", bind, port);
361 }
362 if playground_enabled {
363 log::info!("GraphiQL: {}", playground_url);
364 #[cfg(target_os = "macos")]
366 let _ = std::process::Command::new("open").arg(&playground_url).spawn();
367 #[cfg(target_os = "linux")]
368 let _ = std::process::Command::new("xdg-open").arg(&playground_url).spawn();
369 }
370 l
371 }
372 Err(e) => {
373 return Err(format!(
374 "failed to bind GraphQL port {port} — {e} (another instance running?)"
375 ));
376 }
377 };
378 let gql_server = axum::serve(
379 gql_listener,
380 app.into_make_service_with_connect_info::<std::net::SocketAddr>(),
381 )
382 .with_graceful_shutdown(shutdown_signal());
383
384 let extra_sub_port = subsonic_port.filter(|p| *p != port);
388 if let Some(sub_port) = extra_sub_port
389 && let Some(sub_app) = subsonic_dedicated
390 {
391 let sub_addr = std::net::SocketAddr::new(bind, sub_port);
392 match tokio::net::TcpListener::bind(sub_addr).await {
393 Ok(sub_listener) => {
394 log::info!(
395 "Subsonic REST also on http://{}:{}/rest/ (dedicated port)",
396 bind,
397 sub_port,
398 );
399 let sub_server = axum::serve(sub_listener, sub_app)
400 .with_graceful_shutdown(shutdown_signal());
401
402 tokio::select! {
403 r = gql_server => { if let Err(e) = r { log::error!("GraphQL server error: {e}"); } },
404 r = sub_server => { if let Err(e) = r { log::error!("Subsonic server error: {e}"); } },
405 }
406 return Ok(());
407 }
408 Err(e) => {
409 log::warn!(
410 "Dedicated Subsonic port {} unavailable — {}. Mounted on GraphQL port only.",
411 sub_port,
412 e,
413 );
414 }
415 }
416 }
417
418 if let Err(e) = gql_server.await {
419 log::error!("GraphQL server error: {e}");
420 }
421 Ok(())
422 })
423}
424
425pub fn start_api_background(
431 state: Arc<SharedPlayerState>,
432 cmd_tx: Sender<PlayerCommand>,
433 db_path: PathBuf,
434 port: Option<u16>,
435 bind: Option<std::net::IpAddr>,
436 subsonic_port: Option<u16>,
437 playground: bool,
438) {
439 if let Err(e) = run_api_blocking(ApiServerOpts {
442 state,
443 cmd_tx,
444 pool: Arc::new(Pool::new(db_path)),
445 port,
446 bind,
447 subsonic_port,
448 playground,
449 viz: None,
450 }) {
451 log::error!("API server not started: {}", e);
452 }
453}
454
455pub(crate) struct BrowserPolicy {
464 origins: Vec<String>,
465 hosts: Vec<String>,
466}
467
468impl BrowserPolicy {
469 fn host_allowed(&self, host: &str) -> bool {
470 if self.hosts.iter().any(|h| h.eq_ignore_ascii_case(host)) {
471 return true;
472 }
473 let bare = strip_port(host);
474 if self.hosts.iter().any(|h| h.eq_ignore_ascii_case(bare)) {
475 return true;
476 }
477 bare.eq_ignore_ascii_case("localhost") || bare.parse::<std::net::IpAddr>().is_ok()
480 }
481
482 fn origin_allowed(&self, origin: &str, host: Option<&str>) -> bool {
485 if self.origins.iter().any(|o| o == origin) {
486 return true;
487 }
488 match (origin.split_once("://"), host) {
489 (Some((_, authority)), Some(host)) => authority.eq_ignore_ascii_case(host),
490 _ => false,
491 }
492 }
493}
494
495fn strip_port(host: &str) -> &str {
497 if let Some(rest) = host.strip_prefix('[') {
498 return rest.split(']').next().unwrap_or(rest);
499 }
500 match host.rsplit_once(':') {
501 Some((h, port)) if !port.is_empty() && port.bytes().all(|b| b.is_ascii_digit()) => h,
502 _ => host,
503 }
504}
505
506fn header_str(request: &axum::extract::Request, name: axum::http::HeaderName) -> Option<&str> {
507 request.headers().get(name).and_then(|v| v.to_str().ok())
508}
509
510async fn host_guard(
512 axum::extract::State(policy): axum::extract::State<Arc<BrowserPolicy>>,
513 request: axum::extract::Request,
514 next: axum::middleware::Next,
515) -> axum::response::Response {
516 use axum::response::IntoResponse;
517
518 let host = header_str(&request, axum::http::header::HOST)
521 .map(str::to_owned)
522 .or_else(|| request.uri().host().map(str::to_owned));
523
524 if let Some(ref host) = host
525 && !policy.host_allowed(host)
526 {
527 log::warn!("rejected request for unrecognised Host: {}", host);
528 return (axum::http::StatusCode::FORBIDDEN, "host not allowed").into_response();
529 }
530
531 next.run(request).await
532}
533
534async fn browser_guard(
543 axum::extract::State(policy): axum::extract::State<Arc<BrowserPolicy>>,
544 request: axum::extract::Request,
545 next: axum::middleware::Next,
546) -> axum::response::Response {
547 use axum::response::IntoResponse;
548
549 let host = header_str(&request, axum::http::header::HOST).map(str::to_owned);
550 if let Some(origin) = header_str(&request, axum::http::header::ORIGIN)
552 && !policy.origin_allowed(origin, host.as_deref())
553 {
554 log::warn!(
555 "rejected GraphQL request from disallowed Origin: {}",
556 origin
557 );
558 return (axum::http::StatusCode::FORBIDDEN, "origin not allowed").into_response();
559 }
560
561 if request.method() == axum::http::Method::POST && !is_graphql_content_type(&request) {
562 return (
563 axum::http::StatusCode::UNSUPPORTED_MEDIA_TYPE,
564 "content type must be application/json or application/graphql",
565 )
566 .into_response();
567 }
568
569 next.run(request).await
570}
571
572fn is_graphql_content_type(request: &axum::extract::Request) -> bool {
573 header_str(request, axum::http::header::CONTENT_TYPE).is_some_and(|ct| {
574 let ct = ct.trim().to_ascii_lowercase();
575 ct.starts_with("application/json") || ct.starts_with("application/graphql")
576 })
577}
578
579async fn shutdown_signal() {
580 tokio::signal::ctrl_c()
581 .await
582 .expect("failed to listen for ctrl+c");
583}
584
585async fn graphql_handler(
586 axum::Extension(user): axum::Extension<AuthUser>,
587 axum::extract::State(schema): axum::extract::State<KoanSchema>,
588 req: async_graphql_axum::GraphQLRequest,
589) -> async_graphql_axum::GraphQLResponse {
590 let mut request = req.into_inner();
591 request = request.data(user);
594 schema.execute(request).await.into()
595}
596
597async fn graphql_ws_handler(
598 axum::Extension(user): axum::Extension<AuthUser>,
599 axum::extract::State(schema): axum::extract::State<KoanSchema>,
600 protocol: async_graphql_axum::GraphQLProtocol,
601 websocket: axum::extract::WebSocketUpgrade,
602) -> axum::response::Response {
603 websocket
604 .protocols(async_graphql::http::ALL_WEBSOCKET_PROTOCOLS)
605 .on_upgrade(move |stream| {
606 let stream = async_graphql_axum::GraphQLWebSocket::new(stream, schema, protocol)
607 .on_connection_init(move |_| async move {
608 let mut data = async_graphql::Data::default();
609 data.insert(user);
610 Ok(data)
611 });
612 async move {
613 stream.serve().await;
614 }
615 })
616}
617
618async fn graphql_playground(
619 axum::extract::Query(params): axum::extract::Query<std::collections::HashMap<String, String>>,
620 axum::extract::State(key): axum::extract::State<Option<Arc<String>>>,
621) -> axum::response::Response {
622 use axum::response::IntoResponse;
623
624 if let Some(ref expected) = key {
626 let provided = params.get("introspection-key");
627 if provided.map(|k| k.as_str()) != Some(expected.as_str()) {
628 return (
629 axum::http::StatusCode::FORBIDDEN,
630 "invalid or missing introspection-key",
631 )
632 .into_response();
633 }
634 }
635
636 let mut source = async_graphql::http::GraphiQLSource::build().endpoint("/graphql");
639 if let Some(ref k) = key {
640 source = source.header("X-Introspection-Key", k.as_str());
641 }
642
643 axum::response::Html(source.finish()).into_response()
644}
645
646pub fn cmd_serve_daemon(
648 port: Option<u16>,
649 bind: Option<std::net::IpAddr>,
650 subsonic_port: Option<u16>,
651 playground: bool,
652 mcp_bind: Option<std::net::SocketAddr>,
653) {
654 use std::fs;
655 use std::process::Command;
656
657 let cfg = Config::load().unwrap_or_default();
658 let port_val = port.unwrap_or(cfg.graphql.port);
659 let bind_val = bind.unwrap_or(cfg.graphql.bind);
660
661 let exe = std::env::current_exe().expect("failed to get current exe path");
662 let mut cmd = Command::new(exe);
663 cmd.arg("--headless");
665 cmd.arg("--port").arg(port_val.to_string());
666 cmd.arg("--bind").arg(bind_val.to_string());
667 if let Some(sp) = subsonic_port {
668 cmd.arg("--subsonic").arg(sp.to_string());
669 }
670 if let Some(addr) = mcp_bind {
671 cmd.arg("--mcp-bind").arg(addr.to_string());
672 }
673 if playground || cfg.graphql.playground {
674 cmd.arg("--playground");
675 }
676
677 cmd.stdin(std::process::Stdio::null());
678 cmd.stdout(std::process::Stdio::null());
679 cmd.stderr(std::process::Stdio::null());
680
681 let mut child = cmd.spawn().expect("failed to spawn daemon process");
682 let pid = child.id();
683
684 let pid_path = koan_core::config::config_dir().join("koan-serve.pid");
685 fs::write(&pid_path, pid.to_string()).ok();
686
687 std::thread::spawn(move || {
688 let _ = child.wait();
689 });
690
691 eprintln!("kōan daemon started (pid {}) on port {}", pid, port_val);
692 if let Some(sp) = subsonic_port {
693 eprintln!(" Subsonic REST on port {}", sp);
694 }
695 eprintln!(" PID file: {}", pid_path.display());
696}
697
698pub async fn execute_in_process(
707 schema: &KoanSchema,
708 query: &str,
709 variables: Option<serde_json::Value>,
710 user_id: i64,
711 role: koan_core::auth::Role,
712) -> serde_json::Value {
713 let mut request = async_graphql::Request::new(query);
714 request = request.data(AuthUser {
715 user_id,
716 role,
717 ..AuthUser::anonymous_admin()
718 });
719 if let Some(serde_json::Value::Object(map)) = variables {
720 let mut gql_vars = async_graphql::Variables::default();
721 for (k, v) in map {
722 gql_vars.insert(
723 async_graphql::Name::new(&k),
724 async_graphql::Value::from_json(v).unwrap_or(async_graphql::Value::Null),
725 );
726 }
727 request = request.variables(gql_vars);
728 }
729 let response = schema.execute(request).await;
730 serde_json::to_value(&response).unwrap_or(serde_json::Value::Null)
731}
732
733#[cfg(test)]
738mod tests {
739 use super::*;
740 use axum::body::Body;
741 use axum::http::{Request as HttpRequest, StatusCode};
742 use axum::routing::{get, post};
743 use tower::ServiceExt as _;
744
745 fn policy() -> Arc<BrowserPolicy> {
746 Arc::new(BrowserPolicy {
747 origins: vec!["https://music.example.com".into()],
748 hosts: vec!["koan.local".into()],
749 })
750 }
751
752 async fn ok() -> &'static str {
753 "ok"
754 }
755
756 fn routes() -> axum::Router<Arc<BrowserPolicy>> {
757 axum::Router::new()
758 .route("/graphql", post(ok).get(ok))
759 .route("/graphql/ws", get(ok))
760 }
761
762 async fn run_host(req: HttpRequest<Body>) -> StatusCode {
763 let app = routes()
764 .layer(axum::middleware::from_fn_with_state(policy(), host_guard))
765 .with_state(policy());
766 app.oneshot(req).await.unwrap().status()
767 }
768
769 async fn run_browser(req: HttpRequest<Body>) -> StatusCode {
770 let app = routes()
771 .layer(axum::middleware::from_fn_with_state(
772 policy(),
773 browser_guard,
774 ))
775 .with_state(policy());
776 app.oneshot(req).await.unwrap().status()
777 }
778
779 fn json_post(uri: &str) -> axum::http::request::Builder {
780 HttpRequest::post(uri).header(axum::http::header::CONTENT_TYPE, "application/json")
781 }
782
783 #[test]
786 fn host_policy_accepts_loopback_literals_and_configured_names() {
787 let p = policy();
788 assert!(p.host_allowed("localhost:4000"));
789 assert!(p.host_allowed("127.0.0.1:4000"));
790 assert!(p.host_allowed("192.168.1.20:4000"));
791 assert!(p.host_allowed("[::1]:4000"));
792 assert!(p.host_allowed("koan.local"));
793 assert!(p.host_allowed("koan.local:4000"));
794 }
795
796 #[test]
797 fn host_policy_rejects_attacker_controlled_names() {
798 let p = policy();
799 assert!(!p.host_allowed("evil.com"));
800 assert!(!p.host_allowed("rebind.evil.com:4000"));
801 assert!(!p.host_allowed("koan.local.evil.com"));
802 }
803
804 #[tokio::test]
805 async fn host_guard_rejects_foreign_host() {
806 let req = json_post("/graphql")
807 .header(axum::http::header::HOST, "rebind.evil.com")
808 .body(Body::empty())
809 .unwrap();
810 assert_eq!(run_host(req).await, StatusCode::FORBIDDEN);
811 }
812
813 #[tokio::test]
814 async fn host_guard_allows_known_host_and_missing_host() {
815 let req = json_post("/graphql")
816 .header(axum::http::header::HOST, "127.0.0.1:4000")
817 .body(Body::empty())
818 .unwrap();
819 assert_eq!(run_host(req).await, StatusCode::OK);
820
821 let req = json_post("/graphql").body(Body::empty()).unwrap();
822 assert_eq!(run_host(req).await, StatusCode::OK);
823 }
824
825 #[tokio::test]
828 async fn ws_upgrade_from_foreign_origin_is_rejected() {
829 let req = HttpRequest::get("/graphql/ws")
830 .header(axum::http::header::HOST, "127.0.0.1:4000")
831 .header(axum::http::header::ORIGIN, "https://evil.com")
832 .body(Body::empty())
833 .unwrap();
834 assert_eq!(run_browser(req).await, StatusCode::FORBIDDEN);
835 }
836
837 #[tokio::test]
838 async fn ws_upgrade_without_origin_is_allowed() {
839 let req = HttpRequest::get("/graphql/ws")
840 .header(axum::http::header::HOST, "127.0.0.1:4000")
841 .body(Body::empty())
842 .unwrap();
843 assert_eq!(run_browser(req).await, StatusCode::OK);
844 }
845
846 #[tokio::test]
847 async fn configured_and_same_origin_are_allowed() {
848 let req = HttpRequest::get("/graphql/ws")
849 .header(axum::http::header::HOST, "127.0.0.1:4000")
850 .header(axum::http::header::ORIGIN, "https://music.example.com")
851 .body(Body::empty())
852 .unwrap();
853 assert_eq!(run_browser(req).await, StatusCode::OK);
854
855 let req = json_post("/graphql")
857 .header(axum::http::header::HOST, "127.0.0.1:4000")
858 .header(axum::http::header::ORIGIN, "http://127.0.0.1:4000")
859 .body(Body::empty())
860 .unwrap();
861 assert_eq!(run_browser(req).await, StatusCode::OK);
862 }
863
864 #[tokio::test]
867 async fn text_plain_post_is_rejected() {
868 let req = HttpRequest::post("/graphql")
869 .header(axum::http::header::CONTENT_TYPE, "text/plain")
870 .body(Body::from(r#"{"query":"mutation{clearQueue{ok}}"}"#))
871 .unwrap();
872 assert_eq!(run_browser(req).await, StatusCode::UNSUPPORTED_MEDIA_TYPE);
873 }
874
875 #[tokio::test]
876 async fn post_without_content_type_is_rejected() {
877 let req = HttpRequest::post("/graphql").body(Body::empty()).unwrap();
878 assert_eq!(run_browser(req).await, StatusCode::UNSUPPORTED_MEDIA_TYPE);
879 }
880
881 #[tokio::test]
884 async fn load_perimeter_passes_requests_and_turns_panics_into_500s() {
885 async fn boom() -> &'static str {
886 panic!("resolver exploded");
887 }
888
889 let app = load_perimeter(
890 axum::Router::new()
891 .route("/graphql", post(ok))
892 .route("/boom", post(boom)),
893 );
894
895 let req = json_post("/graphql").body(Body::empty()).unwrap();
896 assert_eq!(
897 app.clone().oneshot(req).await.unwrap().status(),
898 StatusCode::OK
899 );
900
901 let req = json_post("/boom").body(Body::empty()).unwrap();
903 assert_eq!(
904 app.oneshot(req).await.unwrap().status(),
905 StatusCode::INTERNAL_SERVER_ERROR
906 );
907 }
908
909 #[tokio::test]
910 async fn json_post_is_accepted() {
911 let req = json_post("/graphql").body(Body::empty()).unwrap();
912 assert_eq!(run_browser(req).await, StatusCode::OK);
913
914 let req = HttpRequest::post("/graphql")
915 .header(
916 axum::http::header::CONTENT_TYPE,
917 "application/json; charset=utf-8",
918 )
919 .body(Body::empty())
920 .unwrap();
921 assert_eq!(run_browser(req).await, StatusCode::OK);
922 }
923}