[model-gateway] : Rust integration tests for integration_mock replacement (#16441)
This commit is contained in:
@@ -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,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(¬_before)?;
|
||||
cert_builder.set_not_after(¬_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(¬_before)?;
|
||||
cert_builder.set_not_after(¬_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(¬_before)?;
|
||||
cert_builder.set_not_after(¬_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()
|
||||
}
|
||||
Reference in New Issue
Block a user