use log::debug;
use std::sync::Arc;
use tokio::sync::RwLock;
use super::priority::PriorityCalculator;
use super::queue::{QueuedRequest, ServiceRequest};
use super::response_channel::ResponseChannel;
use super::worker::WorkerManager;
use crate::domain::ServiceResponse;
use crate::error::VecboostError;
use crate::i18n;
use crate::service::embedding::EmbeddingService;
use crate::service::rerank::RerankService;
pub struct PipelineScheduler {
priority_calculator: PriorityCalculator,
response_channel: Arc<ResponseChannel>,
worker_manager: Arc<WorkerManager>,
service: Arc<RwLock<EmbeddingService>>,
rerank_service: Option<Arc<RwLock<RerankService>>>,
}
impl PipelineScheduler {
pub fn new(
priority_calculator: PriorityCalculator,
response_channel: Arc<ResponseChannel>,
worker_manager: Arc<WorkerManager>,
service: Arc<RwLock<EmbeddingService>>,
) -> Self {
debug!("Creating PipelineScheduler");
Self {
priority_calculator,
response_channel,
worker_manager,
service,
rerank_service: None,
}
}
pub fn with_rerank_service(mut self, rerank_service: Arc<RwLock<RerankService>>) -> Self {
self.rerank_service = Some(rerank_service);
self
}
pub async fn process_request(
&self,
request: QueuedRequest,
) -> Result<ServiceResponse, VecboostError> {
debug!("Processing request {}", request.request_id);
match request.request {
ServiceRequest::Embed(embed_req) => {
let service = self.service.read().await;
let resp = service.process_text(embed_req, None).await?;
Ok(ServiceResponse::Embed(resp))
}
ServiceRequest::Rerank(rerank_req) => {
let rerank_service = self.rerank_service.as_ref().ok_or_else(|| {
VecboostError::InternalError(i18n::tr("rerank-not-configured"))
})?;
let service = rerank_service.read().await;
let resp = service.process_rerank(rerank_req, 100, 8192).await?;
Ok(ServiceResponse::Rerank(resp))
}
}
}
pub fn worker_manager(&self) -> Arc<WorkerManager> {
Arc::clone(&self.worker_manager)
}
pub fn response_channel(&self) -> Arc<ResponseChannel> {
Arc::clone(&self.response_channel)
}
pub fn priority_calculator(&self) -> &PriorityCalculator {
&self.priority_calculator
}
}
#[cfg(test)]
mod tests {
use super::*;
fn ensure_i18n_init() {
i18n::init();
}
use crate::config::model::{ModelConfig, Precision};
use crate::domain::EmbedRequest;
use crate::engine::InferenceEngine;
use crate::pipeline::config::{PriorityConfig, WorkerConfig};
use crate::pipeline::priority::{Priority, PriorityInput, RequestSource};
use crate::pipeline::queue::PriorityRequestQueue;
use crate::service::rerank::RerankService;
use async_trait::async_trait;
use std::time::{Duration, Instant};
struct TestEngine {
dimension: usize,
}
impl TestEngine {
fn new(dimension: usize) -> Self {
Self { dimension }
}
}
#[async_trait]
impl InferenceEngine for TestEngine {
fn embed(&self, _text: &str) -> Result<Vec<f32>, VecboostError> {
Ok(vec![0.5; self.dimension])
}
fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, VecboostError> {
Ok(texts.iter().map(|_| vec![0.5; self.dimension]).collect())
}
fn precision(&self) -> &Precision {
static PRECISION: Precision = Precision::Fp32;
&PRECISION
}
fn supports_mixed_precision(&self) -> bool {
false
}
async fn try_fallback_to_cpu(
&mut self,
_config: &ModelConfig,
) -> Result<(), VecboostError> {
Ok(())
}
}
fn create_test_service() -> Arc<RwLock<EmbeddingService>> {
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(TestEngine::new(384)));
Arc::new(RwLock::new(EmbeddingService::new(engine, None)))
}
#[tokio::test(flavor = "multi_thread")]
async fn test_scheduler_creation() {
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_test_service();
let worker_manager = Arc::new(WorkerManager::new(
Arc::new(PriorityRequestQueue::new(100)),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = PipelineScheduler::new(
priority_calculator,
response_channel,
worker_manager,
service,
);
assert_eq!(scheduler.worker_manager().current_workers(), 0);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_process_request_success() {
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_test_service();
let worker_manager = Arc::new(WorkerManager::new(
Arc::new(PriorityRequestQueue::new(100)),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = PipelineScheduler::new(
priority_calculator,
response_channel,
worker_manager,
service,
);
let request = QueuedRequest {
request_id: "test-001".to_string(),
request: ServiceRequest::Embed(EmbedRequest {
text: "hello world".to_string(),
normalize: None,
}),
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
};
let result = scheduler.process_request(request).await;
assert!(result.is_ok(), "process_request should succeed");
let response = result.unwrap();
let ServiceResponse::Embed(response) = response else {
panic!("Expected ServiceResponse::Embed");
};
assert_eq!(response.dimension, 384);
assert_eq!(response.embedding.len(), 384);
let expected = 0.5f32 / (384f32 * 0.25f32).sqrt();
assert!(
response
.embedding
.iter()
.all(|&v| (v - expected).abs() < 1e-6),
"all values should equal L2-normalized 0.5, got {:?}",
&response.embedding[..5.min(response.embedding.len())]
);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_process_request_empty_text() {
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_test_service();
let worker_manager = Arc::new(WorkerManager::new(
Arc::new(PriorityRequestQueue::new(100)),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = PipelineScheduler::new(
priority_calculator,
response_channel,
worker_manager,
service,
);
let request = QueuedRequest {
request_id: "test-002".to_string(),
request: ServiceRequest::Embed(EmbedRequest {
text: "".to_string(),
normalize: None,
}),
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
};
let result = scheduler.process_request(request).await;
assert!(result.is_err(), "empty text should return validation error");
}
struct ErrorEngine;
impl ErrorEngine {
fn new() -> Self {
Self
}
}
#[async_trait]
impl InferenceEngine for ErrorEngine {
fn embed(&self, _text: &str) -> Result<Vec<f32>, VecboostError> {
Err(VecboostError::InferenceError(
"mock inference failure".to_string(),
))
}
fn embed_batch(&self, _texts: &[String]) -> Result<Vec<Vec<f32>>, VecboostError> {
Err(VecboostError::InferenceError(
"mock batch inference failure".to_string(),
))
}
fn precision(&self) -> &Precision {
static PRECISION: Precision = Precision::Fp32;
&PRECISION
}
fn supports_mixed_precision(&self) -> bool {
false
}
async fn try_fallback_to_cpu(
&mut self,
_config: &ModelConfig,
) -> Result<(), VecboostError> {
Ok(())
}
}
fn create_error_service() -> Arc<RwLock<EmbeddingService>> {
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(ErrorEngine::new()));
Arc::new(RwLock::new(EmbeddingService::new(engine, None)))
}
#[tokio::test(flavor = "multi_thread")]
async fn test_process_request_engine_error() {
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_error_service();
let worker_manager = Arc::new(WorkerManager::new(
Arc::new(PriorityRequestQueue::new(100)),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = PipelineScheduler::new(
priority_calculator,
response_channel,
worker_manager,
service,
);
let request = QueuedRequest {
request_id: "test-err-001".to_string(),
request: ServiceRequest::Embed(EmbedRequest {
text: "hello world".to_string(),
normalize: None,
}),
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
};
let result = scheduler.process_request(request).await;
assert!(
result.is_err(),
"process_request should propagate engine error"
);
match result.unwrap_err() {
VecboostError::InferenceError(msg) => {
assert!(
msg.contains("mock inference failure"),
"error should come from ErrorEngine, got: {}",
msg
);
}
other => panic!("expected InferenceError, got {:?}", other),
}
}
#[tokio::test(flavor = "multi_thread")]
async fn test_process_request_concurrent() {
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_test_service();
let worker_manager = Arc::new(WorkerManager::new(
Arc::new(PriorityRequestQueue::new(100)),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = Arc::new(PipelineScheduler::new(
priority_calculator,
response_channel,
worker_manager,
service,
));
let total = 10;
let mut handles = Vec::with_capacity(total);
for i in 0..total {
let scheduler = Arc::clone(&scheduler);
handles.push(tokio::spawn(async move {
let request = QueuedRequest {
request_id: format!("test-concurrent-{:03}", i),
request: ServiceRequest::Embed(EmbedRequest {
text: format!("hello world {}", i),
normalize: None,
}),
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
};
scheduler.process_request(request).await
}));
}
let results = futures::future::join_all(handles).await;
assert_eq!(results.len(), total, "all tasks should complete");
for (i, result) in results.into_iter().enumerate() {
assert!(
result.is_ok(),
"task {} panicked or was cancelled: {:?}",
i,
result.err()
);
let response = result.unwrap().unwrap();
let ServiceResponse::Embed(response) = response else {
panic!("Expected ServiceResponse::Embed for task {}", i);
};
assert_eq!(
response.dimension, 384,
"task {} should return 384-dim embedding",
i
);
assert_eq!(
response.embedding.len(),
384,
"task {} embedding length mismatch",
i
);
}
}
#[tokio::test(flavor = "multi_thread")]
async fn test_scheduler_accessors() {
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_test_service();
let queue = Arc::new(PriorityRequestQueue::new(100));
let worker_manager = Arc::new(WorkerManager::new(
queue.clone(),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = PipelineScheduler::new(
priority_calculator,
response_channel.clone(),
worker_manager.clone(),
service,
);
let wm = scheduler.worker_manager();
assert!(Arc::ptr_eq(&wm, &worker_manager));
let rc = scheduler.response_channel();
assert!(Arc::ptr_eq(&rc, &response_channel));
let _pc = scheduler.priority_calculator();
}
#[tokio::test(flavor = "multi_thread")]
async fn test_process_request_with_normalize_true() {
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_test_service();
let worker_manager = Arc::new(WorkerManager::new(
Arc::new(PriorityRequestQueue::new(100)),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = PipelineScheduler::new(
priority_calculator,
response_channel,
worker_manager,
service,
);
let request = QueuedRequest {
request_id: "test-norm-true".to_string(),
request: ServiceRequest::Embed(EmbedRequest {
text: "hello world".to_string(),
normalize: Some(true),
}),
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
};
let result = scheduler.process_request(request).await;
assert!(result.is_ok());
let response = result.unwrap();
let ServiceResponse::Embed(response) = response else {
panic!("Expected ServiceResponse::Embed");
};
assert_eq!(response.dimension, 384);
let norm: f32 = response.embedding.iter().map(|v| v * v).sum::<f32>().sqrt();
assert!(
(norm - 1.0).abs() < 1e-5,
"L2 norm should be 1.0, got {}",
norm
);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_process_request_with_normalize_false() {
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_test_service();
let worker_manager = Arc::new(WorkerManager::new(
Arc::new(PriorityRequestQueue::new(100)),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = PipelineScheduler::new(
priority_calculator,
response_channel,
worker_manager,
service,
);
let request = QueuedRequest {
request_id: "test-norm-false".to_string(),
request: ServiceRequest::Embed(EmbedRequest {
text: "hello world".to_string(),
normalize: Some(false),
}),
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
};
let result = scheduler.process_request(request).await;
assert!(result.is_ok());
let response = result.unwrap();
let ServiceResponse::Embed(response) = response else {
panic!("Expected ServiceResponse::Embed");
};
assert_eq!(response.dimension, 384);
let norm: f32 = response.embedding.iter().map(|v| v * v).sum::<f32>().sqrt();
assert!(
(norm - 1.0).abs() < 1e-5,
"process_text always applies L2 normalization regardless of normalize flag, got {}",
norm
);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_process_request_whitespace_only_returns_error() {
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_test_service();
let worker_manager = Arc::new(WorkerManager::new(
Arc::new(PriorityRequestQueue::new(100)),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = PipelineScheduler::new(
priority_calculator,
response_channel,
worker_manager,
service,
);
let request = QueuedRequest {
request_id: "test-ws".to_string(),
request: ServiceRequest::Embed(EmbedRequest {
text: " ".to_string(),
normalize: None,
}),
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
};
let result = scheduler.process_request(request).await;
assert!(result.is_err());
match result.unwrap_err() {
VecboostError::InvalidInput(_) => {}
other => panic!("expected InvalidInput for whitespace-only, got {:?}", other),
}
}
#[tokio::test(flavor = "multi_thread")]
async fn test_process_request_critical_priority() {
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_test_service();
let worker_manager = Arc::new(WorkerManager::new(
Arc::new(PriorityRequestQueue::new(100)),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = PipelineScheduler::new(
priority_calculator,
response_channel,
worker_manager,
service,
);
let request = QueuedRequest {
request_id: "test-critical".to_string(),
request: ServiceRequest::Embed(EmbedRequest {
text: "urgent request".to_string(),
normalize: Some(true),
}),
priority: Priority::Critical,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Internal,
};
let result = scheduler.process_request(request).await;
assert!(result.is_ok(), "critical priority request should succeed");
let response = result.unwrap();
let ServiceResponse::Embed(embed_resp) = response else {
panic!("Expected Embed response");
};
assert_eq!(embed_resp.dimension, 384);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_process_request_low_priority() {
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_test_service();
let worker_manager = Arc::new(WorkerManager::new(
Arc::new(PriorityRequestQueue::new(100)),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = PipelineScheduler::new(
priority_calculator,
response_channel,
worker_manager,
service,
);
let request = QueuedRequest {
request_id: "test-low".to_string(),
request: ServiceRequest::Embed(EmbedRequest {
text: "low priority request".to_string(),
normalize: Some(true),
}),
priority: Priority::Low,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Grpc {
client_id: "client-1".to_string(),
},
};
let result = scheduler.process_request(request).await;
assert!(result.is_ok(), "low priority request should succeed");
let response = result.unwrap();
let ServiceResponse::Embed(embed_resp) = response else {
panic!("Expected Embed response");
};
assert_eq!(embed_resp.dimension, 384);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_process_request_long_text() {
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_test_service();
let worker_manager = Arc::new(WorkerManager::new(
Arc::new(PriorityRequestQueue::new(100)),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = PipelineScheduler::new(
priority_calculator,
response_channel,
worker_manager,
service,
);
let request = QueuedRequest {
request_id: "test-long".to_string(),
request: ServiceRequest::Embed(EmbedRequest {
text: "X".repeat(1000),
normalize: Some(true),
}),
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
};
let result = scheduler.process_request(request).await;
assert!(result.is_ok(), "long text should succeed");
let response = result.unwrap();
let ServiceResponse::Embed(response) = response else {
panic!("Expected ServiceResponse::Embed");
};
assert_eq!(response.dimension, 384);
assert_eq!(response.embedding.len(), 384);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_priority_calculator_accessor_usable() {
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_test_service();
let worker_manager = Arc::new(WorkerManager::new(
Arc::new(PriorityRequestQueue::new(100)),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = PipelineScheduler::new(
priority_calculator,
response_channel,
worker_manager,
service,
);
let pc = scheduler.priority_calculator();
let input = PriorityInput {
base_priority: Priority::Normal,
time_until_timeout: Duration::from_millis(50),
user_tier: None,
source: RequestSource::Internal,
queue_length: 0,
};
let result = pc.calculate(input);
assert_eq!(result, Priority::Critical);
}
struct RerankCapableEngine;
#[async_trait]
impl InferenceEngine for RerankCapableEngine {
fn embed(&self, _text: &str) -> Result<Vec<f32>, VecboostError> {
Ok(vec![0.0; 128])
}
fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, VecboostError> {
Ok(texts.iter().map(|_| vec![0.0; 128]).collect())
}
fn precision(&self) -> &Precision {
&Precision::Fp32
}
fn supports_mixed_precision(&self) -> bool {
false
}
fn rerank(&self, _query: &str, document: &str) -> Result<f32, VecboostError> {
Ok(document.len() as f32 / 100.0)
}
fn rerank_batch(
&self,
query: &str,
documents: &[String],
) -> Result<Vec<f32>, VecboostError> {
documents
.iter()
.map(|doc| self.rerank(query, doc))
.collect()
}
fn supports_rerank(&self) -> bool {
true
}
async fn try_fallback_to_cpu(
&mut self,
_config: &ModelConfig,
) -> Result<(), VecboostError> {
Ok(())
}
}
fn create_rerank_service() -> Arc<RwLock<RerankService>> {
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(RerankCapableEngine));
Arc::new(RwLock::new(RerankService::new(engine, None)))
}
#[tokio::test(flavor = "multi_thread")]
async fn test_service_request_embed_routes_to_embedding() {
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_test_service();
let worker_manager = Arc::new(WorkerManager::new(
Arc::new(PriorityRequestQueue::new(100)),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = PipelineScheduler::new(
priority_calculator,
response_channel,
worker_manager,
service,
);
let request = QueuedRequest {
request_id: "embed-route-test".to_string(),
request: ServiceRequest::Embed(EmbedRequest {
text: "hello".to_string(),
normalize: None,
}),
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
};
let result = scheduler.process_request(request).await.unwrap();
match result {
ServiceResponse::Embed(resp) => {
assert_eq!(resp.dimension, 384);
}
_ => panic!("Expected ServiceResponse::Embed"),
}
}
#[tokio::test(flavor = "multi_thread")]
async fn test_service_request_rerank_routes_to_rerank() {
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_test_service();
let rerank_service = create_rerank_service();
let worker_manager = Arc::new(WorkerManager::new(
Arc::new(PriorityRequestQueue::new(100)),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = PipelineScheduler::new(
priority_calculator,
response_channel,
worker_manager,
service,
)
.with_rerank_service(rerank_service);
let request = QueuedRequest {
request_id: "rerank-route-test".to_string(),
request: ServiceRequest::Rerank(crate::domain::RerankRequest {
query: "what is rust?".to_string(),
documents: vec!["short".to_string(), "a longer document".to_string()],
top_k: None,
return_documents: None,
}),
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
};
let result = scheduler.process_request(request).await.unwrap();
match result {
ServiceResponse::Rerank(resp) => {
assert_eq!(resp.results.len(), 2);
assert!(resp.results[0].score >= resp.results[1].score);
}
_ => panic!("Expected ServiceResponse::Rerank"),
}
}
#[tokio::test(flavor = "multi_thread")]
async fn test_rerank_without_service_returns_error() {
ensure_i18n_init();
let priority_calculator = PriorityCalculator::new(PriorityConfig::default());
let response_channel = Arc::new(ResponseChannel::new());
let service = create_test_service();
let worker_manager = Arc::new(WorkerManager::new(
Arc::new(PriorityRequestQueue::new(100)),
response_channel.clone(),
WorkerConfig::default(),
service.clone(),
));
let scheduler = PipelineScheduler::new(
priority_calculator,
response_channel,
worker_manager,
service,
);
let request = QueuedRequest {
request_id: "rerank-no-service".to_string(),
request: ServiceRequest::Rerank(crate::domain::RerankRequest {
query: "test".to_string(),
documents: vec!["doc".to_string()],
top_k: None,
return_documents: None,
}),
priority: Priority::Normal,
submitted_at: Instant::now(),
timeout: Duration::from_secs(30),
source: RequestSource::Http {
ip: "127.0.0.1".to_string(),
},
};
let result = scheduler.process_request(request).await;
assert!(result.is_err());
match result.unwrap_err() {
VecboostError::InternalError(msg) => {
assert!(msg.contains("not configured"), "got: {}", msg);
}
other => panic!("Expected InternalError, got: {:?}", other),
}
}
#[test]
fn test_test_engine_embed_batch() {
let engine = TestEngine::new(128);
let texts = vec!["a".to_string(), "b".to_string()];
let vecs = engine.embed_batch(&texts).unwrap();
assert_eq!(vecs.len(), 2);
assert_eq!(vecs[0].len(), 128);
}
#[test]
fn test_test_engine_precision() {
let engine = TestEngine::new(4);
assert_eq!(*engine.precision(), Precision::Fp32);
}
#[test]
fn test_test_engine_supports_mixed_precision() {
let engine = TestEngine::new(4);
assert!(!engine.supports_mixed_precision());
}
#[tokio::test]
async fn test_test_engine_try_fallback() {
let mut engine = TestEngine::new(4);
let config = ModelConfig {
name: "test".to_string(),
engine_type: crate::config::model::EngineType::Candle,
model_path: std::path::PathBuf::from("/tmp"),
tokenizer_path: None,
device: crate::config::model::DeviceType::Cpu,
max_batch_size: 1,
pooling_mode: None,
expected_dimension: None,
memory_limit_bytes: None,
oom_fallback_enabled: false,
model_sha256: None,
quantized: false,
};
assert!(engine.try_fallback_to_cpu(&config).await.is_ok());
}
#[test]
fn test_error_engine_embed_batch() {
let engine = ErrorEngine::new();
let texts = vec!["a".to_string()];
assert!(engine.embed_batch(&texts).is_err());
}
#[test]
fn test_error_engine_precision() {
let engine = ErrorEngine::new();
assert_eq!(*engine.precision(), Precision::Fp32);
}
#[test]
fn test_error_engine_supports_mixed_precision() {
let engine = ErrorEngine::new();
assert!(!engine.supports_mixed_precision());
}
#[tokio::test]
async fn test_error_engine_try_fallback() {
let mut engine = ErrorEngine::new();
let config = ModelConfig {
name: "test".to_string(),
engine_type: crate::config::model::EngineType::Candle,
model_path: std::path::PathBuf::from("/tmp"),
tokenizer_path: None,
device: crate::config::model::DeviceType::Cpu,
max_batch_size: 1,
pooling_mode: None,
expected_dimension: None,
memory_limit_bytes: None,
oom_fallback_enabled: false,
model_sha256: None,
quantized: false,
};
assert!(engine.try_fallback_to_cpu(&config).await.is_ok());
}
}