#![allow(clippy::arithmetic_side_effects, clippy::expect_used, clippy::indexing_slicing, clippy::unwrap_used)]
use super::client::{GrpcClient, GrpcClientConfig};
use super::proto::acorn_rpc_server::AcornRpc;
use super::proto::{invoke_response, request_id, BatchRequest, InvokeRequest, NotifyRequest, RequestId};
use super::server::{validate_transport, GrpcServer, GrpcService};
use super::{json_value, protobuf_value};
use crate::io::api::json_rpc::{InvocationContext, MethodName, OperationDefinition, OperationEffects, OperationRegistry, RpcError, INVALID_PARAMS};
use crate::io::api::Secret;
use alloc::sync::Arc;
use core::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use core::time::Duration;
use futures::channel::oneshot;
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use std::net::TcpListener;
use tonic::{Code, Request};
#[derive(Deserialize, JsonSchema)]
#[serde(deny_unknown_fields)]
struct AddInput {
left: i64,
right: i64,
}
#[derive(JsonSchema, Serialize)]
#[serde(deny_unknown_fields)]
struct AddOutput {
total: i64,
}
struct CancellationGuard(Arc<AtomicBool>);
impl Drop for CancellationGuard {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
fn add_request(identifier: i64, left: i64, right: i64) -> InvokeRequest {
InvokeRequest {
id: Some(RequestId {
value: Some(request_id::Value::Number(identifier)),
}),
method: MethodName::from(["add"]).to_string(),
params: Some(protobuf_value(json!({"left": left, "right": right})).unwrap()),
}
}
fn authenticated<T>(message: T, token: &str) -> Request<T> {
let mut request = Request::new(message);
request.metadata_mut().insert("authorization", format!("Bearer {token}").parse().unwrap());
request
}
fn registry() -> OperationRegistry {
let definition = OperationDefinition::new::<AddInput, AddOutput>(MethodName::from(["add"]), OperationEffects::default()).unwrap();
OperationRegistry::default()
.register(definition, |input: AddInput, _| async move {
Ok(AddOutput {
total: input.left + input.right,
})
})
.unwrap()
}
fn response_value(response: super::proto::InvokeResponse) -> Result<Value, RpcError> {
match response.outcome {
| Some(invoke_response::Outcome::Error(error)) => Err(error.into()),
| Some(invoke_response::Outcome::Result(value)) => json_value(value),
| None => Err(RpcError::new(crate::io::api::json_rpc::INTERNAL_ERROR, "Missing gRPC outcome")),
}
}
fn slow_registry(cancelled: Arc<AtomicBool>) -> OperationRegistry {
let definition = OperationDefinition::new::<AddInput, AddOutput>(MethodName::from(["slow"]), OperationEffects::default()).unwrap();
registry()
.register(definition, move |input: AddInput, _| {
let cancelled = Arc::clone(&cancelled);
async move {
let _guard = CancellationGuard(cancelled);
tokio::time::sleep(Duration::from_secs(1)).await;
Ok(AddOutput {
total: input.left + input.right,
})
}
})
.unwrap()
}
#[tokio::test]
async fn authentication_and_batch_limits_use_transport_statuses() {
let service = GrpcService::new(registry(), Secret::from("secret".to_string()), InvocationContext::default())
.unwrap()
.with_max_batch_length(1);
let unauthenticated = service.invoke(Request::new(add_request(1, 1, 2))).await.unwrap_err();
assert_eq!(unauthenticated.code(), Code::Unauthenticated);
let oversized = service
.batch(authenticated(
BatchRequest {
requests: vec![add_request(1, 1, 2), add_request(2, 3, 4)],
},
"secret",
))
.await
.unwrap_err();
assert_eq!(oversized.code(), Code::ResourceExhausted);
}
#[tokio::test]
async fn batch_preserves_request_order() {
let service = GrpcService::new(registry(), Secret::from("secret".to_string()), InvocationContext::default()).unwrap();
let response = service
.batch(authenticated(
BatchRequest {
requests: vec![add_request(1, 1, 2), add_request(2, 10, 20)],
},
"secret",
))
.await
.unwrap()
.into_inner();
assert_eq!(response.responses.len(), 2);
assert_eq!(response_value(response.responses[0].clone()).unwrap(), json!({"total": 3}));
assert_eq!(response_value(response.responses[1].clone()).unwrap(), json!({"total": 30}));
}
#[test]
fn client_and_server_transport_policies_are_enforced() {
let token = Secret::from("secret".to_string());
assert!(GrpcClientConfig::new("http://127.0.0.1:50051", token.clone())
.with_offline(true)
.endpoint()
.is_ok());
assert!(GrpcClientConfig::new("http://example.com:50051", token.clone()).endpoint().is_err());
assert!(GrpcClientConfig::new("https://example.com:50051", token)
.with_offline(true)
.endpoint()
.is_err());
assert!(validate_transport("127.0.0.1:50051".parse().unwrap(), None).is_ok());
assert!(validate_transport("0.0.0.0:50051".parse().unwrap(), None).is_err());
}
#[tokio::test]
async fn grpc_and_json_rpc_produce_equivalent_results_and_errors() {
let registry = registry();
let service = GrpcService::new(registry.clone(), Secret::from("secret".to_string()), InvocationContext::default()).unwrap();
let direct = registry
.invoke(&MethodName::from(["add"]), json!({"left": 2, "right": 5}), InvocationContext::default())
.await
.unwrap();
let json_request = serde_json::to_vec(&json!({
"jsonrpc": "2.0",
"id": 1,
"method": MethodName::from(["add"]),
"params": {"left": 2, "right": 5}
}))
.unwrap();
let dispatched = registry.dispatch(&json_request, InvocationContext::default()).await;
let json_response: Value = serde_json::from_slice(dispatched.body.as_deref().unwrap()).unwrap();
let grpc_response = service.invoke(authenticated(add_request(1, 2, 5), "secret")).await.unwrap().into_inner();
assert_eq!(response_value(grpc_response).unwrap(), direct);
assert_eq!(json_response["result"], direct);
let invalid = InvokeRequest {
params: Some(protobuf_value(json!({"left": "bad", "right": 5})).unwrap()),
..add_request(2, 0, 0)
};
let error = response_value(service.invoke(authenticated(invalid, "secret")).await.unwrap().into_inner()).unwrap_err();
assert_eq!(error.code, INVALID_PARAMS);
}
#[tokio::test]
async fn network_client_server_enforce_auth_deadlines_and_cancellation() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let address = listener.local_addr().unwrap();
drop(listener);
let cancelled = Arc::new(AtomicBool::new(false));
let server = GrpcServer::new(
slow_registry(Arc::clone(&cancelled)),
Secret::from("secret".to_string()),
InvocationContext::default(),
)
.unwrap()
.with_request_timeout(Duration::from_millis(50));
let (shutdown_tx, shutdown_rx) = oneshot::channel();
let server_task = tokio::spawn(server.serve_with_shutdown(address, async move {
let _ = shutdown_rx.await;
}));
tokio::time::sleep(Duration::from_millis(30)).await;
let endpoint = format!("http://{address}");
let mut client =
GrpcClient::connect(GrpcClientConfig::new(endpoint.clone(), Secret::from("secret".to_string())).with_request_timeout(Duration::from_secs(1)))
.await
.unwrap();
let result = client.invoke(add_request(1, 4, 6)).await.unwrap();
assert_eq!(response_value(result).unwrap(), json!({"total": 10}));
let mut unauthorized =
GrpcClient::connect(GrpcClientConfig::new(endpoint, Secret::from("wrong".to_string())).with_request_timeout(Duration::from_secs(1)))
.await
.unwrap();
let status = unauthorized.invoke(add_request(2, 1, 1)).await.unwrap_err();
assert_eq!(status.code(), Code::Unauthenticated);
let slow = InvokeRequest {
method: MethodName::from(["slow"]).to_string(),
..add_request(3, 1, 2)
};
let status = client.invoke(slow).await.unwrap_err();
assert_eq!(status.code(), Code::Cancelled);
tokio::time::sleep(Duration::from_millis(20)).await;
assert!(cancelled.load(Ordering::SeqCst));
shutdown_tx.send(()).unwrap();
server_task.await.unwrap().unwrap();
}
#[tokio::test]
async fn notification_invokes_without_an_operation_response() {
let invocations = Arc::new(AtomicUsize::new(0));
let definition = OperationDefinition::new::<AddInput, AddOutput>(MethodName::from(["notify"]), OperationEffects::default()).unwrap();
let registry = OperationRegistry::default()
.register(definition, {
let invocations = Arc::clone(&invocations);
move |input: AddInput, _| {
let invocations = Arc::clone(&invocations);
async move {
invocations.fetch_add(1, Ordering::SeqCst);
Ok(AddOutput {
total: input.left + input.right,
})
}
}
})
.unwrap();
let service = GrpcService::new(registry, Secret::from("secret".to_string()), InvocationContext::default()).unwrap();
let notification = NotifyRequest {
method: MethodName::from(["notify"]).to_string(),
params: Some(protobuf_value(json!({"left": 1, "right": 2})).unwrap()),
};
service.notify(authenticated(notification, "secret")).await.unwrap();
assert_eq!(invocations.load(Ordering::SeqCst), 1);
}
#[test]
fn protobuf_values_round_trip_json_shapes() {
let value = json!({
"array": [null, true, 1.5, "value"],
"object": {"nested": false}
});
assert_eq!(json_value(protobuf_value(value.clone()).unwrap()).unwrap(), value);
assert!(protobuf_value(json!(9_007_199_254_740_992_u64)).is_err());
}
#[tokio::test]
async fn transport_failures_and_empty_credentials_are_reported_by_client() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let address = listener.local_addr().unwrap();
drop(listener);
let endpoint = format!("http://{address}");
let empty = GrpcClient::connect(GrpcClientConfig::new(endpoint.clone(), Secret::from(String::new()))).await;
assert!(empty.is_err());
let unavailable =
GrpcClient::connect(GrpcClientConfig::new(endpoint, Secret::from("secret".to_string())).with_request_timeout(Duration::from_millis(50)))
.await;
assert!(unavailable.is_err());
}