use std::collections::HashMap;
use fraiseql_core::{
db::{
dialect::RowViewColumnType,
types::{ColumnSpec, ColumnValue},
},
schema::{CompiledSchema, FieldType, TypeDefinition},
security::SecurityContext,
};
use fraiseql_error::FraiseQLError;
use prost_reflect::{DynamicMessage, MessageDescriptor, ReflectMessage, Value};
use tracing::{debug, warn};
const MAX_GRPC_RESULT_ROWS: u32 = 10_000;
const DEFAULT_GRPC_LIMIT: u32 = 100;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum RpcKind {
Query {
returns_list: bool,
columns: Vec<ColumnSpec>,
row_descriptor: MessageDescriptor,
},
ServerStream {
columns: Vec<ColumnSpec>,
row_descriptor: MessageDescriptor,
},
Mutation {
function_name: String,
},
}
#[derive(Debug, Clone)]
pub struct RpcOperation {
pub operation_name: String,
pub type_name: String,
pub kind: RpcKind,
pub response_descriptor: MessageDescriptor,
}
pub type RpcDispatchTable = HashMap<String, RpcOperation>;
#[must_use]
pub const fn field_type_to_column_type(ft: &FieldType) -> Option<RowViewColumnType> {
match ft {
FieldType::String | FieldType::Scalar(_) | FieldType::Decimal | FieldType::Time => {
Some(RowViewColumnType::Text)
},
FieldType::Int => Some(RowViewColumnType::Int32),
FieldType::Float => Some(RowViewColumnType::Float64),
FieldType::Boolean => Some(RowViewColumnType::Boolean),
FieldType::Id | FieldType::Uuid => Some(RowViewColumnType::Uuid),
FieldType::DateTime => Some(RowViewColumnType::Timestamptz),
FieldType::Date => Some(RowViewColumnType::Date),
FieldType::Json => Some(RowViewColumnType::Json),
FieldType::Enum(_) => Some(RowViewColumnType::Text),
_ => None,
}
}
#[must_use]
pub fn column_specs_from_type(type_def: &TypeDefinition) -> Vec<ColumnSpec> {
type_def
.fields
.iter()
.filter_map(|f| {
field_type_to_column_type(&f.field_type).map(|ct| ColumnSpec {
name: f.name.to_string(),
column_type: ct,
})
})
.collect()
}
pub(crate) fn proto_value_to_json(value: &Value) -> serde_json::Value {
match value {
Value::Bool(b) => serde_json::Value::Bool(*b),
Value::I32(n) | Value::EnumNumber(n) => serde_json::json!(*n),
Value::I64(n) => serde_json::json!(*n),
Value::U32(n) => serde_json::json!(*n),
Value::U64(n) => serde_json::json!(*n),
Value::F32(f) => serde_json::json!(*f),
Value::F64(f) => serde_json::json!(*f),
Value::String(s) => serde_json::Value::String(s.clone()),
Value::Bytes(b) => serde_json::Value::String(base64_encode(b)),
Value::List(items) => {
serde_json::Value::Array(items.iter().map(proto_value_to_json).collect())
},
Value::Map(entries) => {
let obj: serde_json::Map<std::string::String, serde_json::Value> = entries
.iter()
.map(|(k, v)| (map_key_to_string(k), proto_value_to_json(v)))
.collect();
serde_json::Value::Object(obj)
},
Value::Message(inner) => dynamic_message_to_json(inner),
}
}
pub(crate) fn recase_keys_to_snake(value: serde_json::Value) -> serde_json::Value {
match value {
serde_json::Value::Object(map) => serde_json::Value::Object(
map.into_iter()
.map(|(k, v)| (fraiseql_core::utils::to_snake_case(&k), recase_keys_to_snake(v)))
.collect(),
),
serde_json::Value::Array(items) => {
serde_json::Value::Array(items.into_iter().map(recase_keys_to_snake).collect())
},
other => other,
}
}
fn base64_encode(bytes: &prost::bytes::Bytes) -> String {
use base64::Engine;
base64::engine::general_purpose::STANDARD.encode(bytes)
}
fn map_key_to_string(key: &prost_reflect::MapKey) -> String {
match key {
prost_reflect::MapKey::Bool(b) => b.to_string(),
prost_reflect::MapKey::I32(n) => n.to_string(),
prost_reflect::MapKey::I64(n) => n.to_string(),
prost_reflect::MapKey::U32(n) => n.to_string(),
prost_reflect::MapKey::U64(n) => n.to_string(),
prost_reflect::MapKey::String(s) => s.clone(),
}
}
fn dynamic_message_to_json(msg: &DynamicMessage) -> serde_json::Value {
serde_json::to_value(msg).unwrap_or(serde_json::Value::Null)
}
#[must_use]
pub fn extract_limit(msg: &DynamicMessage) -> u32 {
for field_desc in msg.descriptor().fields() {
if field_desc.name() == "limit" && msg.has_field(&field_desc) {
let val = msg.get_field(&field_desc);
if let Value::I32(n) = val.as_ref() {
let n = u32::try_from(*n).unwrap_or(DEFAULT_GRPC_LIMIT);
return n.min(MAX_GRPC_RESULT_ROWS);
}
if let Value::U32(n) = val.as_ref() {
return (*n).min(MAX_GRPC_RESULT_ROWS);
}
}
}
DEFAULT_GRPC_LIMIT
}
#[must_use]
pub fn extract_offset(msg: &DynamicMessage) -> Option<u32> {
for field_desc in msg.descriptor().fields() {
if field_desc.name() == "offset" && msg.has_field(&field_desc) {
let val = msg.get_field(&field_desc);
if let Value::I32(n) = val.as_ref() {
return u32::try_from(*n).ok();
}
if let Value::U32(n) = val.as_ref() {
return Some(*n);
}
}
}
None
}
#[must_use]
pub fn read_arguments(
msg: &DynamicMessage,
type_def: &TypeDefinition,
returns_list: bool,
) -> HashMap<String, serde_json::Value> {
let mut arguments = HashMap::new();
let mut filters = serde_json::Map::new();
for field_desc in msg.descriptor().fields() {
let field_name = field_desc.name();
if matches!(field_name, "limit" | "offset" | "order_by") {
continue;
}
if type_def.find_field(field_name).is_none() || !msg.has_field(&field_desc) {
continue;
}
let value = msg.get_field(&field_desc);
filters.insert(
field_name.to_string(),
serde_json::json!({ "eq": proto_value_to_json(&value) }),
);
}
if !filters.is_empty() {
arguments.insert("where".to_string(), serde_json::Value::Object(filters));
}
let limit = if returns_list { extract_limit(msg) } else { 1 };
arguments.insert("limit".to_string(), serde_json::json!(limit));
if let Some(offset) = extract_offset(msg) {
arguments.insert("offset".to_string(), serde_json::json!(offset));
}
if let Some((field, direction)) = extract_order_by_pair(msg, type_def) {
arguments.insert("orderBy".to_string(), serde_json::json!({ field: direction }));
}
arguments
}
#[must_use]
pub fn extract_order_by_pair(
msg: &DynamicMessage,
type_def: &TypeDefinition,
) -> Option<(String, String)> {
for field_desc in msg.descriptor().fields() {
if field_desc.name() != "order_by" || !msg.has_field(&field_desc) {
continue;
}
let val = msg.get_field(&field_desc);
let Value::String(s) = val.as_ref() else {
continue;
};
let mut parts = s.split_whitespace();
let col_name = parts.next()?;
if type_def.find_field(col_name).is_none() {
warn!(column = %col_name, "gRPC order_by references unknown column — ignoring");
return None;
}
let direction = parts
.next()
.filter(|d| d.eq_ignore_ascii_case("asc") || d.eq_ignore_ascii_case("desc"))
.map_or("ASC", |d| {
if d.eq_ignore_ascii_case("desc") {
"DESC"
} else {
"ASC"
}
});
return Some((col_name.to_string(), direction.to_string()));
}
None
}
pub async fn execute_grpc_read(
executor: &fraiseql_core::runtime::Executor,
query_name: &str,
columns: &[ColumnSpec],
returns_list: bool,
request_msg: &DynamicMessage,
type_def: &TypeDefinition,
security_context: Option<&SecurityContext>,
) -> Result<fraiseql_core::runtime::RowRead, FraiseQLError> {
let query_match = grpc_query_match(
executor.schema(),
query_name,
columns,
returns_list,
request_msg,
type_def,
)?;
debug!(
query = %query_name,
columns = columns.len(),
user_id = ?security_context.map(|c| &c.user_id),
"Executing gRPC row read through the engine"
);
executor.execute_row_read(&query_match, None, security_context, columns).await
}
pub(super) fn grpc_query_match(
schema: &CompiledSchema,
query_name: &str,
columns: &[ColumnSpec],
returns_list: bool,
request_msg: &DynamicMessage,
type_def: &TypeDefinition,
) -> Result<fraiseql_core::runtime::QueryMatch, FraiseQLError> {
let query_def = schema.find_query(query_name).cloned().ok_or_else(|| {
FraiseQLError::validation(format!("Query '{query_name}' not found in schema"))
})?;
let fields: Vec<String> = columns.iter().map(|c| c.name.clone()).collect();
let arguments = read_arguments(request_msg, type_def, returns_list);
fraiseql_core::runtime::QueryMatch::from_operation(query_def, fields, arguments, Some(type_def))
}
pub async fn execute_grpc_mutation(
executor: &fraiseql_core::runtime::Executor,
mutation_name: &str,
request_msg: &DynamicMessage,
recase_input_keys: bool,
security_context: Option<&fraiseql_core::security::SecurityContext>,
) -> Result<MutationResult, FraiseQLError> {
let variables =
grpc_mutation_variables(executor.schema(), mutation_name, request_msg, recase_input_keys);
debug!(
mutation = %mutation_name,
arg_count = variables.as_object().map_or(0, serde_json::Map::len),
"Executing gRPC mutation through the chokepoint"
);
let selections =
fraiseql_core::runtime::mutation_return_selections(executor.schema(), mutation_name);
let execution = executor
.execute_mutation_as(
mutation_name,
Some(&variables),
security_context,
fraiseql_core::runtime::WriteSelections::new(&selections)?,
)
.await?;
Ok(mutation_result_from_outcome(&execution.outcome))
}
fn grpc_mutation_variables(
schema: &fraiseql_core::schema::CompiledSchema,
mutation_name: &str,
request_msg: &DynamicMessage,
recase_input_keys: bool,
) -> serde_json::Value {
let declared: Vec<String> = schema
.find_mutation(mutation_name)
.map(|m| m.arguments.iter().map(|a| a.name.clone()).collect())
.unwrap_or_default();
let mut out = serde_json::Map::new();
for field in request_msg.descriptor().fields() {
if !request_msg.has_field(&field) {
continue;
}
let value = proto_value_to_json(request_msg.get_field(&field).as_ref());
let value = if recase_input_keys {
recase_keys_to_snake(value)
} else {
value
};
let proto_name = field.name();
let matched = declared.iter().find(|name| {
name.as_str() == proto_name || fraiseql_core::utils::to_snake_case(name) == proto_name
});
match matched {
Some(name) => {
out.insert(name.clone(), value);
},
None => {
out.insert(proto_name.to_string(), value);
},
}
}
serde_json::Value::Object(out)
}
fn mutation_result_from_outcome(
outcome: &fraiseql_core::runtime::mutation_result::MutationOutcome,
) -> MutationResult {
use fraiseql_core::runtime::mutation_result::MutationOutcome;
match outcome {
MutationOutcome::Success { entity_id, .. } => MutationResult {
success: true,
id: entity_id.clone(),
error: None,
},
MutationOutcome::Error { message, .. } => MutationResult {
success: false,
id: None,
error: Some(message.clone()),
},
_ => MutationResult {
success: false,
id: None,
error: Some("mutation outcome not recognised by this build".to_string()),
},
}
}
#[derive(Debug)]
pub struct MutationResult {
pub success: bool,
pub id: Option<String>,
pub error: Option<String>,
}
#[must_use]
pub fn encode_mutation_response(
result: &MutationResult,
response_desc: &MessageDescriptor,
) -> DynamicMessage {
let mut msg = DynamicMessage::new(response_desc.clone());
if let Some(field) = response_desc.get_field_by_name("success") {
msg.set_field(&field, Value::Bool(result.success));
}
if let (Some(field), Some(id)) = (response_desc.get_field_by_name("id"), &result.id) {
msg.set_field(&field, Value::String(id.clone()));
}
if let (Some(field), Some(err)) = (response_desc.get_field_by_name("error"), &result.error) {
msg.set_field(&field, Value::String(err.clone()));
}
msg
}
#[must_use]
pub fn encode_row(
row: &[ColumnValue],
columns: &[ColumnSpec],
row_desc: &MessageDescriptor,
) -> DynamicMessage {
let mut msg = DynamicMessage::new(row_desc.clone());
for (col_val, col_spec) in row.iter().zip(columns.iter()) {
if let Some(field_desc) = row_desc.get_field_by_name(&col_spec.name) {
let proto_val = column_value_to_proto(col_val);
if let Some(v) = proto_val {
msg.set_field(&field_desc, v);
}
}
}
msg
}
#[must_use]
pub fn column_value_to_proto(col: &ColumnValue) -> Option<Value> {
match col {
ColumnValue::Null => None,
ColumnValue::Text(s) => Some(Value::String(s.clone())),
ColumnValue::Int32(n) => Some(Value::I32(*n)),
ColumnValue::Int64(n) => Some(Value::I64(*n)),
ColumnValue::Float64(f) => Some(Value::F64(*f)),
ColumnValue::Boolean(b) => Some(Value::Bool(*b)),
ColumnValue::Uuid(u) => Some(Value::String(u.clone())),
ColumnValue::Timestamptz(ts) => {
Some(Value::String(ts.clone()))
},
ColumnValue::Date(d) => Some(Value::String(d.clone())),
ColumnValue::Json(v) => Some(Value::String(v.clone())),
}
}
#[must_use]
pub fn encode_response(
rows: Vec<Vec<ColumnValue>>,
columns: &[ColumnSpec],
returns_list: bool,
row_descriptor: &MessageDescriptor,
response_descriptor: &MessageDescriptor,
) -> DynamicMessage {
let mut response = DynamicMessage::new(response_descriptor.clone());
if returns_list {
let items: Vec<Value> = rows
.iter()
.map(|row| {
let row_msg = encode_row(row, columns, row_descriptor);
Value::Message(row_msg)
})
.collect();
for field_desc in response_descriptor.fields() {
if field_desc.is_list() && field_desc.kind().as_message().is_some() {
response.set_field(&field_desc, Value::List(items));
break;
}
}
} else {
if let Some(row) = rows.into_iter().next() {
for (col_val, col_spec) in row.iter().zip(columns.iter()) {
if let Some(field_desc) = response_descriptor.get_field_by_name(&col_spec.name) {
if let Some(v) = column_value_to_proto(col_val) {
response.set_field(&field_desc, v);
}
}
}
}
}
response
}
pub fn build_dispatch_table(
schema: &CompiledSchema,
service_name: &str,
pool: &prost_reflect::DescriptorPool,
) -> Result<RpcDispatchTable, FraiseQLError> {
let mut table = HashMap::new();
let service_desc = pool.get_service_by_name(service_name).ok_or_else(|| {
FraiseQLError::validation(format!(
"gRPC service '{service_name}' not found in descriptor pool"
))
})?;
for method_desc in service_desc.methods() {
let method_name = method_desc.name().to_string();
let full_method = format!("/{service_name}/{method_name}");
let response_desc = method_desc.output();
if method_name.starts_with("Get") || method_name.starts_with("List") {
let query_name = grpc_method_to_query_name(&method_name);
if let Some(query_def) = schema.find_query(&query_name) {
if query_def.function.is_some() {
warn!(
method = %method_name,
query = %query_name,
"gRPC does not carry function-backed queries (#1329) — the method \
would read the type's view instead of invoking the function; \
skipping"
);
continue;
}
let type_name = &query_def.return_type;
let Some(type_def) = schema.find_type(type_name) else {
warn!(
method = %method_name,
type_name = %type_name,
"gRPC query return type not found in schema — skipping"
);
continue;
};
let columns = column_specs_from_type(type_def);
let is_server_streaming = method_desc.is_server_streaming();
let kind = if is_server_streaming && query_def.returns_list {
RpcKind::ServerStream {
columns,
row_descriptor: response_desc.clone(),
}
} else {
let row_desc = if query_def.returns_list {
response_desc
.fields()
.find(|f| f.is_list() && f.kind().as_message().is_some())
.and_then(|f| f.kind().as_message().cloned())
.unwrap_or_else(|| response_desc.clone())
} else {
response_desc.clone()
};
RpcKind::Query {
returns_list: query_def.returns_list,
columns,
row_descriptor: row_desc,
}
};
table.insert(
full_method,
RpcOperation {
operation_name: query_name,
type_name: type_name.clone(),
kind,
response_descriptor: response_desc,
},
);
continue;
}
}
let mutation_name = grpc_method_to_mutation_name(&method_name);
if let Some(mutation_def) = schema.find_mutation(&mutation_name) {
let function_name =
mutation_def.sql_source.clone().unwrap_or_else(|| format!("fn_{mutation_name}"));
table.insert(
full_method,
RpcOperation {
operation_name: mutation_name,
type_name: mutation_def.return_type.clone(),
kind: RpcKind::Mutation { function_name },
response_descriptor: response_desc,
},
);
continue;
}
debug!(
method = %method_name,
"gRPC method has no matching query or mutation — skipping"
);
}
Ok(table)
}
pub(crate) fn grpc_method_to_query_name(method: &str) -> String {
let name = method
.strip_prefix("Get")
.or_else(|| method.strip_prefix("List"))
.unwrap_or(method);
let mut result = String::with_capacity(name.len());
for (i, ch) in name.chars().enumerate() {
if ch.is_uppercase() && i > 0 {
result.push('_');
}
result.push(ch.to_ascii_lowercase());
}
result
}
pub(crate) fn grpc_method_to_mutation_name(method: &str) -> String {
let mut chars = method.chars();
match chars.next() {
Some(first) => {
let mut result = first.to_lowercase().to_string();
result.extend(chars);
result
},
None => String::new(),
}
}