use super::*;
use crate::core::{ApiError, HandlerArgs, HandlerFn, HandlerState, extract_value};
use crate::grpc::handler::GrpcHandlerRegistration;
#[cfg(feature = "grpc")]
use std::collections::HashMap;
#[cfg(feature = "grpc")]
use std::sync::OnceLock;
#[cfg(feature = "grpc")]
use tonic::{Request, Response, Status, transport::Server};
#[cfg(feature = "grpc")]
use sdforge_v1::{
CallRequest, CallResponse, InfoRequest, InfoResponse,
sd_forge_service_server::{SdForgeService, SdForgeServiceServer},
};
#[cfg(feature = "grpc")]
const MAX_GRPC_ARGUMENTS_SIZE_BYTES: usize = 0x10_0000;
#[cfg(feature = "grpc")]
#[derive(Clone)]
pub struct SdForgeGrpcService {
state: HandlerState,
handlers: OnceLock<HashMap<&'static str, HandlerFn>>,
body_params: OnceLock<HashMap<&'static str, Option<&'static str>>>,
#[cfg(feature = "ratelimit")]
rate_limiter: Option<std::sync::Arc<dyn crate::security::ratelimit::RateLimiter>>,
}
#[cfg(feature = "grpc")]
impl Default for SdForgeGrpcService {
fn default() -> Self {
Self {
state: None,
handlers: OnceLock::new(),
body_params: OnceLock::new(),
#[cfg(feature = "ratelimit")]
rate_limiter: None,
}
}
}
#[cfg(feature = "grpc")]
impl SdForgeGrpcService {
#[must_use]
pub fn with_state(state: HandlerState) -> Self {
Self {
state,
handlers: OnceLock::new(),
body_params: OnceLock::new(),
#[cfg(feature = "ratelimit")]
rate_limiter: None,
}
}
#[cfg(feature = "ratelimit")]
#[must_use]
pub fn with_state_and_rate_limiter(
state: HandlerState,
rate_limiter: Option<std::sync::Arc<dyn crate::security::ratelimit::RateLimiter>>,
) -> Self {
Self {
state,
handlers: OnceLock::new(),
body_params: OnceLock::new(),
rate_limiter,
}
}
#[must_use]
fn handlers(&self) -> &HashMap<&'static str, HandlerFn> {
self.handlers.get_or_init(|| {
inventory::iter::<GrpcHandlerRegistration>()
.map(|r| (r.method, r.handler))
.collect()
})
}
#[must_use]
fn body_params(&self) -> &HashMap<&'static str, Option<&'static str>> {
self.body_params.get_or_init(|| {
inventory::iter::<GrpcHandlerRegistration>()
.map(|r| (r.method, r.body_param))
.collect()
})
}
}
#[cfg(feature = "grpc")]
#[tonic::async_trait]
impl SdForgeService for SdForgeGrpcService {
async fn call(&self, request: Request<CallRequest>) -> Result<Response<CallResponse>, Status> {
#[cfg(feature = "ratelimit")]
if let Some(ref limiter) = self.rate_limiter {
let identifier = request
.remote_addr()
.map(|addr| addr.ip().to_string())
.unwrap_or_else(|| "unknown".to_string());
if let Err(e) = limiter.check(&identifier).await {
use crate::security::ratelimit::RateLimitError;
let msg = match e {
RateLimitError::Exceeded {
limit,
window_seconds,
} => {
format!(
"rate limit exceeded: {} per {}s (client: {})",
limit, window_seconds, identifier
)
}
RateLimitError::Banned { reason } => {
format!("client banned: {} (client: {})", reason, identifier)
}
RateLimitError::CircuitOpen => {
format!("circuit breaker open (client: {})", identifier)
}
RateLimitError::QuotaExhausted { used, total } => {
format!(
"quota exhausted: {}/{} (client: {})",
used, total, identifier
)
}
RateLimitError::Limiteron(e) => {
format!("rate limiter error: {} (client: {})", e, identifier)
}
};
return Err(Status::resource_exhausted(msg));
}
}
let req = request.into_inner();
let payload_size = req.parameters.values().map(|v| v.len()).sum::<usize>() + req.data.len();
if payload_size > MAX_GRPC_ARGUMENTS_SIZE_BYTES {
return Err(Status::invalid_argument(format!(
"arguments payload size ({}) exceeds maximum allowed size ({})",
payload_size, MAX_GRPC_ARGUMENTS_SIZE_BYTES
)));
}
let handler = self.handlers().get(req.method.as_str()).copied().ok_or_else(
|| {
Status::not_found(format!(
"method '{}' not registered (no matching #[forge(grpc_method = \"...\")] declaration)",
req.method
))
},
)?;
let mut args: HandlerArgs = req.parameters.into_iter().collect();
if !req.data.is_empty() {
match self
.body_params()
.get(req.method.as_str())
.copied()
.flatten()
{
Some(bp) => {
args.insert(bp.to_string(), req.data);
}
None => {
return Err(Status::invalid_argument(format!(
"method '{}' has no body parameter but CallRequest.data is non-empty",
req.method
)));
}
}
}
use futures_util::FutureExt;
use std::panic::AssertUnwindSafe;
let outcome = AssertUnwindSafe(handler(args, self.state.clone()))
.catch_unwind()
.await;
match outcome {
Ok(Ok(value)) => {
Ok(Response::new(CallResponse {
success: true,
data: extract_value(&value),
error: String::new(),
status_code: 200,
}))
}
Ok(Err(e)) => {
let status_code = map_error_to_http(&e);
Ok(Response::new(CallResponse {
success: false,
data: String::new(),
error: e.to_string(),
status_code,
}))
}
Err(_panic) => {
Err(Status::internal("handler panicked"))
}
}
}
async fn get_info(
&self,
_request: Request<InfoRequest>,
) -> Result<Response<InfoResponse>, Status> {
let response = InfoResponse {
name: "SdForge Service".to_string(),
version: "0.1.0".to_string(),
methods: self
.handlers()
.keys()
.map(|k| (*k).to_string())
.collect::<Vec<_>>(),
description: "SdForge Multi-Protocol SDK Framework".to_string(),
};
Ok(Response::new(response))
}
}
#[cfg(feature = "grpc")]
fn map_error_to_http(e: &ApiError) -> i32 {
match e {
ApiError::NotFound { .. } => 404,
ApiError::InvalidInput { .. } | ApiError::ValidationError { .. } => 422,
ApiError::AuthenticationFailed { .. } => 401,
ApiError::AccessDenied { .. } => 403,
ApiError::RateLimitExceeded { .. } => 429,
ApiError::ServiceUnavailable { .. } => 503,
ApiError::Internal { .. } => 500,
}
}
#[cfg(feature = "grpc")]
impl GrpcRoute {
#[allow(missing_docs)]
pub fn new(service_name: String, metadata: ApiMetadata) -> Self {
Self {
service_name,
metadata,
}
}
#[cfg(test)]
pub(crate) fn service_name(&self) -> &str {
&self.service_name
}
#[cfg(test)]
pub(crate) fn metadata(&self) -> &ApiMetadata {
&self.metadata
}
}
#[cfg(feature = "grpc")]
#[deprecated(
note = "use build_server_with_config with auth configured; build_server starts an unauthenticated server"
)]
pub async fn build_server(addr: &str) -> Result<(), Box<dyn std::error::Error>> {
let addr = match addr.parse::<std::net::SocketAddr>() {
Ok(addr) => addr,
Err(e) => {
return Err(Box::new(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!("Invalid gRPC server address format: {}", e),
)));
}
};
let service = SdForgeGrpcService::default();
Server::builder()
.add_service(SdForgeServiceServer::new(service).max_decoding_message_size(4 * 1024 * 1024))
.serve(addr)
.await?;
Ok(())
}
#[cfg(feature = "grpc")]
pub async fn build_server_with_config(
addr: &str,
config: GrpcServerConfig,
) -> Result<(), Box<dyn std::error::Error>> {
let addr = match addr.parse::<std::net::SocketAddr>() {
Ok(addr) => addr,
Err(e) => {
return Err(Box::new(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!("Invalid gRPC server address format: {}", e),
)));
}
};
#[cfg(feature = "security")]
if config.require_auth && config.auth.is_none() {
return Err(Box::new(std::io::Error::new(
std::io::ErrorKind::PermissionDenied,
"gRPC server requires authentication but no auth configured \
(set GrpcServerConfig.require_auth = false to override)",
)));
}
#[cfg(feature = "ratelimit")]
let service =
SdForgeGrpcService::with_state_and_rate_limiter(config.state, config.rate_limiter);
#[cfg(not(feature = "ratelimit"))]
let service = SdForgeGrpcService::with_state(config.state);
#[cfg(feature = "security")]
let mut builder = {
let auth_interceptor = make_auth_interceptor(config.auth.clone());
Server::builder().layer(tonic::service::InterceptorLayer::new(auth_interceptor))
};
#[cfg(not(feature = "security"))]
let mut builder = { Server::builder() };
if config.max_connections > 0 {
builder = builder.concurrency_limit_per_connection(config.max_connections);
}
if config.timeout_seconds > 0 {
builder = builder.timeout(std::time::Duration::from_secs(config.timeout_seconds));
}
builder
.add_service(SdForgeServiceServer::new(service).max_decoding_message_size(4 * 1024 * 1024))
.serve(addr)
.await?;
Ok(())
}
#[cfg(feature = "grpc")]
impl Default for GrpcServerConfig {
fn default() -> Self {
Self {
max_connections: 1000,
timeout_seconds: 30,
require_auth: true, #[cfg(feature = "security")]
auth: None,
state: None,
#[cfg(feature = "ratelimit")]
rate_limiter: None,
}
}
}
#[cfg(all(feature = "grpc", feature = "security"))]
pub(crate) fn make_auth_interceptor(
auth: Option<crate::security::BearerAuth>,
) -> AuthGrpcInterceptor {
AuthGrpcInterceptor { auth }
}
#[cfg(all(feature = "grpc", feature = "security"))]
impl tonic::service::Interceptor for AuthGrpcInterceptor {
fn call(&mut self, req: tonic::Request<()>) -> Result<tonic::Request<()>, Status> {
let Some(ref bearer_auth) = self.auth else {
return Ok(req);
};
let token = req
.metadata()
.get("authorization")
.and_then(|v| v.to_str().ok())
.and_then(|h| h.strip_prefix("Bearer "))
.map(String::from);
match token {
Some(token_str) => {
if bearer_auth.validate_token(&token_str).is_some() {
Ok(req)
} else {
Err(Status::unauthenticated("Invalid or expired token"))
}
}
None => Err(Status::unauthenticated("Missing authorization header")),
}
}
}
#[cfg(feature = "grpc")]
impl SdForgeGrpcService {
#[cfg(test)]
pub(crate) fn body_params_map(&self) -> &HashMap<&'static str, Option<&'static str>> {
self.body_params()
}
}
#[cfg(all(test, feature = "grpc"))]
mod tests {
use super::*;
use crate::core::HandlerArgs;
use serde_json::Value;
use std::sync::Arc;
use tonic::Request;
fn echo_handler(args: HandlerArgs, _state: HandlerState) -> crate::core::HandlerFuture {
let msg = args.get("msg").cloned().unwrap_or_default();
Box::pin(async move { Ok(Value::String(msg)) })
}
inventory::submit! {
GrpcHandlerRegistration {
method: "test_echo",
handler: echo_handler,
body_param: None,
}
}
fn not_found_handler(_args: HandlerArgs, _state: HandlerState) -> crate::core::HandlerFuture {
Box::pin(async {
Err(ApiError::NotFound {
resource: "test_resource".to_string(),
resource_id: Some("123".to_string()),
})
})
}
inventory::submit! {
GrpcHandlerRegistration {
method: "test_not_found",
handler: not_found_handler,
body_param: None,
}
}
fn panic_handler(_args: HandlerArgs, _state: HandlerState) -> crate::core::HandlerFuture {
Box::pin(async {
panic!("boom — must not leak to client");
})
}
inventory::submit! {
GrpcHandlerRegistration {
method: "test_panic",
handler: panic_handler,
body_param: None,
}
}
fn body_handler(args: HandlerArgs, _state: HandlerState) -> crate::core::HandlerFuture {
let payload = args.get("payload").cloned().unwrap_or_default();
Box::pin(async move { Ok(Value::String(payload)) })
}
inventory::submit! {
GrpcHandlerRegistration {
method: "test_body",
handler: body_handler,
body_param: Some("payload"),
}
}
#[test]
fn lookup_builds_cache_from_inventory() {
let service = SdForgeGrpcService::default();
let table = service.handlers();
assert!(table.contains_key("test_echo"));
assert!(table.contains_key("test_not_found"));
assert!(table.contains_key("test_panic"));
assert!(table.contains_key("test_body"));
}
#[test]
fn lookup_cache_is_idempotent() {
let service = SdForgeGrpcService::default();
let first = service.handlers();
let second = service.handlers();
assert!(std::ptr::eq(first, second));
}
#[test]
fn body_params_cache_built_correctly() {
let service = SdForgeGrpcService::default();
let map = service.body_params_map();
assert_eq!(map.get("test_echo"), Some(&None));
assert_eq!(map.get("test_body"), Some(&Some("payload")));
}
#[tokio::test]
async fn call_routes_to_registered_handler() {
let service = SdForgeGrpcService::default();
let mut params = HashMap::new();
params.insert("msg".to_string(), "hello world".to_string());
let req = Request::new(CallRequest {
method: "test_echo".to_string(),
parameters: params,
data: String::new(),
});
let resp = service.call(req).await.unwrap().into_inner();
assert!(resp.success);
assert_eq!(resp.data, "hello world");
assert_eq!(resp.status_code, 200);
assert!(resp.error.is_empty());
}
#[tokio::test]
async fn call_unknown_method_returns_not_found() {
let service = SdForgeGrpcService::default();
let req = Request::new(CallRequest {
method: "no_such_method".to_string(),
parameters: HashMap::new(),
data: String::new(),
});
let err = service.call(req).await.unwrap_err();
assert_eq!(err.code(), tonic::Code::NotFound);
assert!(err.message().contains("no_such_method"));
}
#[tokio::test]
async fn call_data_without_body_param_returns_invalid_argument() {
let service = SdForgeGrpcService::default();
let req = Request::new(CallRequest {
method: "test_echo".to_string(),
parameters: HashMap::new(),
data: "unexpected payload".to_string(),
});
let err = service.call(req).await.unwrap_err();
assert_eq!(err.code(), tonic::Code::InvalidArgument);
}
#[tokio::test]
async fn call_routes_data_to_body_param() {
let service = SdForgeGrpcService::default();
let req = Request::new(CallRequest {
method: "test_body".to_string(),
parameters: HashMap::new(),
data: "{\"x\":1}".to_string(),
});
let resp = service.call(req).await.unwrap().into_inner();
assert!(resp.success);
assert_eq!(resp.data, "{\"x\":1}");
}
#[tokio::test]
async fn call_business_error_returns_success_false_with_status_code() {
let service = SdForgeGrpcService::default();
let req = Request::new(CallRequest {
method: "test_not_found".to_string(),
parameters: HashMap::new(),
data: String::new(),
});
let resp = service.call(req).await.unwrap().into_inner();
assert!(!resp.success);
assert_eq!(resp.status_code, 404);
assert!(resp.error.contains("test_resource"));
assert!(resp.data.is_empty());
}
#[tokio::test]
async fn call_panic_handler_returns_status_internal() {
let service = SdForgeGrpcService::default();
let req = Request::new(CallRequest {
method: "test_panic".to_string(),
parameters: HashMap::new(),
data: String::new(),
});
let err = service.call(req).await.unwrap_err();
assert_eq!(err.code(), tonic::Code::Internal);
assert!(!err.message().contains("boom"));
assert!(!err.message().contains("leak"));
}
#[test]
fn map_error_to_http_covers_all_variants() {
let cases: Vec<(ApiError, i32)> = vec![
(
ApiError::NotFound {
resource: "x".into(),
resource_id: None,
},
404,
),
(
ApiError::InvalidInput {
message: "x".into(),
field: None,
value: None,
},
422,
),
(
ApiError::ValidationError {
field: "x".into(),
constraint: "required".into(),
},
422,
),
(ApiError::AuthenticationFailed { reason: "x".into() }, 401),
(
ApiError::AccessDenied {
permission: "x".into(),
user_id: None,
},
403,
),
(
ApiError::RateLimitExceeded {
limit: 10,
window_seconds: 60,
},
429,
),
(
ApiError::ServiceUnavailable {
service: "x".into(),
retry_after: None,
source: None,
},
503,
),
(
ApiError::Internal {
message: "x".into(),
error_id: "x".into(),
source: None,
context: None,
},
500,
),
];
for (e, expected) in cases {
assert_eq!(map_error_to_http(&e), expected, "mismatch for {e:?}");
}
}
#[test]
fn default_state_is_none() {
let config = GrpcServerConfig::default();
assert!(config.state.is_none());
}
#[tokio::test]
async fn state_injected_to_service_can_be_downcast() {
use std::any::Any;
let state: Arc<dyn Any + Send + Sync> = Arc::new(42_i32);
let service = SdForgeGrpcService::with_state(Some(state));
let borrowed = service.state.clone().unwrap();
let downcast = borrowed.downcast_ref::<i32>();
assert_eq!(downcast, Some(&42_i32));
}
#[cfg(feature = "ratelimit")]
mod vuln_0006_ratelimit_tests {
use super::*;
use crate::security::ratelimit::{RateLimitError, RateLimiter};
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
struct AlwaysAllowLimiter;
impl RateLimiter for AlwaysAllowLimiter {
fn check<'a>(
&'a self,
_identifier: &'a str,
) -> Pin<Box<dyn Future<Output = Result<(), RateLimitError>> + Send + 'a>> {
Box::pin(async { Ok(()) })
}
}
struct AlwaysRejectLimiter;
impl RateLimiter for AlwaysRejectLimiter {
fn check<'a>(
&'a self,
_identifier: &'a str,
) -> Pin<Box<dyn Future<Output = Result<(), RateLimitError>> + Send + 'a>> {
Box::pin(async {
Err(RateLimitError::Exceeded {
limit: 10,
window_seconds: 60,
})
})
}
}
struct CountingLimiter {
count: AtomicU32,
}
impl RateLimiter for CountingLimiter {
fn check<'a>(
&'a self,
_identifier: &'a str,
) -> Pin<Box<dyn Future<Output = Result<(), RateLimitError>> + Send + 'a>> {
self.count.fetch_add(1, Ordering::SeqCst);
Box::pin(async { Ok(()) })
}
}
#[tokio::test]
async fn call_without_rate_limiter_proceeds_normally() {
let service = SdForgeGrpcService::default();
let mut params = HashMap::new();
params.insert("msg".to_string(), "hello".to_string());
let req = Request::new(CallRequest {
method: "test_echo".to_string(),
parameters: params,
data: String::new(),
});
let resp = service.call(req).await.unwrap().into_inner();
assert!(resp.success);
assert_eq!(resp.data, "hello");
}
#[tokio::test]
async fn call_with_allowing_limiter_proceeds_normally() {
let limiter: Arc<dyn RateLimiter> = Arc::new(AlwaysAllowLimiter);
let service = SdForgeGrpcService::with_state_and_rate_limiter(None, Some(limiter));
let mut params = HashMap::new();
params.insert("msg".to_string(), "allowed".to_string());
let req = Request::new(CallRequest {
method: "test_echo".to_string(),
parameters: params,
data: String::new(),
});
let resp = service.call(req).await.unwrap().into_inner();
assert!(resp.success);
assert_eq!(resp.data, "allowed");
}
#[tokio::test]
async fn call_with_rejecting_limiter_returns_resource_exhausted() {
let limiter: Arc<dyn RateLimiter> = Arc::new(AlwaysRejectLimiter);
let service = SdForgeGrpcService::with_state_and_rate_limiter(None, Some(limiter));
let req = Request::new(CallRequest {
method: "test_echo".to_string(),
parameters: HashMap::new(),
data: String::new(),
});
let err = service.call(req).await.unwrap_err();
assert_eq!(
err.code(),
tonic::Code::ResourceExhausted,
"vuln-0006: rejected rate limit must return ResourceExhausted, got {:?}",
err.code()
);
assert!(
err.message().contains("rate limit exceeded"),
"error message should mention rate limit, got: {}",
err.message()
);
}
#[tokio::test]
async fn call_with_rate_limiter_invokes_check_once_per_call() {
let limiter = Arc::new(CountingLimiter {
count: AtomicU32::new(0),
});
let count_clone = Arc::clone(&limiter);
let limiter_dyn: Arc<dyn RateLimiter> = limiter as Arc<dyn RateLimiter>;
let service = SdForgeGrpcService::with_state_and_rate_limiter(None, Some(limiter_dyn));
let req = Request::new(CallRequest {
method: "test_echo".to_string(),
parameters: HashMap::new(),
data: String::new(),
});
let _ = service.call(req).await;
assert_eq!(
count_clone.count.load(Ordering::SeqCst),
1,
"vuln-0006: rate limiter check must be called exactly once per gRPC call"
);
}
#[tokio::test]
async fn call_with_banned_limiter_returns_resource_exhausted() {
struct AlwaysBannedLimiter;
impl RateLimiter for AlwaysBannedLimiter {
fn check<'a>(
&'a self,
_identifier: &'a str,
) -> Pin<Box<dyn Future<Output = Result<(), RateLimitError>> + Send + 'a>>
{
Box::pin(async {
Err(RateLimitError::Banned {
reason: "abuse detected".to_string(),
})
})
}
}
let limiter: Arc<dyn RateLimiter> = Arc::new(AlwaysBannedLimiter);
let service = SdForgeGrpcService::with_state_and_rate_limiter(None, Some(limiter));
let req = Request::new(CallRequest {
method: "test_echo".to_string(),
parameters: HashMap::new(),
data: String::new(),
});
let err = service.call(req).await.unwrap_err();
assert_eq!(err.code(), tonic::Code::ResourceExhausted);
assert!(
err.message().contains("banned") && err.message().contains("abuse detected"),
"error should mention ban reason, got: {}",
err.message()
);
}
#[test]
fn default_config_has_no_rate_limiter() {
let config = GrpcServerConfig::default();
assert!(
config.rate_limiter.is_none(),
"default config should not enable rate limiting"
);
}
}
}