use async_graphql::dynamic::{
Field, FieldFuture, InputValue, Object, ResolverContext, Schema, TypeRef,
};
use async_graphql::Value;
use crate::{GraphQLField, GraphQLSchema, GraphQLType};
fn parse_type_ref(type_name: &str) -> Result<TypeRef, String> {
let unsupported = || format!("Unsupported type reference '{type_name}'");
let trimmed = type_name.trim();
let (inner, non_null) = match trimmed.strip_suffix('!') {
Some(rest) => (rest.trim(), true),
None => (trimmed, false),
};
if let Some(list_inner) = inner.strip_prefix('[').and_then(|s| s.strip_suffix(']')) {
let list_inner = list_inner.trim();
let (item, item_non_null) = match list_inner.strip_suffix('!') {
Some(rest) => (rest.trim(), true),
None => (list_inner, false),
};
if item.is_empty() || item.starts_with('[') {
return Err(unsupported());
}
Ok(match (item_non_null, non_null) {
(false, false) => TypeRef::named_list(item),
(true, false) => TypeRef::named_nn_list(item),
(false, true) => TypeRef::named_list_nn(item),
(true, true) => TypeRef::named_nn_list_nn(item),
})
} else if inner.is_empty() {
Err(unsupported())
} else if non_null {
Ok(TypeRef::named_nn(inner))
} else {
Ok(TypeRef::named(inner))
}
}
fn mock_payload(field_name: &str, is_list: bool) -> Value {
let object = |id: &str| {
serde_json::json!({
"id": id,
"name": format!("{field_name}_{id}"),
"createdAt": "2024-01-01T00:00:00Z",
"updatedAt": "2024-01-01T00:00:00Z",
})
};
let json = if is_list {
serde_json::json!([object("1"), object("2")])
} else {
object("1")
};
Value::from_json(json).unwrap_or(Value::Null)
}
fn mock_root_field(field: &GraphQLField) -> Result<Field, String> {
let type_ref = parse_type_ref(&field.type_name)?;
let is_list = field.type_name.trim_start().starts_with('[');
let value = mock_payload(&field.name, is_list);
let mut root = Field::new(field.name.clone(), type_ref, move |_ctx| {
FieldFuture::from_value(Some(value.clone()))
});
if !is_list {
root = root.argument(InputValue::new("id", TypeRef::named(TypeRef::ID)));
}
Ok(root)
}
fn resolver_root_field(
field: &GraphQLField,
resolver: crate::resolver::SharedDbResolver,
) -> Result<Field, String> {
let type_ref = parse_type_ref(&field.type_name)?;
let is_list = field.type_name.trim_start().starts_with('[');
let field_name = field.name.clone();
let type_name = field.type_name.clone();
let resolver = resolver;
let mut root = Field::new(
field.name.clone(),
type_ref,
move |ctx: ResolverContext<'_>| {
let mut args = serde_json::Map::new();
for (key, val) in ctx.args.iter() {
if let Ok(json_v) = val.as_value().clone().into_json() {
args.insert(key.to_string(), json_v);
}
}
let resolver_ctx = crate::resolver::ResolverContext {
field_name: field_name.clone(),
type_name: type_name.clone(),
is_list,
args: serde_json::Value::Object(args),
};
let resolver = resolver.clone();
FieldFuture::new(async move {
match resolver.resolve_query(&resolver_ctx).await {
Ok(value) => {
let gql_value = Value::from_json(value).unwrap_or(Value::Null);
Ok(Some(gql_value))
}
Err(msg) => {
tracing::error!(
field = %resolver_ctx.field_name,
error = %msg,
"DB resolver failed"
);
Ok(Some(Value::Null))
}
}
})
},
);
if !is_list {
root = root.argument(InputValue::new("id", TypeRef::named(TypeRef::ID)));
}
Ok(root)
}
fn object_type(t: &GraphQLType) -> Result<Object, String> {
let mut obj = Object::new(t.name.clone());
for field in &t.fields {
let type_ref = parse_type_ref(&field.type_name)?;
let field_name = field.name.clone();
obj = obj.field(Field::new(
field.name.clone(),
type_ref,
move |ctx: ResolverContext<'_>| {
let value = ctx
.parent_value
.try_to_value()
.ok()
.and_then(|parent| match parent {
Value::Object(map) => map.get(field_name.as_str()).cloned(),
_ => None,
});
FieldFuture::from_value(value)
},
));
}
Ok(obj)
}
pub fn build_dynamic_schema(
schema: &GraphQLSchema,
resolver: Option<&crate::resolver::SharedDbResolver>,
) -> Result<Schema, String> {
let mutation_name = if schema.mutations.is_empty() {
None
} else {
Some("Mutation")
};
let mut builder = Schema::build("Query", mutation_name, None);
for t in &schema.types {
builder = builder.register(object_type(t)?);
}
let mut query = Object::new("Query");
for field in &schema.queries {
let f = match resolver {
Some(r) => resolver_root_field(field, std::sync::Arc::clone(r))?,
None => mock_root_field(field)?,
};
query = query.field(f);
}
builder = builder.register(query);
if mutation_name.is_some() {
let mut mutation = Object::new("Mutation");
for field in &schema.mutations {
let f = match resolver {
Some(r) => resolver_root_field(field, std::sync::Arc::clone(r))?,
None => mock_root_field(field)?,
};
mutation = mutation.field(f);
}
builder = builder.register(mutation);
}
builder.finish().map_err(|e| e.to_string())
}
pub fn router(schema: Schema) -> axum::Router {
async fn graphql_handler(
axum::extract::State(schema): axum::extract::State<Schema>,
request: async_graphql_axum::GraphQLRequest,
) -> async_graphql_axum::GraphQLResponse {
let inner = request.into_inner();
#[cfg(feature = "graphql-complexity")]
{
if let Ok(ir) = crate::query_ir::parse_query(&inner.query, None) {
let calculator = crate::complexity::ComplexityCalculator::with_defaults();
let result = calculator.calculate(&ir);
if let Some(err) = result.exceeded {
let resp = async_graphql::Response::from_errors(vec![
async_graphql::ServerError::new(format!("query rejected: {err}"), None),
]);
return resp.into();
}
}
}
schema.execute(inner).await.into()
}
axum::Router::new()
.route("/graphql", axum::routing::post(graphql_handler))
.with_state(schema)
}
#[allow(dead_code)]
pub async fn execute_async(schema: &Schema, query: &str) -> Result<serde_json::Value, String> {
let response = schema.execute(query).await;
response_to_json(response)
}
fn response_to_json(response: async_graphql::Response) -> Result<serde_json::Value, String> {
if !response.errors.is_empty() {
return Err(response
.errors
.iter()
.map(|e| e.message.as_str())
.collect::<Vec<_>>()
.join("; "));
}
let data = response.data.into_json().map_err(|e| e.to_string())?;
match data {
serde_json::Value::Object(map) => Ok(map
.into_iter()
.next()
.map(|(_, value)| value)
.unwrap_or(serde_json::Value::Null)),
other => Ok(other),
}
}
pub fn execute(schema: &Schema, query: &str) -> Result<serde_json::Value, String> {
let response = std::thread::scope(|scope| {
scope
.spawn(|| -> Result<async_graphql::Response, String> {
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.map_err(|e| format!("Failed to build tokio runtime: {e}"))?;
Ok(runtime.block_on(schema.execute(query)))
})
.join()
.map_err(|_| "GraphQL executor thread panicked".to_string())?
})?;
response_to_json(response)
}
#[cfg(test)]
mod tests {
use crate::resolver::{DbResolver, ResolverContext, SharedDbResolver};
use crate::{GraphQLSchemaGenerator, GraphQLServer};
use serde_json::json;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
struct TestDbResolver;
impl DbResolver for TestDbResolver {
fn resolve_query(
&self,
ctx: &ResolverContext,
) -> Pin<Box<dyn Future<Output = Result<serde_json::Value, String>> + Send>> {
let field_name = ctx.field_name.clone();
let is_list = ctx.is_list;
Box::pin(async move {
if is_list {
Ok(json!([
{"id": "100", "name": format!("{}_real_100", field_name), "createdAt": "2024-01-01T00:00:00Z", "updatedAt": "2024-01-01T00:00:00Z"},
{"id": "200", "name": format!("{}_real_200", field_name), "createdAt": "2024-01-01T00:00:00Z", "updatedAt": "2024-01-01T00:00:00Z"}
]))
} else {
Ok(json!({
"id": "100",
"name": format!("{}_real", field_name),
"createdAt": "2024-01-01T00:00:00Z",
"updatedAt": "2024-01-01T00:00:00Z"
}))
}
})
}
}
#[test]
fn test_resolver_returns_real_data_single() {
let resolver: SharedDbResolver = Arc::new(TestDbResolver);
let srv = GraphQLServer::new(4501)
.with_schema(GraphQLSchemaGenerator::generate_schema(&["users"]))
.with_db_resolver(resolver);
let result = srv.execute_query("{ getUser(id: 1) { id name } }");
assert!(result.is_ok(), "expected ok, got {:?}", result);
let v = result.unwrap();
assert_eq!(v["id"], "100");
assert!(
v["name"].as_str().unwrap().contains("real"),
"name should contain 'real': {}",
v["name"]
);
}
#[test]
fn test_resolver_returns_real_data_list() {
let resolver: SharedDbResolver = Arc::new(TestDbResolver);
let srv = GraphQLServer::new(4502)
.with_schema(GraphQLSchemaGenerator::generate_schema(&["users"]))
.with_db_resolver(resolver);
let result = srv.execute_query("{ listUsers { id name } }");
assert!(result.is_ok(), "expected ok, got {:?}", result);
let v = result.unwrap();
assert!(v.is_array());
let arr = v.as_array().unwrap();
assert_eq!(arr.len(), 2);
assert_eq!(arr[0]["id"], "100");
assert_eq!(arr[1]["id"], "200");
assert!(arr[0]["name"].as_str().unwrap().contains("real"));
}
#[test]
fn test_no_resolver_falls_back_to_mock() {
let srv = GraphQLServer::new(4503)
.with_schema(GraphQLSchemaGenerator::generate_schema(&["users"]));
let result = srv.execute_query("{ getUser(id: 1) { id name } }");
assert!(result.is_ok(), "expected ok, got {:?}", result);
let v = result.unwrap();
assert_eq!(v["id"], "1"); assert!(
v["name"].as_str().unwrap().contains("getUser"),
"mock name should contain field name"
);
}
}