use ferrum_engine::{ContinuousBatchEngine, LlmInferenceEngine};
use ferrum_scheduler::implementations::ContinuousBatchScheduler;
use ferrum_testkit::{
ConfigurableModelExecutor, MockKvCacheManager, MockModelExecutor, MockSampler,
MockTensorFactory, MockTokenizer,
};
use ferrum_types::{FinishReason, InferenceRequest, SchedulerConfig};
use std::sync::Arc;
const VOCAB: usize = 1000;
fn make_engine_with_executor(
executor: Arc<dyn ferrum_interfaces::ModelExecutor>,
) -> ContinuousBatchEngine {
let config = ferrum_types::EngineConfig::default();
let scheduler = Arc::new(ContinuousBatchScheduler::new(SchedulerConfig::default()));
let tokenizer = Arc::new(MockTokenizer::new(VOCAB));
let sampler = Arc::new(MockSampler);
let kv_cache = Arc::new(MockKvCacheManager::new(1024));
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")
}
fn make_engine() -> ContinuousBatchEngine {
make_engine_with_executor(Arc::new(MockModelExecutor::instant(VOCAB)))
}
fn make_request(prompt: &str, max_tokens: usize) -> InferenceRequest {
let mut req = InferenceRequest::new(prompt, "mock-model");
req.sampling_params.max_tokens = max_tokens;
req.sampling_params.temperature = 0.0;
req
}
#[tokio::test]
async fn greedy_determinism_same_output() {
let engine = make_engine();
let r1 = engine.infer(make_request("Hello", 10)).await.unwrap();
let r2 = engine.infer(make_request("Hello", 10)).await.unwrap();
assert_eq!(r1.text, r2.text, "same prompt + greedy = same output");
assert_eq!(r1.tokens.len(), r2.tokens.len());
}
#[tokio::test]
async fn different_prompts_different_output() {
let engine = make_engine();
let r1 = engine.infer(make_request("Hello", 5)).await.unwrap();
let r2 = engine.infer(make_request("Goodbye", 5)).await.unwrap();
assert_eq!(r1.tokens.len(), 5);
assert_eq!(r2.tokens.len(), 5);
}
#[tokio::test]
async fn max_tokens_1_generates_one_token() {
let engine = make_engine();
let resp = engine.infer(make_request("Hi", 1)).await.unwrap();
assert_eq!(resp.tokens.len(), 1);
assert!(
resp.finish_reason == FinishReason::Length || resp.finish_reason == FinishReason::EOS,
"unexpected finish_reason: {:?}",
resp.finish_reason
);
}
#[tokio::test]
async fn max_tokens_exact_count() {
for n in [1, 5, 10, 20, 50] {
let engine = make_engine();
let resp = engine.infer(make_request("Test", n)).await.unwrap();
assert_eq!(
resp.tokens.len(),
n,
"max_tokens={n} but got {}",
resp.tokens.len()
);
}
}
#[tokio::test]
async fn eos_token_stops_generation() {
let executor = Arc::new(ConfigurableModelExecutor::with_eos_after(VOCAB, 3, 2));
let engine = make_engine_with_executor(executor);
let mut req = make_request("Hello", 100);
req.sampling_params.temperature = 0.0;
let resp = engine.infer(req).await.unwrap();
assert!(
resp.tokens.len() <= 100,
"should not exceed max_tokens, got {} tokens",
resp.tokens.len()
);
}
#[tokio::test]
async fn configurable_executor_produces_expected_tokens() {
let executor = Arc::new(ConfigurableModelExecutor::with_token_sequence(
VOCAB,
vec![10, 20, 30],
));
let engine = make_engine_with_executor(executor);
let resp = engine.infer(make_request("Test", 6)).await.unwrap();
assert_eq!(resp.tokens.len(), 6);
assert!(!resp.text.is_empty());
}
#[tokio::test]
async fn single_token_prompt() {
let engine = make_engine();
let resp = engine.infer(make_request("A", 5)).await.unwrap();
assert_eq!(resp.tokens.len(), 5);
}
#[tokio::test]
async fn unicode_prompt_no_panic() {
let engine = make_engine();
let resp = engine
.infer(make_request("你好世界 🚀 こんにちは", 5))
.await
.unwrap();
assert_eq!(resp.tokens.len(), 5);
}
#[tokio::test]
async fn long_prompt_completes() {
let engine = make_engine();
let long_prompt = "word ".repeat(1000); let resp = engine.infer(make_request(&long_prompt, 5)).await.unwrap();
assert_eq!(resp.tokens.len(), 5);
}
#[tokio::test]
async fn empty_prompt_handled() {
let engine = make_engine();
let result = engine.infer(make_request("", 5)).await;
match result {
Ok(resp) => assert!(resp.tokens.len() <= 5),
Err(_) => {} }
}