[model-gateway] : Rust integration tests for integration_mock replacement (#16441)

This commit is contained in:
Simo Lin
2026-01-04 21:01:47 -08:00
committed by GitHub
parent 078270473a
commit f84487af59
48 changed files with 4856 additions and 57 deletions
@@ -23,6 +23,12 @@ pub struct MockSearchServer {
tool_router: ToolRouter<MockSearchServer>,
}
impl Default for MockSearchServer {
fn default() -> Self {
Self::new()
}
}
#[tool_router]
impl MockSearchServer {
pub fn new() -> Self {
+6
View File
@@ -6,7 +6,11 @@ pub mod mock_openai_server;
pub mod mock_worker;
pub mod streaming_helpers;
pub mod test_app;
pub mod test_certs;
pub mod test_config;
pub mod tls_mock_worker;
// Re-export commonly used test builders
use std::{
fs,
path::PathBuf,
@@ -33,6 +37,8 @@ use smg::{
tokenizer::registry::TokenizerRegistry,
tool_parser::ParserFactory as ToolParserFactory,
};
#[allow(unused_imports)]
pub use test_config::{TestRouterConfig, TestWorkerConfig};
/// Test context for directly testing mock workers without full router setup.
pub struct WorkerTestContext {
@@ -0,0 +1,299 @@
// TLS certificate generation for integration tests
#![allow(dead_code)]
use std::path::PathBuf;
use openssl::{
asn1::Asn1Time,
bn::{BigNum, MsbOption},
hash::MessageDigest,
pkey::{PKey, Private},
rsa::Rsa,
x509::{
extension::{BasicConstraints, ExtendedKeyUsage, KeyUsage, SubjectAlternativeName},
X509Builder, X509NameBuilder, X509,
},
};
use tempfile::TempDir;
/// Container for generated test certificates
pub struct TestCertificates {
/// Temporary directory containing the certificate files
pub temp_dir: TempDir,
/// Path to CA certificate
pub ca_cert_path: PathBuf,
/// Path to CA private key
pub ca_key_path: PathBuf,
/// Path to server certificate
pub server_cert_path: PathBuf,
/// Path to server private key
pub server_key_path: PathBuf,
/// Path to client certificate
pub client_cert_path: PathBuf,
/// Path to client private key
pub client_key_path: PathBuf,
}
impl TestCertificates {
/// Generate a complete set of test certificates for mTLS testing.
///
/// Creates:
/// - CA certificate and key (self-signed root CA)
/// - Server certificate and key (signed by CA, for localhost)
/// - Client certificate and key (signed by CA, for client authentication)
pub fn generate() -> Result<Self, Box<dyn std::error::Error>> {
let temp_dir = TempDir::new()?;
let base_path = temp_dir.path();
// Generate CA key pair
let ca_key = generate_rsa_key()?;
let ca_cert = generate_ca_certificate(&ca_key)?;
// Generate server key pair and certificate
let server_key = generate_rsa_key()?;
let server_cert = generate_server_certificate(&server_key, &ca_cert, &ca_key)?;
// Generate client key pair and certificate
let client_key = generate_rsa_key()?;
let client_cert = generate_client_certificate(&client_key, &ca_cert, &ca_key)?;
// Write all files
let ca_cert_path = base_path.join("ca_cert.pem");
let ca_key_path = base_path.join("ca_key.pem");
let server_cert_path = base_path.join("server_cert.pem");
let server_key_path = base_path.join("server_key.pem");
let client_cert_path = base_path.join("client_cert.pem");
let client_key_path = base_path.join("client_key.pem");
std::fs::write(&ca_cert_path, ca_cert.to_pem()?)?;
std::fs::write(&ca_key_path, ca_key.private_key_to_pem_pkcs8()?)?;
std::fs::write(&server_cert_path, server_cert.to_pem()?)?;
std::fs::write(&server_key_path, server_key.private_key_to_pem_pkcs8()?)?;
std::fs::write(&client_cert_path, client_cert.to_pem()?)?;
std::fs::write(&client_key_path, client_key.private_key_to_pem_pkcs8()?)?;
Ok(Self {
temp_dir,
ca_cert_path,
ca_key_path,
server_cert_path,
server_key_path,
client_cert_path,
client_key_path,
})
}
/// Get paths as string references for use with RouterConfig builder
pub fn ca_cert_str(&self) -> &str {
self.ca_cert_path.to_str().unwrap()
}
pub fn server_cert_str(&self) -> &str {
self.server_cert_path.to_str().unwrap()
}
pub fn server_key_str(&self) -> &str {
self.server_key_path.to_str().unwrap()
}
pub fn client_cert_str(&self) -> &str {
self.client_cert_path.to_str().unwrap()
}
pub fn client_key_str(&self) -> &str {
self.client_key_path.to_str().unwrap()
}
}
/// Generate a 2048-bit RSA key pair
fn generate_rsa_key() -> Result<PKey<Private>, Box<dyn std::error::Error>> {
let rsa = Rsa::generate(2048)?;
Ok(PKey::from_rsa(rsa)?)
}
/// Generate a self-signed CA certificate
fn generate_ca_certificate(key: &PKey<Private>) -> Result<X509, Box<dyn std::error::Error>> {
let mut name_builder = X509NameBuilder::new()?;
name_builder.append_entry_by_text("C", "US")?;
name_builder.append_entry_by_text("ST", "California")?;
name_builder.append_entry_by_text("L", "Test City")?;
name_builder.append_entry_by_text("O", "Test CA Organization")?;
name_builder.append_entry_by_text("CN", "Test CA")?;
let name = name_builder.build();
let mut cert_builder = X509Builder::new()?;
cert_builder.set_version(2)?; // X509 v3
// Serial number
let serial = {
let mut bn = BigNum::new()?;
bn.rand(128, MsbOption::MAYBE_ZERO, false)?;
bn.to_asn1_integer()?
};
cert_builder.set_serial_number(&serial)?;
cert_builder.set_subject_name(&name)?;
cert_builder.set_issuer_name(&name)?; // Self-signed
cert_builder.set_pubkey(key)?;
// Validity: 1 year from now
let not_before = Asn1Time::days_from_now(0)?;
let not_after = Asn1Time::days_from_now(365)?;
cert_builder.set_not_before(&not_before)?;
cert_builder.set_not_after(&not_after)?;
// Extensions for CA
let basic_constraints = BasicConstraints::new().critical().ca().build()?;
cert_builder.append_extension(basic_constraints)?;
let key_usage = KeyUsage::new()
.critical()
.key_cert_sign()
.crl_sign()
.build()?;
cert_builder.append_extension(key_usage)?;
cert_builder.sign(key, MessageDigest::sha256())?;
Ok(cert_builder.build())
}
/// Generate a server certificate signed by the CA
fn generate_server_certificate(
key: &PKey<Private>,
ca_cert: &X509,
ca_key: &PKey<Private>,
) -> Result<X509, Box<dyn std::error::Error>> {
let mut name_builder = X509NameBuilder::new()?;
name_builder.append_entry_by_text("C", "US")?;
name_builder.append_entry_by_text("ST", "California")?;
name_builder.append_entry_by_text("L", "Test City")?;
name_builder.append_entry_by_text("O", "Test Server Organization")?;
name_builder.append_entry_by_text("CN", "localhost")?;
let name = name_builder.build();
let mut cert_builder = X509Builder::new()?;
cert_builder.set_version(2)?;
let serial = {
let mut bn = BigNum::new()?;
bn.rand(128, MsbOption::MAYBE_ZERO, false)?;
bn.to_asn1_integer()?
};
cert_builder.set_serial_number(&serial)?;
cert_builder.set_subject_name(&name)?;
cert_builder.set_issuer_name(ca_cert.subject_name())?;
cert_builder.set_pubkey(key)?;
let not_before = Asn1Time::days_from_now(0)?;
let not_after = Asn1Time::days_from_now(365)?;
cert_builder.set_not_before(&not_before)?;
cert_builder.set_not_after(&not_after)?;
// Extensions for server certificate
let basic_constraints = BasicConstraints::new().build()?;
cert_builder.append_extension(basic_constraints)?;
let key_usage = KeyUsage::new()
.critical()
.digital_signature()
.key_encipherment()
.build()?;
cert_builder.append_extension(key_usage)?;
let ext_key_usage = ExtendedKeyUsage::new().server_auth().build()?;
cert_builder.append_extension(ext_key_usage)?;
// Subject Alternative Names for localhost
let san = SubjectAlternativeName::new()
.dns("localhost")
.ip("127.0.0.1")
.ip("::1")
.build(&cert_builder.x509v3_context(Some(ca_cert), None))?;
cert_builder.append_extension(san)?;
cert_builder.sign(ca_key, MessageDigest::sha256())?;
Ok(cert_builder.build())
}
/// Generate a client certificate signed by the CA
fn generate_client_certificate(
key: &PKey<Private>,
ca_cert: &X509,
ca_key: &PKey<Private>,
) -> Result<X509, Box<dyn std::error::Error>> {
let mut name_builder = X509NameBuilder::new()?;
name_builder.append_entry_by_text("C", "US")?;
name_builder.append_entry_by_text("ST", "California")?;
name_builder.append_entry_by_text("L", "Test City")?;
name_builder.append_entry_by_text("O", "Test Client Organization")?;
name_builder.append_entry_by_text("CN", "Test Client")?;
let name = name_builder.build();
let mut cert_builder = X509Builder::new()?;
cert_builder.set_version(2)?;
let serial = {
let mut bn = BigNum::new()?;
bn.rand(128, MsbOption::MAYBE_ZERO, false)?;
bn.to_asn1_integer()?
};
cert_builder.set_serial_number(&serial)?;
cert_builder.set_subject_name(&name)?;
cert_builder.set_issuer_name(ca_cert.subject_name())?;
cert_builder.set_pubkey(key)?;
let not_before = Asn1Time::days_from_now(0)?;
let not_after = Asn1Time::days_from_now(365)?;
cert_builder.set_not_before(&not_before)?;
cert_builder.set_not_after(&not_after)?;
// Extensions for client certificate
let basic_constraints = BasicConstraints::new().build()?;
cert_builder.append_extension(basic_constraints)?;
let key_usage = KeyUsage::new().critical().digital_signature().build()?;
cert_builder.append_extension(key_usage)?;
let ext_key_usage = ExtendedKeyUsage::new().client_auth().build()?;
cert_builder.append_extension(ext_key_usage)?;
cert_builder.sign(ca_key, MessageDigest::sha256())?;
Ok(cert_builder.build())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_certificate_generation() {
let certs = TestCertificates::generate().expect("Failed to generate certificates");
// Verify all files exist
assert!(certs.ca_cert_path.exists(), "CA cert should exist");
assert!(certs.ca_key_path.exists(), "CA key should exist");
assert!(certs.server_cert_path.exists(), "Server cert should exist");
assert!(certs.server_key_path.exists(), "Server key should exist");
assert!(certs.client_cert_path.exists(), "Client cert should exist");
assert!(certs.client_key_path.exists(), "Client key should exist");
// Verify files are not empty
assert!(
std::fs::metadata(&certs.ca_cert_path).unwrap().len() > 0,
"CA cert should not be empty"
);
assert!(
std::fs::metadata(&certs.server_cert_path).unwrap().len() > 0,
"Server cert should not be empty"
);
assert!(
std::fs::metadata(&certs.client_cert_path).unwrap().len() > 0,
"Client cert should not be empty"
);
}
}
@@ -0,0 +1,316 @@
//! Test configuration builders to reduce duplication across tests
//!
//! Provides pre-configured RouterConfig and MockWorkerConfig builders
//! for common test scenarios.
use smg::config::{CircuitBreakerConfig, PolicyConfig, RetryConfig, RouterConfig};
use super::mock_worker::{HealthStatus, MockWorkerConfig, WorkerType};
/// Default test configuration values
pub mod defaults {
pub const HOST: &str = "127.0.0.1";
pub const MAX_PAYLOAD_SIZE: usize = 256 * 1024 * 1024; // 256MB
pub const REQUEST_TIMEOUT_SECS: u64 = 600;
pub const WORKER_STARTUP_TIMEOUT_SECS: u64 = 5;
pub const WORKER_STARTUP_CHECK_INTERVAL_SECS: u64 = 1;
pub const MAX_CONCURRENT_REQUESTS: i32 = 64;
pub const QUEUE_TIMEOUT_SECS: u64 = 60;
}
/// Builder for common test RouterConfig patterns
pub struct TestRouterConfig;
impl TestRouterConfig {
/// Create a basic round-robin config for routing tests
pub fn round_robin(port: u16) -> RouterConfig {
RouterConfig::builder()
.regular_mode(vec![])
.round_robin_policy()
.host(defaults::HOST)
.port(port)
.max_payload_size(defaults::MAX_PAYLOAD_SIZE)
.request_timeout_secs(defaults::REQUEST_TIMEOUT_SECS)
.worker_startup_timeout_secs(defaults::WORKER_STARTUP_TIMEOUT_SECS)
.worker_startup_check_interval_secs(defaults::WORKER_STARTUP_CHECK_INTERVAL_SECS)
.max_concurrent_requests(defaults::MAX_CONCURRENT_REQUESTS)
.queue_timeout_secs(defaults::QUEUE_TIMEOUT_SECS)
.build_unchecked()
}
/// Create a random load balancing config
pub fn random(port: u16) -> RouterConfig {
RouterConfig::builder()
.regular_mode(vec![])
.random_policy()
.host(defaults::HOST)
.port(port)
.max_payload_size(defaults::MAX_PAYLOAD_SIZE)
.request_timeout_secs(defaults::REQUEST_TIMEOUT_SECS)
.worker_startup_timeout_secs(defaults::WORKER_STARTUP_TIMEOUT_SECS)
.worker_startup_check_interval_secs(defaults::WORKER_STARTUP_CHECK_INTERVAL_SECS)
.max_concurrent_requests(defaults::MAX_CONCURRENT_REQUESTS)
.queue_timeout_secs(defaults::QUEUE_TIMEOUT_SECS)
.build_unchecked()
}
/// Create a cache-aware config for routing tests
pub fn cache_aware(port: u16) -> RouterConfig {
RouterConfig::builder()
.regular_mode(vec![])
.cache_aware_policy(
0.5, // cache_threshold
32, // balance_abs_threshold
1.5, // balance_rel_threshold
60, // eviction_interval_secs
1000, // max_tree_size
)
.host(defaults::HOST)
.port(port)
.max_payload_size(defaults::MAX_PAYLOAD_SIZE)
.request_timeout_secs(defaults::REQUEST_TIMEOUT_SECS)
.worker_startup_timeout_secs(defaults::WORKER_STARTUP_TIMEOUT_SECS)
.worker_startup_check_interval_secs(defaults::WORKER_STARTUP_CHECK_INTERVAL_SECS)
.max_concurrent_requests(defaults::MAX_CONCURRENT_REQUESTS)
.queue_timeout_secs(defaults::QUEUE_TIMEOUT_SECS)
.build_unchecked()
}
/// Create a power-of-two config
pub fn power_of_two(port: u16) -> RouterConfig {
RouterConfig::builder()
.regular_mode(vec![])
.power_of_two_policy(5) // load_check_interval_secs
.host(defaults::HOST)
.port(port)
.max_payload_size(defaults::MAX_PAYLOAD_SIZE)
.request_timeout_secs(defaults::REQUEST_TIMEOUT_SECS)
.worker_startup_timeout_secs(defaults::WORKER_STARTUP_TIMEOUT_SECS)
.worker_startup_check_interval_secs(defaults::WORKER_STARTUP_CHECK_INTERVAL_SECS)
.max_concurrent_requests(defaults::MAX_CONCURRENT_REQUESTS)
.queue_timeout_secs(defaults::QUEUE_TIMEOUT_SECS)
.build_unchecked()
}
/// Create a manual routing config (for sticky routing tests)
pub fn manual(port: u16) -> RouterConfig {
RouterConfig::builder()
.regular_mode(vec![])
.policy(PolicyConfig::Manual {
eviction_interval_secs: 60,
max_idle_secs: 3600,
})
.host(defaults::HOST)
.port(port)
.max_payload_size(defaults::MAX_PAYLOAD_SIZE)
.request_timeout_secs(defaults::REQUEST_TIMEOUT_SECS)
.worker_startup_timeout_secs(defaults::WORKER_STARTUP_TIMEOUT_SECS)
.worker_startup_check_interval_secs(defaults::WORKER_STARTUP_CHECK_INTERVAL_SECS)
.max_concurrent_requests(defaults::MAX_CONCURRENT_REQUESTS)
.queue_timeout_secs(defaults::QUEUE_TIMEOUT_SECS)
.build_unchecked()
}
/// Create a config with custom concurrent request limit (for rate limiting tests)
pub fn with_concurrency(port: u16, max_concurrent: i32) -> RouterConfig {
RouterConfig::builder()
.regular_mode(vec![])
.round_robin_policy()
.host(defaults::HOST)
.port(port)
.max_payload_size(defaults::MAX_PAYLOAD_SIZE)
.request_timeout_secs(defaults::REQUEST_TIMEOUT_SECS)
.worker_startup_timeout_secs(defaults::WORKER_STARTUP_TIMEOUT_SECS)
.worker_startup_check_interval_secs(defaults::WORKER_STARTUP_CHECK_INTERVAL_SECS)
.max_concurrent_requests(max_concurrent)
.queue_timeout_secs(defaults::QUEUE_TIMEOUT_SECS)
.build_unchecked()
}
/// Create a config with custom payload size limit
pub fn with_payload_limit(port: u16, max_payload_size: usize) -> RouterConfig {
RouterConfig::builder()
.regular_mode(vec![])
.round_robin_policy()
.host(defaults::HOST)
.port(port)
.max_payload_size(max_payload_size)
.request_timeout_secs(defaults::REQUEST_TIMEOUT_SECS)
.worker_startup_timeout_secs(defaults::WORKER_STARTUP_TIMEOUT_SECS)
.worker_startup_check_interval_secs(defaults::WORKER_STARTUP_CHECK_INTERVAL_SECS)
.max_concurrent_requests(defaults::MAX_CONCURRENT_REQUESTS)
.queue_timeout_secs(defaults::QUEUE_TIMEOUT_SECS)
.build_unchecked()
}
/// Create a config with short timeouts (for timeout/retry tests)
pub fn with_short_timeouts(port: u16) -> RouterConfig {
RouterConfig::builder()
.regular_mode(vec![])
.round_robin_policy()
.host(defaults::HOST)
.port(port)
.max_payload_size(defaults::MAX_PAYLOAD_SIZE)
.request_timeout_secs(5)
.worker_startup_timeout_secs(2)
.worker_startup_check_interval_secs(1)
.max_concurrent_requests(defaults::MAX_CONCURRENT_REQUESTS)
.queue_timeout_secs(5)
.build_unchecked()
}
/// Create a round-robin config with retry settings
pub fn round_robin_with_retry(port: u16, retry_config: RetryConfig) -> RouterConfig {
RouterConfig::builder()
.regular_mode(vec![])
.round_robin_policy()
.host(defaults::HOST)
.port(port)
.max_payload_size(defaults::MAX_PAYLOAD_SIZE)
.request_timeout_secs(defaults::REQUEST_TIMEOUT_SECS)
.worker_startup_timeout_secs(defaults::WORKER_STARTUP_TIMEOUT_SECS)
.worker_startup_check_interval_secs(defaults::WORKER_STARTUP_CHECK_INTERVAL_SECS)
.max_concurrent_requests(defaults::MAX_CONCURRENT_REQUESTS)
.queue_timeout_secs(defaults::QUEUE_TIMEOUT_SECS)
.retry_config(retry_config)
.build_unchecked()
}
/// Create a round-robin config with circuit breaker
pub fn round_robin_with_circuit_breaker(
port: u16,
circuit_breaker: CircuitBreakerConfig,
) -> RouterConfig {
RouterConfig::builder()
.regular_mode(vec![])
.round_robin_policy()
.host(defaults::HOST)
.port(port)
.max_payload_size(defaults::MAX_PAYLOAD_SIZE)
.request_timeout_secs(defaults::REQUEST_TIMEOUT_SECS)
.worker_startup_timeout_secs(defaults::WORKER_STARTUP_TIMEOUT_SECS)
.worker_startup_check_interval_secs(defaults::WORKER_STARTUP_CHECK_INTERVAL_SECS)
.max_concurrent_requests(defaults::MAX_CONCURRENT_REQUESTS)
.queue_timeout_secs(defaults::QUEUE_TIMEOUT_SECS)
.circuit_breaker_config(circuit_breaker)
.build_unchecked()
}
/// Create a round-robin config with both retry and circuit breaker
pub fn round_robin_with_reliability(
port: u16,
retry_config: RetryConfig,
circuit_breaker: CircuitBreakerConfig,
) -> RouterConfig {
RouterConfig::builder()
.regular_mode(vec![])
.round_robin_policy()
.host(defaults::HOST)
.port(port)
.max_payload_size(defaults::MAX_PAYLOAD_SIZE)
.request_timeout_secs(defaults::REQUEST_TIMEOUT_SECS)
.worker_startup_timeout_secs(defaults::WORKER_STARTUP_TIMEOUT_SECS)
.worker_startup_check_interval_secs(defaults::WORKER_STARTUP_CHECK_INTERVAL_SECS)
.max_concurrent_requests(defaults::MAX_CONCURRENT_REQUESTS)
.queue_timeout_secs(defaults::QUEUE_TIMEOUT_SECS)
.retry_config(retry_config)
.circuit_breaker_config(circuit_breaker)
.build_unchecked()
}
}
/// Builder for common MockWorkerConfig patterns
pub struct TestWorkerConfig;
impl TestWorkerConfig {
/// Create a healthy worker config
pub fn healthy(port: u16) -> MockWorkerConfig {
MockWorkerConfig {
port,
worker_type: WorkerType::Regular,
health_status: HealthStatus::Healthy,
response_delay_ms: 0,
fail_rate: 0.0,
}
}
/// Create multiple healthy workers with sequential ports
pub fn healthy_workers(start_port: u16, count: u16) -> Vec<MockWorkerConfig> {
(0..count).map(|i| Self::healthy(start_port + i)).collect()
}
/// Create an unhealthy worker config
pub fn unhealthy(port: u16) -> MockWorkerConfig {
MockWorkerConfig {
port,
worker_type: WorkerType::Regular,
health_status: HealthStatus::Unhealthy,
response_delay_ms: 0,
fail_rate: 0.0,
}
}
/// Create a slow worker config (for timeout tests)
pub fn slow(port: u16, delay_ms: u64) -> MockWorkerConfig {
MockWorkerConfig {
port,
worker_type: WorkerType::Regular,
health_status: HealthStatus::Healthy,
response_delay_ms: delay_ms,
fail_rate: 0.0,
}
}
/// Create a flaky worker config (for retry/fault tolerance tests)
pub fn flaky(port: u16, fail_rate: f32) -> MockWorkerConfig {
MockWorkerConfig {
port,
worker_type: WorkerType::Regular,
health_status: HealthStatus::Healthy,
response_delay_ms: 0,
fail_rate,
}
}
/// Create a decode worker config (for PD routing tests)
pub fn decode(port: u16) -> MockWorkerConfig {
MockWorkerConfig {
port,
worker_type: WorkerType::Decode,
health_status: HealthStatus::Healthy,
response_delay_ms: 0,
fail_rate: 0.0,
}
}
/// Create a prefill worker config (for PD routing tests)
pub fn prefill(port: u16) -> MockWorkerConfig {
MockWorkerConfig {
port,
worker_type: WorkerType::Prefill,
health_status: HealthStatus::Healthy,
response_delay_ms: 0,
fail_rate: 0.0,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_round_robin_config() {
let config = TestRouterConfig::round_robin(3000);
assert_eq!(config.port, 3000);
}
#[test]
fn test_healthy_workers() {
let workers = TestWorkerConfig::healthy_workers(8000, 3);
assert_eq!(workers.len(), 3);
assert_eq!(workers[0].port, 8000);
assert_eq!(workers[1].port, 8001);
assert_eq!(workers[2].port, 8002);
}
}
@@ -0,0 +1,370 @@
// TLS-enabled mock worker for mTLS integration tests
#![allow(dead_code)]
use std::{
net::SocketAddr,
path::Path,
sync::{Arc, Once},
time::{SystemTime, UNIX_EPOCH},
};
use axum::{
extract::{Json, State},
http::StatusCode,
response::{IntoResponse, Response},
routing::{get, post},
Router,
};
use axum_server::tls_rustls::RustlsConfig;
use rustls::{server::WebPkiClientVerifier, RootCertStore};
use serde_json::json;
use tokio::sync::RwLock;
// Ensure crypto provider is installed exactly once
static CRYPTO_PROVIDER_INIT: Once = Once::new();
fn ensure_crypto_provider() {
CRYPTO_PROVIDER_INIT.call_once(|| {
let _ = rustls::crypto::ring::default_provider().install_default();
});
}
/// Configuration for TLS mock worker behavior
#[derive(Clone)]
pub struct TlsMockWorkerConfig {
pub port: u16,
/// Require client certificate (mTLS) or just server TLS
pub require_client_cert: bool,
/// Response delay in milliseconds
pub response_delay_ms: u64,
/// Fail rate (0.0 - 1.0)
pub fail_rate: f32,
}
impl Default for TlsMockWorkerConfig {
fn default() -> Self {
Self {
port: 0,
require_client_cert: true,
response_delay_ms: 0,
fail_rate: 0.0,
}
}
}
/// TLS-enabled mock worker server for mTLS testing
pub struct TlsMockWorker {
config: Arc<RwLock<TlsMockWorkerConfig>>,
shutdown_handle: Option<tokio::task::JoinHandle<()>>,
shutdown_tx: Option<tokio::sync::oneshot::Sender<()>>,
}
impl TlsMockWorker {
pub fn new(config: TlsMockWorkerConfig) -> Self {
Self {
config: Arc::new(RwLock::new(config)),
shutdown_handle: None,
shutdown_tx: None,
}
}
/// Start the TLS mock worker server
///
/// # Arguments
/// * `server_cert_path` - Path to server certificate PEM file
/// * `server_key_path` - Path to server private key PEM file
/// * `ca_cert_path` - Path to CA certificate for client verification (optional for TLS-only mode)
pub async fn start(
&mut self,
server_cert_path: &Path,
server_key_path: &Path,
ca_cert_path: Option<&Path>,
) -> Result<String, Box<dyn std::error::Error + Send + Sync>> {
// Ensure crypto provider is installed before using rustls
ensure_crypto_provider();
let config = self.config.clone();
let port = config.read().await.port;
let require_client_cert = config.read().await.require_client_cert;
// If port is 0, find an available port
let port = if port == 0 {
let listener = std::net::TcpListener::bind("127.0.0.1:0")?;
let port = listener.local_addr()?.port();
drop(listener);
config.write().await.port = port;
port
} else {
port
};
let app = Router::new()
.route("/health", get(health_handler))
.route("/health_generate", get(health_generate_handler))
.route("/get_server_info", get(server_info_handler))
.route("/generate", post(generate_handler))
.route("/v1/chat/completions", post(chat_completions_handler))
.with_state(config);
let (shutdown_tx, mut shutdown_rx) = tokio::sync::oneshot::channel::<()>();
self.shutdown_tx = Some(shutdown_tx);
// Build TLS configuration
let rustls_config = if require_client_cert {
// mTLS: require client certificate
let ca_cert_path = ca_cert_path.ok_or("CA cert path required for mTLS")?;
build_mtls_config(server_cert_path, server_key_path, ca_cert_path).await?
} else {
// TLS only: no client cert required
build_tls_config(server_cert_path, server_key_path).await?
};
let addr = SocketAddr::from(([127, 0, 0, 1], port));
// Spawn the server in a separate task
let handle = tokio::spawn(async move {
let server =
axum_server::bind_rustls(addr, rustls_config).serve(app.into_make_service());
tokio::select! {
result = server => {
if let Err(e) = result {
eprintln!("TLS Server error: {}", e);
}
}
_ = &mut shutdown_rx => {
// Graceful shutdown
}
}
});
self.shutdown_handle = Some(handle);
// Wait for the server to start
tokio::time::sleep(tokio::time::Duration::from_millis(200)).await;
let url = format!("https://127.0.0.1:{}", port);
Ok(url)
}
/// Stop the TLS mock worker server
pub async fn stop(&mut self) {
if let Some(shutdown_tx) = self.shutdown_tx.take() {
let _ = shutdown_tx.send(());
}
if let Some(handle) = self.shutdown_handle.take() {
let _ = tokio::time::timeout(tokio::time::Duration::from_secs(5), handle).await;
}
}
}
impl Drop for TlsMockWorker {
fn drop(&mut self) {
if let Some(shutdown_tx) = self.shutdown_tx.take() {
let _ = shutdown_tx.send(());
}
}
}
/// Build TLS config for server-only TLS (no client cert required)
async fn build_tls_config(
cert_path: &Path,
key_path: &Path,
) -> Result<RustlsConfig, Box<dyn std::error::Error + Send + Sync>> {
let config = RustlsConfig::from_pem_file(cert_path, key_path).await?;
Ok(config)
}
/// Build mTLS config requiring client certificate
async fn build_mtls_config(
cert_path: &Path,
key_path: &Path,
ca_cert_path: &Path,
) -> Result<RustlsConfig, Box<dyn std::error::Error + Send + Sync>> {
use std::io::BufReader;
use rustls_pemfile::{certs, pkcs8_private_keys};
// Read server certificate
let cert_file = std::fs::File::open(cert_path)?;
let mut reader = BufReader::new(cert_file);
let cert_chain: Vec<_> = certs(&mut reader).filter_map(|r| r.ok()).collect();
// Read server private key
let key_file = std::fs::File::open(key_path)?;
let mut reader = BufReader::new(key_file);
let private_key = pkcs8_private_keys(&mut reader)
.next()
.ok_or("No private key found")??;
// Read CA certificate for client verification
let ca_file = std::fs::File::open(ca_cert_path)?;
let mut reader = BufReader::new(ca_file);
let ca_certs: Vec<_> = certs(&mut reader).filter_map(|r| r.ok()).collect();
// Build root certificate store for client verification
let mut root_store = RootCertStore::empty();
for cert in ca_certs {
root_store.add(cert)?;
}
// Create client certificate verifier
let client_verifier = WebPkiClientVerifier::builder(Arc::new(root_store))
.build()
.map_err(|e| format!("Failed to build client verifier: {}", e))?;
// Build server config with client verification
let server_config = rustls::ServerConfig::builder()
.with_client_cert_verifier(client_verifier)
.with_single_cert(cert_chain, private_key.into())
.map_err(|e| format!("Failed to build server config: {}", e))?;
Ok(RustlsConfig::from_config(Arc::new(server_config)))
}
// Handler implementations (simplified versions of mock_worker handlers)
async fn should_fail(config: &TlsMockWorkerConfig) -> bool {
rand::random::<f32>() < config.fail_rate
}
async fn health_handler(State(config): State<Arc<RwLock<TlsMockWorkerConfig>>>) -> Response {
let config = config.read().await;
if should_fail(&config).await {
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({ "error": "Random failure" })),
)
.into_response();
}
Json(json!({
"status": "healthy",
"timestamp": SystemTime::now().duration_since(UNIX_EPOCH).unwrap().as_secs(),
"tls_enabled": true
}))
.into_response()
}
async fn health_generate_handler(
State(config): State<Arc<RwLock<TlsMockWorkerConfig>>>,
) -> Response {
let config = config.read().await;
if should_fail(&config).await {
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({ "error": "Random failure" })),
)
.into_response();
}
Json(json!({
"status": "ok",
"queue_length": 0,
"processing_time_ms": config.response_delay_ms
}))
.into_response()
}
async fn server_info_handler(State(config): State<Arc<RwLock<TlsMockWorkerConfig>>>) -> Response {
let config = config.read().await;
if should_fail(&config).await {
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({ "error": "Random failure" })),
)
.into_response();
}
Json(json!({
"model_path": "mock-tls-model",
"port": config.port,
"host": "127.0.0.1",
"tls_enabled": true,
"version": "0.3.0"
}))
.into_response()
}
async fn generate_handler(
State(config): State<Arc<RwLock<TlsMockWorkerConfig>>>,
Json(_payload): Json<serde_json::Value>,
) -> Response {
let config = config.read().await;
if should_fail(&config).await {
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({ "error": "Random failure" })),
)
.into_response();
}
if config.response_delay_ms > 0 {
tokio::time::sleep(tokio::time::Duration::from_millis(config.response_delay_ms)).await;
}
Json(json!({
"text": "This is a mock TLS response.",
"meta_info": {
"prompt_tokens": 10,
"completion_tokens": 5,
"tls_verified": true
}
}))
.into_response()
}
async fn chat_completions_handler(
State(config): State<Arc<RwLock<TlsMockWorkerConfig>>>,
Json(_payload): Json<serde_json::Value>,
) -> Response {
let config = config.read().await;
if should_fail(&config).await {
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({
"error": {
"message": "Random failure",
"type": "internal_error"
}
})),
)
.into_response();
}
if config.response_delay_ms > 0 {
tokio::time::sleep(tokio::time::Duration::from_millis(config.response_delay_ms)).await;
}
let timestamp = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs();
Json(json!({
"id": format!("chatcmpl-{}", uuid::Uuid::new_v4()),
"object": "chat.completion",
"created": timestamp,
"model": "mock-tls-model",
"choices": [{
"index": 0,
"message": {
"role": "assistant",
"content": "This is a mock TLS chat response."
},
"finish_reason": "stop"
}],
"usage": {
"prompt_tokens": 10,
"completion_tokens": 5,
"total_tokens": 15
}
}))
.into_response()
}