diff --git a/nix/shells/rust.nix b/nix/shells/rust.nix index 7803c0c..bedc76f 100644 --- a/nix/shells/rust.nix +++ b/nix/shells/rust.nix @@ -9,6 +9,7 @@ mkShell, bacon, cargo-workspaces, + diesel-cli, rust-bin, defaultShell, }: let @@ -23,6 +24,7 @@ packages = [ bacon cargo-workspaces + diesel-cli rust rustfmt ]; diff --git a/rust/f7c1a2d347e4c52d5fb8d10cb4d94b5884e546fb.tar.gz b/rust/f7c1a2d347e4c52d5fb8d10cb4d94b5884e546fb.tar.gz new file mode 100644 index 0000000..e69de29 diff --git a/rust/henka-api/src/main.rs b/rust/henka-api/src/main.rs index 318cc22..d57ed8d 100644 --- a/rust/henka-api/src/main.rs +++ b/rust/henka-api/src/main.rs @@ -1,4 +1,4 @@ -use std::{net::SocketAddr, sync::Arc, time::Duration}; +use std::{net::SocketAddr, sync::Arc}; use axum::{ Extension, Router, @@ -8,14 +8,16 @@ use axum::{ response::{Html, Response}, routing::{MethodFilter, MethodRouter, get, on}, }; -use futures::stream::{BoxStream, StreamExt as _}; -use juniper::{EmptyMutation, FieldError, RootNode, graphql_object, graphql_subscription}; +use juniper::EmptyMutation; use juniper_axum::{ extract::JuniperRequest, graphiql, graphql, playground, response::JuniperResponse, ws, }; use juniper_graphql_ws::ConnectionConfig; -use tokio::{net::TcpListener, time::interval}; -use tokio_stream::wrappers::IntervalStream; +use tokio::net::TcpListener; + +use crate::public::{Query as PubQuery, Schema as PubSchema, Subscription as PubSub}; + +mod public; const AUTH_TOKEN: &str = "henka-demo-token"; @@ -26,42 +28,6 @@ struct Context { impl juniper::Context for Context {} -#[derive(Clone, Copy, Debug)] -struct Query; - -#[graphql_object(context = Context)] -impl Query { - fn add(a: i32, b: i32) -> i32 { - a + b - } - - fn mul(context: &Context, a: i32, b: i32) -> Result { - if !context.authorized { - return Err(FieldError::from("unauthorized")); - } - Ok(a * b) - } -} - -#[derive(Clone, Copy, Debug)] -struct Subscription; - -type NumberStream = BoxStream<'static, Result>; - -#[graphql_subscription(context = Context)] -impl Subscription { - async fn count() -> NumberStream { - let mut value = 0; - let stream = IntervalStream::new(interval(Duration::from_secs(1))).map(move |_| { - value += 1; - Ok(value) - }); - Box::pin(stream) - } -} - -type Schema = RootNode, Subscription>; - async fn homepage() -> Html<&'static str> { "

juniper_axum/simple example

\
visit GraphiQL
\ @@ -111,7 +77,7 @@ async fn auth_middleware(mut request: Request, next: Next) -> Result>, + Extension(schema): Extension>, Extension(context): Extension, JuniperRequest(req): JuniperRequest, ) -> JuniperResponse { @@ -124,10 +90,10 @@ async fn main() { .with_max_level(tracing::Level::INFO) .init(); - let schema = Arc::new(Schema::new( - Query, + let schema = Arc::new(PubSchema::new( + PubQuery, EmptyMutation::::new(), - Subscription, + PubSub, )); let app = Router::new() @@ -136,9 +102,9 @@ async fn main() { on(MethodFilter::GET.or(MethodFilter::POST), graphql_handler), ) .route_layer(middleware::from_fn(auth_middleware)) - .route("/priv/graphql", schema_router!(Schema)) - .route("/subscriptions", subscription_router!(Schema)) - .route("/priv/subscriptions", subscription_router!(Schema)) + .route("/priv/graphql", schema_router!(PubSchema)) + .route("/subscriptions", subscription_router!(PubSchema)) + .route("/priv/subscriptions", subscription_router!(PubSchema)) .route("/graphiql", gw("/graphql", "/subscriptions")) .route("/priv/graphiql", gw("/priv/graphql", "/priv/subscriptions")) .route("/playground", pg("/graphql", "/subscriptions")) @@ -158,89 +124,3 @@ async fn main() { .await .unwrap_or_else(|e| panic!("failed to run `axum::serve`: {e}")); } - -#[cfg(test)] -mod tests { - use super::*; - use axum::{body::to_bytes, http::StatusCode}; - use serde_json::Value; - use tower::ServiceExt; - - fn build_app() -> Router { - let schema = Arc::new(Schema::new( - Query, - EmptyMutation::::new(), - Subscription, - )); - - Router::new() - .route( - "/graphql", - on(MethodFilter::GET.or(MethodFilter::POST), graphql_handler), - ) - .route_layer(middleware::from_fn(auth_middleware)) - .layer(Extension(schema)) - } - - #[tokio::test] - async fn test_add_without_auth() { - let (status, body) = query(build_app(), "{ add(a: 1, b: 2) }", None).await; - assert_eq!(status, StatusCode::OK); - assert_eq!(body["data"]["add"], 3); - } - - #[tokio::test] - async fn test_mul_without_auth_returns_error() { - let (status, body) = query(build_app(), "{ mul(a: 3, b: 4) }", None).await; - assert_eq!(status, StatusCode::OK); - assert!( - body.get("errors").is_some(), - "expected GraphQL error for unauthorized access" - ); - assert!(body["data"]["mul"].is_null()); - } - - #[tokio::test] - async fn test_mul_with_wrong_token_returns_error() { - let (status, body) = query( - build_app(), - "{ mul(a: 3, b: 4) }", - Some("Bearer wrong-token"), - ) - .await; - assert_eq!(status, StatusCode::OK); - assert!( - body.get("errors").is_some(), - "expected GraphQL error for wrong token" - ); - } - - #[tokio::test] - async fn test_mul_with_valid_token() { - let (status, body) = query(build_app(), "{ mul(a: 3, b: 4) }", Some(AUTH_TOKEN)).await; - assert_eq!(status, StatusCode::OK); - assert_eq!(body["data"]["mul"], 12); - } - - async fn query(app: Router, q: &str, token: Option<&str>) -> (StatusCode, Value) { - let mut req_builder = axum::http::Request::builder() - .method("POST") - .uri("/graphql") - .header("content-type", "application/json"); - - if let Some(token) = token { - let bearer = format!("bearer {}", token); - req_builder = req_builder.header("authorization", bearer); - } - - let req = req_builder - .body(Body::from(serde_json::json!({"query": q}).to_string())) - .unwrap(); - - let res = app.oneshot(req).await.unwrap(); - let status = res.status(); - let bytes = to_bytes(res.into_body(), usize::MAX).await.unwrap(); - let body: Value = serde_json::from_slice(&bytes).unwrap(); - (status, body) - } -} diff --git a/rust/henka-api/src/public.rs b/rust/henka-api/src/public.rs new file mode 100644 index 0000000..95aa2f3 --- /dev/null +++ b/rust/henka-api/src/public.rs @@ -0,0 +1,145 @@ +use std::time::Duration; + +use futures::stream::BoxStream; +use futures::stream::StreamExt; +use juniper::{EmptyMutation, FieldError, RootNode, graphql_object, graphql_subscription}; +use tokio::time::interval; +use tokio_stream::wrappers::IntervalStream; + +use crate::Context; + +#[derive(Clone, Copy, Debug)] +pub(crate) struct Query; + +#[graphql_object(context = Context)] +impl Query { + fn add(a: i32, b: i32) -> i32 { + a + b + } + + fn mul(context: &Context, a: i32, b: i32) -> Result { + if !context.authorized { + return Err(FieldError::from("unauthorized")); + } + Ok(a * b) + } +} + +#[derive(Clone, Copy, Debug)] +pub(crate) struct Subscription; + +type NumberStream = BoxStream<'static, Result>; + +#[graphql_subscription(context = Context)] +impl Subscription { + async fn count() -> NumberStream { + let mut value = 0; + let stream = IntervalStream::new(interval(Duration::from_secs(1))).map(move |_| { + value += 1; + Ok(value) + }); + Box::pin(stream) + } +} + +pub(crate) type Schema = RootNode, Subscription>; + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use axum::{ + Extension, Router, + body::Body, + body::to_bytes, + http::StatusCode, + middleware, + routing::{MethodFilter, on}, + }; + use serde_json::Value; + use tower::ServiceExt; + + use super::*; + + fn build_app() -> Router { + let schema = Arc::new(Schema::new( + Query, + EmptyMutation::::new(), + Subscription, + )); + + Router::new() + .route( + "/graphql", + on( + MethodFilter::GET.or(MethodFilter::POST), + crate::graphql_handler, + ), + ) + .route_layer(middleware::from_fn(crate::auth_middleware)) + .layer(Extension(schema)) + } + + #[tokio::test] + async fn test_add_without_auth() { + let (status, body) = query(build_app(), "{ add(a: 1, b: 2) }", None).await; + assert_eq!(status, StatusCode::OK); + assert_eq!(body["data"]["add"], 3); + } + + #[tokio::test] + async fn test_mul_without_auth_returns_error() { + let (status, body) = query(build_app(), "{ mul(a: 3, b: 4) }", None).await; + assert_eq!(status, StatusCode::OK); + assert!( + body.get("errors").is_some(), + "expected GraphQL error for unauthorized access" + ); + assert!(body["data"]["mul"].is_null()); + } + + #[tokio::test] + async fn test_mul_with_wrong_token_returns_error() { + let (status, body) = query( + build_app(), + "{ mul(a: 3, b: 4) }", + Some("Bearer wrong-token"), + ) + .await; + assert_eq!(status, StatusCode::OK); + assert!( + body.get("errors").is_some(), + "expected GraphQL error for wrong token" + ); + } + + #[tokio::test] + async fn test_mul_with_valid_token() { + let (status, body) = + query(build_app(), "{ mul(a: 3, b: 4) }", Some(crate::AUTH_TOKEN)).await; + assert_eq!(status, StatusCode::OK); + assert_eq!(body["data"]["mul"], 12); + } + + async fn query(app: Router, q: &str, token: Option<&str>) -> (StatusCode, Value) { + let mut req_builder = axum::http::Request::builder() + .method("POST") + .uri("/graphql") + .header("content-type", "application/json"); + + if let Some(token) = token { + let bearer = format!("bearer {}", token); + req_builder = req_builder.header("authorization", bearer); + } + + let req = req_builder + .body(Body::from(serde_json::json!({"query": q}).to_string())) + .unwrap(); + + let res = app.oneshot(req).await.unwrap(); + let status = res.status(); + let bytes = to_bytes(res.into_body(), usize::MAX).await.unwrap(); + let body: Value = serde_json::from_slice(&bytes).unwrap(); + (status, body) + } +}