212 lines
6.2 KiB
Rust
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());
|
|
}
|
|
}
|