Files
rustytorch/crates/production/rtx-serving-api/tests/grpc_tests.rs
T
2026-03-04 00:08:42 +00:00

212 lines
6.2 KiB
Rust

//! gRPC service tests
#[cfg(test)]
mod tests {
use rtx_serving_api::grpc::{
GrpcConfig, GrpcServer, HealthCheckRequest, HealthCheckResponse, InferenceRequest,
InferenceResponse, ModelInfoRequest, ModelInfoResponse,
inference_service_server::InferenceService,
};
use tonic::{Request, Response, Status};
#[tokio::test]
async fn test_grpc_server_creation() {
let config = GrpcConfig::default();
let server = GrpcServer::new(config);
assert_eq!(server.port(), 50051);
assert!(server.is_ready());
}
#[tokio::test]
async fn test_grpc_health_check() {
let config = GrpcConfig::default();
let server = GrpcServer::new(config);
let request = Request::new(HealthCheckRequest {
service: "inference".to_string(),
});
let response = server.health_check(request).await.unwrap();
let inner = response.into_inner();
assert_eq!(inner.status, "healthy");
assert!(inner.uptime_seconds > 0);
}
#[tokio::test]
async fn test_grpc_model_info() {
let config = GrpcConfig::default();
let mut server = GrpcServer::new(config);
// Register a test model
server.register_model("gpt2", "transformer", 1024);
let request = Request::new(ModelInfoRequest {
model_name: "gpt2".to_string(),
});
let response = server.model_info(request).await.unwrap();
let info = response.into_inner();
assert_eq!(info.name, "gpt2");
assert_eq!(info.model_type, "transformer");
assert_eq!(info.max_batch_size, 1024);
assert!(info.is_loaded);
}
#[tokio::test]
async fn test_grpc_inference_request() {
let config = GrpcConfig::default();
let mut server = GrpcServer::new(config);
server.register_model("test_model", "classifier", 32);
let request = Request::new(InferenceRequest {
model_name: "test_model".to_string(),
input_data: vec![1.0, 2.0, 3.0, 4.0],
batch_size: 1,
options: Default::default(),
});
let response = server.predict(request).await.unwrap();
let result = response.into_inner();
assert_eq!(result.model_name, "test_model");
assert!(!result.outputs.is_empty());
assert!(result.latency_ms > 0.0);
}
#[tokio::test]
async fn test_grpc_streaming_inference() {
let config = GrpcConfig::default();
let mut server = GrpcServer::new(config);
server.register_model("streaming_model", "generator", 128);
let request = Request::new(InferenceRequest {
model_name: "streaming_model".to_string(),
input_data: vec![1.0, 2.0],
batch_size: 1,
options: Default::default(),
});
let mut stream = server.predict_stream(request).await.unwrap().into_inner();
use futures::StreamExt;
let mut token_count = 0;
while let Some(result) = stream.next().await {
let result = result.unwrap();
assert_eq!(result.model_name, "streaming_model");
assert!(!result.outputs.is_empty());
token_count += 1;
}
assert!(token_count > 0);
}
#[tokio::test]
async fn test_grpc_batch_inference() {
let config = GrpcConfig::default();
let mut server = GrpcServer::new(config);
server.register_model("batch_model", "classifier", 64);
let batch_size = 4;
let input_size = 10;
let mut input_data = Vec::new();
for i in 0..batch_size * input_size {
input_data.push(i as f32);
}
let request = Request::new(InferenceRequest {
model_name: "batch_model".to_string(),
input_data,
batch_size: batch_size as i32,
options: Default::default(),
});
let response = server.predict(request).await.unwrap();
let result = response.into_inner();
assert_eq!(result.batch_size, batch_size as i32);
assert!(result.outputs.len() >= batch_size);
}
#[tokio::test]
async fn test_grpc_model_not_found() {
let config = GrpcConfig::default();
let server = GrpcServer::new(config);
let request = Request::new(InferenceRequest {
model_name: "nonexistent".to_string(),
input_data: vec![1.0],
batch_size: 1,
options: Default::default(),
});
let result = server.predict(request).await;
assert!(result.is_err());
if let Err(status) = result {
assert_eq!(status.code(), tonic::Code::NotFound);
}
}
#[tokio::test]
async fn test_grpc_concurrent_requests() {
let config = GrpcConfig::default();
let server = std::sync::Arc::new(tokio::sync::Mutex::new(GrpcServer::new(config)));
{
let mut srv = server.lock().await;
srv.register_model("concurrent_model", "classifier", 256);
}
let mut handles = vec![];
for i in 0..10 {
let server_clone = server.clone();
let handle = tokio::spawn(async move {
let request = Request::new(InferenceRequest {
model_name: "concurrent_model".to_string(),
input_data: vec![i as f32],
batch_size: 1,
options: Default::default(),
});
let srv = server_clone.lock().await;
srv.predict(request).await
});
handles.push(handle);
}
for handle in handles {
let result = handle.await.unwrap();
assert!(result.is_ok());
}
}
#[tokio::test]
async fn test_grpc_server_shutdown() {
let config = GrpcConfig::default();
let server = GrpcServer::new(config);
assert!(server.is_ready());
server.shutdown().await.unwrap();
assert!(!server.is_ready());
}
#[tokio::test]
async fn test_grpc_tls_configuration() {
let mut config = GrpcConfig::default();
config.enable_tls("cert.pem", "key.pem");
let server = GrpcServer::new(config);
assert!(server.is_tls_enabled());
}
}