use std::sync::Arc;
use axum::extract::{Request, State};
use axum::response::Response;
use axum::routing::any;
use axum::{Extension, Router};
use surrealdb_core::dbs::Session;
use surrealdb_protocol::proto::rpc::v1::surreal_db_service_server::SurrealDbServiceServer;
use surrealdb_rpc::capabilities::RouteTarget;
use tonic::codec::CompressionEncoding;
use tower_service::Service;
use super::AppState;
use crate::cnf::GRPC_MAX_MESSAGE_SIZE;
use crate::ntw::error::Error as NetError;
use crate::rpc::RpcState;
use crate::rpc::grpc::GrpcService;
const SERVICE_ROUTE: &str = "/surrealdb.protocol.rpc.v1.SurrealDBService/{*method}";
const COMPRESSED_ENCODINGS: &[(&str, CompressionEncoding)] =
&[("zstd", CompressionEncoding::Zstd), ("gzip", CompressionEncoding::Gzip)];
pub fn accepted_message_encodings() -> Vec<String> {
COMPRESSED_ENCODINGS
.iter()
.map(|(token, _)| (*token).to_string())
.chain(std::iter::once("identity".to_string()))
.collect()
}
pub fn router() -> Router<Arc<RpcState>> {
Router::new().route(SERVICE_ROUTE, any(handler))
}
async fn handler(
Extension(state): Extension<AppState>,
Extension(session): Extension<Session>,
State(rpc_state): State<Arc<RpcState>>,
request: Request,
) -> Result<Response, NetError> {
if !state.datastore.allows_http_route(&RouteTarget::Rpc) {
warn!("Capabilities denied gRPC route request attempt, target: '{}'", &RouteTarget::Rpc);
return Err(NetError::ForbiddenRoute(RouteTarget::Rpc.to_string()));
}
let mut service = SurrealDbServiceServer::new(GrpcService::new(rpc_state, session))
.max_decoding_message_size(*GRPC_MAX_MESSAGE_SIZE)
.max_encoding_message_size(*GRPC_MAX_MESSAGE_SIZE);
for (_, encoding) in COMPRESSED_ENCODINGS {
service = service.accept_compressed(*encoding).send_compressed(*encoding);
}
let response = match service.call(request).await {
Ok(response) => response,
Err(error) => match error {},
};
Ok(response.map(axum::body::Body::new))
}
#[cfg(test)]
mod tests {
use surrealdb_protocol::proto::rpc::v1::surreal_db_service_server::SERVICE_NAME;
use super::{COMPRESSED_ENCODINGS, SERVICE_ROUTE, accepted_message_encodings};
#[test]
fn the_advertised_codecs_are_the_enabled_ones() {
let advertised = accepted_message_encodings();
let (named, identity) = advertised.split_at(advertised.len() - 1);
assert_eq!(identity, ["identity"], "identity is always accepted and is named last");
let enabled: Vec<&str> = COMPRESSED_ENCODINGS.iter().map(|(token, _)| *token).collect();
assert_eq!(named, enabled.as_slice());
}
#[test]
fn the_route_matches_the_generated_service_name() {
assert_eq!(SERVICE_ROUTE, format!("/{SERVICE_NAME}/{{*method}}"));
}
}