//! 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()); } }