use ferrum_engine::{ContinuousBatchEngine, LlmInferenceEngine, Scheduler, Tokenizer};
use ferrum_scheduler::implementations::ContinuousBatchScheduler;
use ferrum_testkit::{
MockKvCacheManager, MockModelExecutor, MockSampler, MockTensorFactory, MockTokenizer,
};
use ferrum_types::{InferenceRequest, InferenceResponse, SchedulerConfig, SpecialTokens, TokenId};
use std::sync::Arc;
use std::time::Duration;
const VOCAB_SIZE: usize = 1000;
const JSON_OBJECT_TOKEN_ID: u32 = 123;
fn make_engine() -> ContinuousBatchEngine {
make_engine_with_tokenizer(Arc::new(MockTokenizer::new(VOCAB_SIZE)))
}
fn make_engine_with_tokenizer(
tokenizer: Arc<dyn Tokenizer + Send + Sync>,
) -> ContinuousBatchEngine {
let config = ferrum_types::EngineConfig::default();
let scheduler = Arc::new(ContinuousBatchScheduler::new(SchedulerConfig::default()));
let sampler = Arc::new(MockSampler);
let kv_cache = Arc::new(MockKvCacheManager::new(1024));
let executor = Arc::new(MockModelExecutor::instant(VOCAB_SIZE));
let tensor_factory = Arc::new(MockTensorFactory);
ContinuousBatchEngine::new(
config,
scheduler,
tokenizer,
sampler,
kv_cache,
executor,
tensor_factory,
)
.expect("legacy engine composition must match executor authority")
}
struct JsonObjectTokenizer {
inner: MockTokenizer,
}
impl JsonObjectTokenizer {
fn new() -> Self {
Self {
inner: MockTokenizer::new(VOCAB_SIZE),
}
}
}
impl Tokenizer for JsonObjectTokenizer {
fn encode(&self, text: &str, add_special: bool) -> ferrum_types::Result<Vec<TokenId>> {
self.inner.encode(text, add_special)
}
fn decode(&self, tokens: &[TokenId], skip_special: bool) -> ferrum_types::Result<String> {
let mut output = String::new();
for token in tokens {
if token.get() == JSON_OBJECT_TOKEN_ID {
output.push_str("{}");
} else {
if !output.is_empty() {
output.push(' ');
}
output.push_str(&self.inner.decode(&[*token], skip_special)?);
}
}
Ok(output)
}
fn decode_incremental(
&self,
previous: &[TokenId],
next: TokenId,
) -> ferrum_types::Result<String> {
if next.get() == JSON_OBJECT_TOKEN_ID {
Ok("{}".to_string())
} else {
self.inner.decode_incremental(previous, next)
}
}
fn vocab_size(&self) -> usize {
self.inner.vocab_size()
}
fn special_tokens(&self) -> &SpecialTokens {
self.inner.special_tokens()
}
fn token_id(&self, text: &str) -> Option<TokenId> {
(text == "{}")
.then_some(TokenId::new(JSON_OBJECT_TOKEN_ID))
.or_else(|| self.inner.token_id(text))
}
fn token_text(&self, token_id: TokenId) -> Option<&str> {
self.inner.token_text(token_id)
}
fn token_bytes(&self, token_id: TokenId) -> Option<Vec<u8>> {
if token_id.get() == JSON_OBJECT_TOKEN_ID {
Some(b"{}".to_vec())
} else {
self.inner.token_bytes(token_id)
}
}
fn info(&self) -> ferrum_interfaces::tokenizer::TokenizerInfo {
self.inner.info()
}
}
fn make_request(prompt: &str) -> InferenceRequest {
let mut req = InferenceRequest::new(prompt, "mock-model");
req.sampling_params.max_tokens = 5;
req.sampling_params.temperature = 0.0; req
}
#[tokio::test]
async fn single_request_completes() {
let engine = make_engine();
let request = make_request("Hello world");
let response = engine.infer(request).await.unwrap();
assert_eq!(response.finish_reason, ferrum_types::FinishReason::Length);
assert!(!response.text.is_empty());
}
#[tokio::test]
async fn multiple_requests_complete_sequentially() {
let engine = make_engine();
for i in 0..5 {
let request = make_request(&format!("Request number {}", i));
let response = engine.infer(request).await.unwrap();
assert_eq!(response.finish_reason, ferrum_types::FinishReason::Length);
}
}
#[tokio::test]
async fn streaming_produces_chunks() {
use futures::StreamExt;
let engine = make_engine();
let request = make_request("Stream test");
let stream = engine.infer_stream(request).await.unwrap();
let chunks: Vec<_> = stream
.collect::<Vec<_>>()
.await
.into_iter()
.collect::<Result<Vec<_>, _>>()
.unwrap();
assert!(!chunks.is_empty(), "Stream should produce chunks");
let last = chunks.last().unwrap();
assert!(
last.finish_reason.is_some(),
"Final chunk should have finish_reason"
);
}
#[tokio::test]
async fn non_streaming_usage_matches_tokenizer_and_generated_tokens() {
let engine = make_engine();
let tokenizer = MockTokenizer::new(VOCAB_SIZE);
let prompt = "usage fixture prompt";
let expected_prompt_tokens = tokenizer.encode(prompt, true).unwrap().len();
let response = engine.infer(make_request(prompt)).await.unwrap();
assert_eq!(response.usage.prompt_tokens, expected_prompt_tokens);
assert_eq!(response.usage.completion_tokens, response.tokens.len());
assert_eq!(
response.usage.total_tokens,
response.usage.prompt_tokens + response.usage.completion_tokens
);
}
#[tokio::test]
async fn streaming_final_usage_matches_tokenizer_and_emitted_tokens() {
use futures::StreamExt;
let engine = make_engine();
let tokenizer = MockTokenizer::new(VOCAB_SIZE);
let prompt = "stream usage fixture";
let expected_prompt_tokens = tokenizer.encode(prompt, true).unwrap().len();
let stream = engine.infer_stream(make_request(prompt)).await.unwrap();
let chunks: Vec<_> = stream
.collect::<Vec<_>>()
.await
.into_iter()
.collect::<Result<Vec<_>, _>>()
.unwrap();
let emitted_tokens = chunks.iter().filter(|chunk| chunk.token.is_some()).count();
let final_usage = chunks
.last()
.and_then(|chunk| chunk.usage.as_ref())
.expect("final stream chunk should include usage");
assert_eq!(final_usage.prompt_tokens, expected_prompt_tokens);
assert_eq!(final_usage.completion_tokens, emitted_tokens);
assert_eq!(
final_usage.total_tokens,
final_usage.prompt_tokens + final_usage.completion_tokens
);
}
#[tokio::test]
async fn concurrent_submit_tracked_by_scheduler() {
let scheduler = Arc::new(ContinuousBatchScheduler::new(SchedulerConfig::default()));
for i in 0..3 {
let req = make_request(&format!("Request {}", i));
scheduler.submit(req).await.unwrap();
}
let metrics = scheduler.metrics();
assert_eq!(metrics.waiting_requests, 3);
}
#[tokio::test]
async fn kv_cache_allocated_and_deallocated() {
let kv_cache = Arc::new(MockKvCacheManager::new(1024));
let config = ferrum_types::EngineConfig::default();
let scheduler = Arc::new(ContinuousBatchScheduler::new(SchedulerConfig::default()));
let tokenizer = Arc::new(MockTokenizer::new(VOCAB_SIZE));
let sampler = Arc::new(MockSampler);
let executor = Arc::new(MockModelExecutor::instant(VOCAB_SIZE));
let tensor_factory = Arc::new(MockTensorFactory);
let engine = ContinuousBatchEngine::new(
config,
scheduler,
tokenizer,
sampler,
kv_cache.clone(),
executor,
tensor_factory,
)
.expect("legacy engine composition must match executor authority");
assert_eq!(kv_cache.active_count(), 0);
let request = make_request("KV cache test");
let _response = engine.infer(request).await.unwrap();
assert_eq!(kv_cache.active_count(), 0);
}
#[tokio::test]
async fn mock_executor_tracks_operations() {
let executor = Arc::new(MockModelExecutor::instant(VOCAB_SIZE));
let config = ferrum_types::EngineConfig::default();
let scheduler = Arc::new(ContinuousBatchScheduler::new(SchedulerConfig::default()));
let tokenizer = Arc::new(MockTokenizer::new(VOCAB_SIZE));
let sampler = Arc::new(MockSampler);
let kv_cache = Arc::new(MockKvCacheManager::new(1024));
let tensor_factory = Arc::new(MockTensorFactory);
let engine = ContinuousBatchEngine::new(
config,
scheduler,
tokenizer,
sampler,
kv_cache,
executor.clone(),
tensor_factory,
)
.expect("legacy engine composition must match executor authority");
assert_eq!(executor.prefill_count(), 0);
assert_eq!(executor.decode_count(), 0);
let request = make_request("Track ops");
let _response = engine.infer(request).await.unwrap();
assert_eq!(executor.prefill_count(), 1);
assert_eq!(executor.decode_count(), 4); }
#[tokio::test]
async fn engine_with_latency_still_completes() {
let config = ferrum_types::EngineConfig::default();
let scheduler = Arc::new(ContinuousBatchScheduler::new(SchedulerConfig::default()));
let tokenizer = Arc::new(MockTokenizer::new(VOCAB_SIZE));
let sampler = Arc::new(MockSampler);
let kv_cache = Arc::new(MockKvCacheManager::new(1024));
let executor = Arc::new(MockModelExecutor::new(
VOCAB_SIZE,
Duration::from_millis(5),
Duration::from_millis(2),
));
let tensor_factory = Arc::new(MockTensorFactory);
let engine = ContinuousBatchEngine::new(
config,
scheduler,
tokenizer,
sampler,
kv_cache,
executor,
tensor_factory,
)
.expect("legacy engine composition must match executor authority");
let request = make_request("Latency test");
let response = engine.infer(request).await.unwrap();
assert_eq!(response.finish_reason, ferrum_types::FinishReason::Length);
assert!(response.latency_ms >= 10);
}
fn make_engine_shared() -> Arc<ContinuousBatchEngine> {
Arc::new(make_engine())
}
#[tokio::test]
async fn concurrent_requests_all_complete() {
let engine = make_engine_shared();
let mut handles = Vec::new();
for i in 0..5 {
let e = engine.clone();
handles.push(tokio::spawn(async move {
let req = make_request(&format!("Concurrent request {}", i));
e.infer(req).await
}));
}
let results: Vec<InferenceResponse> = futures::future::join_all(handles)
.await
.into_iter()
.map(|r| r.unwrap().unwrap())
.collect();
assert_eq!(results.len(), 5);
for resp in &results {
assert_eq!(resp.finish_reason, ferrum_types::FinishReason::Length);
assert!(!resp.tokens.is_empty());
}
}
#[tokio::test]
async fn concurrent_requests_deallocate_kv() {
let kv_cache = Arc::new(MockKvCacheManager::new(1024));
let config = ferrum_types::EngineConfig::default();
let scheduler = Arc::new(ContinuousBatchScheduler::new(SchedulerConfig::default()));
let tokenizer = Arc::new(MockTokenizer::new(VOCAB_SIZE));
let sampler = Arc::new(MockSampler);
let executor = Arc::new(MockModelExecutor::instant(VOCAB_SIZE));
let tensor_factory = Arc::new(MockTensorFactory);
let engine = Arc::new(
ContinuousBatchEngine::new(
config,
scheduler,
tokenizer,
sampler,
kv_cache.clone(),
executor,
tensor_factory,
)
.expect("legacy engine composition must match executor authority"),
);
assert_eq!(kv_cache.active_count(), 0);
let mut handles = Vec::new();
for i in 0..3 {
let e = engine.clone();
handles.push(tokio::spawn(async move {
let req = make_request(&format!("KV test {}", i));
e.infer(req).await.unwrap()
}));
}
futures::future::join_all(handles).await;
assert_eq!(kv_cache.active_count(), 0, "All KV caches should be freed");
}
#[tokio::test]
async fn concurrent_streams_all_complete() {
use futures::StreamExt;
let engine = make_engine_shared();
let mut handles = Vec::new();
for i in 0..3 {
let e = engine.clone();
handles.push(tokio::spawn(async move {
let req = make_request(&format!("Stream {}", i));
let stream = e.infer_stream(req).await.unwrap();
let chunks: Vec<_> = stream
.collect::<Vec<_>>()
.await
.into_iter()
.collect::<Result<Vec<_>, _>>()
.unwrap();
chunks
}));
}
let all_chunks: Vec<Vec<_>> = futures::future::join_all(handles)
.await
.into_iter()
.map(|r| r.unwrap())
.collect();
assert_eq!(all_chunks.len(), 3);
for chunks in &all_chunks {
assert!(!chunks.is_empty(), "Each stream should produce chunks");
let last = chunks.last().unwrap();
assert!(
last.finish_reason.is_some(),
"Final chunk should have finish_reason"
);
}
}
#[ignore = "prefix cache defaults OFF; opt in with FERRUM_PREFIX_CACHE=1 + --ignored"]
#[tokio::test]
async fn prefix_cache_avoids_second_prefill() {
let executor = Arc::new(MockModelExecutor::instant(VOCAB_SIZE));
let config = ferrum_types::EngineConfig::default();
let scheduler = Arc::new(ContinuousBatchScheduler::new(SchedulerConfig::default()));
let tokenizer = Arc::new(MockTokenizer::new(VOCAB_SIZE));
let sampler = Arc::new(MockSampler);
let kv_cache = Arc::new(MockKvCacheManager::new(1024));
let tensor_factory = Arc::new(MockTensorFactory);
let engine = ContinuousBatchEngine::new(
config,
scheduler,
tokenizer,
sampler,
kv_cache,
executor.clone(),
tensor_factory,
)
.expect("legacy engine composition must match executor authority");
let req1 = make_request("Identical prompt for prefix cache");
let resp1 = engine.infer(req1).await.unwrap();
assert_eq!(executor.prefill_count(), 1);
let req2 = make_request("Identical prompt for prefix cache");
let resp2 = engine.infer(req2).await.unwrap();
assert_eq!(
executor.prefill_count(),
1,
"Prefix cache should skip second prefill"
);
assert_eq!(resp1.tokens, resp2.tokens);
assert!(!resp1.tokens.is_empty());
}
#[tokio::test]
async fn json_mode_biases_first_token() {
use ferrum_types::ResponseFormat;
let engine = make_engine_with_tokenizer(Arc::new(JsonObjectTokenizer::new()));
let mut plain_req = make_request("Hello");
plain_req.sampling_params.max_tokens = 1;
let plain_resp = engine.infer(plain_req).await.unwrap();
assert_eq!(
plain_resp.tokens[0].get(),
42,
"Without JSON mode, greedy should pick token 42"
);
let mut json_req = make_request("Hello");
json_req.sampling_params.max_tokens = 1;
json_req.sampling_params.response_format = ResponseFormat::JsonObject;
let json_resp = engine.infer(json_req).await.unwrap();
assert_eq!(
json_resp.tokens[0].get(),
JSON_OBJECT_TOKEN_ID,
"JSON mode should select the complete JSON object token"
);
assert_eq!(json_resp.text, "{}");
}