mod common;
use std::net::SocketAddr;
use modelc::runtime::serve::Runtime;
use modelc::runtime::transformer;
use modelc::serve::run_server;
const GPT2_HIDDEN: usize = 12;
const GPT2_VOCAB: usize = 10;
const LLAMA_HIDDEN: usize = 12;
const LLAMA_VOCAB: usize = 10;
#[test]
fn forward_gpt2_produces_vocab_sized_logits() {
let model = common::create_gpt2_test_model();
let runtime = Runtime::from_raw(&model.tensors);
let input = vec![0.1f32; GPT2_HIDDEN];
let out = transformer::forward_gpt2(&runtime, &input, false).expect("gpt2 forward");
assert_eq!(out.len(), GPT2_VOCAB);
assert!(out.iter().all(|v| v.is_finite()), "non-finite logits");
assert_ne!(out, input);
assert!(out.iter().any(|v| *v != 0.0), "output is all zeros");
}
#[test]
fn forward_llama_produces_vocab_sized_logits() {
let model = common::create_llama_test_model();
let runtime = Runtime::from_raw(&model.tensors);
let input = vec![0.05f32; LLAMA_HIDDEN];
let out = transformer::forward_llama(&runtime, &input, false).expect("llama forward");
assert_eq!(out.len(), LLAMA_VOCAB);
assert!(out.iter().all(|v| v.is_finite()));
assert_ne!(out, input);
assert!(out.iter().any(|v| *v != 0.0));
}
#[test]
fn forward_gpt2_is_deterministic() {
let model = common::create_gpt2_test_model();
let runtime = Runtime::from_raw(&model.tensors);
let input = vec![0.2f32; GPT2_HIDDEN];
let a = transformer::forward_gpt2(&runtime, &input, false).unwrap();
let b = transformer::forward_gpt2(&runtime, &input, false).unwrap();
assert_eq!(a, b);
}
#[test]
fn forward_is_input_sensitive() {
let model = common::create_gpt2_test_model();
let runtime = Runtime::from_raw(&model.tensors);
let out_a = transformer::forward_gpt2(&runtime, &[0.1f32; GPT2_HIDDEN], false).unwrap();
let out_b = transformer::forward_gpt2(&runtime, &[0.9f32; GPT2_HIDDEN], false).unwrap();
assert_ne!(out_a, out_b, "forward must respond to input changes");
}
#[test]
fn forward_gpt2_returns_none_without_output_head() {
let model = common::create_gpt2_test_model();
let mut tensors = model.tensors.clone();
tensors.remove("transformer.wte.weight");
let runtime = Runtime::from_raw(&tensors);
let input = vec![0.0f32; GPT2_HIDDEN];
assert!(transformer::forward_gpt2(&runtime, &input, false).is_none());
}
#[test]
fn forward_llama_returns_none_without_output_head() {
let model = common::create_llama_test_model();
let mut tensors = model.tensors.clone();
tensors.remove("lm_head.weight");
tensors.remove("model.embed_tokens.weight");
let runtime = Runtime::from_raw(&tensors);
let input = vec![0.0f32; LLAMA_HIDDEN];
assert!(transformer::forward_llama(&runtime, &input, false).is_none());
}
fn ephemeral_addr() -> SocketAddr {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
listener.local_addr().unwrap()
}
fn wait_for_server(url: &str) -> serde_json::Value {
for _ in 0..50 {
if let Ok(resp) = ureq::get(url).call()
&& let Ok(bytes) = read_body(resp)
&& let Ok(val) = serde_json::from_slice::<serde_json::Value>(&bytes)
{
return val;
}
std::thread::sleep(std::time::Duration::from_millis(20));
}
panic!("server never came up at {url}");
}
fn read_body(resp: ureq::http::Response<ureq::Body>) -> std::io::Result<Vec<u8>> {
use std::io::Read;
let mut buf = Vec::new();
resp.into_body().into_reader().read_to_end(&mut buf)?;
Ok(buf)
}
fn post_json(url: &str, body: &serde_json::Value) -> serde_json::Value {
let body_str = serde_json::to_string(body).unwrap();
let resp = ureq::post(url)
.content_type("application/json")
.send(&body_str)
.expect("POST failed");
let bytes = read_body(resp).expect("read body");
serde_json::from_slice(&bytes).expect("parse json")
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_serves_gpt2_inference() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), None).await;
});
let info = wait_for_server(&format!("{base}/info"));
assert_eq!(info["architecture"], "gpt2");
let body = serde_json::json!({ "input": vec![0.1f32; GPT2_HIDDEN] });
let val = post_json(&format!("{base}/infer"), &body);
let out = val["output"].as_array().expect("output array");
assert_eq!(out.len(), GPT2_VOCAB, "expected vocab-sized logits");
assert!(out.iter().any(|v| v.as_f64().unwrap_or(0.0) != 0.0));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_resizes_oversized_infer_input() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), None).await;
});
wait_for_server(&format!("{base}/info"));
let oversized: Vec<f32> = vec![0.3; GPT2_HIDDEN * 4];
let body = serde_json::json!({ "input": oversized });
let val = post_json(&format!("{base}/infer"), &body);
assert_eq!(val["output"].as_array().unwrap().len(), GPT2_VOCAB);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_serves_health() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), None).await;
});
wait_for_server(&format!("{base}/info"));
let val = wait_for_server(&format!("{base}/health"));
assert_eq!(val["status"], "ok");
assert_eq!(val["model"], "mini_gpt2");
assert_eq!(val["architecture"], "gpt2");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_serves_openai_models() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), None).await;
});
wait_for_server(&format!("{base}/info"));
let val = wait_for_server(&format!("{base}/v1/models"));
assert_eq!(val["object"], "list");
let data = val["data"].as_array().expect("data array");
assert_eq!(data.len(), 1);
assert_eq!(data[0]["object"], "model");
assert_eq!(data[0]["owned_by"], "modelc");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_serves_embeddings() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), None).await;
});
wait_for_server(&format!("{base}/info"));
let body = serde_json::json!({ "input": "hello world" });
let val = post_json(&format!("{base}/embeddings"), &body);
let embedding = val["embedding"].as_array().expect("embedding array");
assert_eq!(
embedding.len(),
GPT2_HIDDEN,
"embedding must be hidden-size"
);
assert!(embedding.iter().any(|v| v.as_f64().unwrap_or(0.0) != 0.0));
assert_eq!(val["model"], "mini_gpt2");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_serves_openai_chat_completions() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), None).await;
});
wait_for_server(&format!("{base}/info"));
let body = serde_json::json!({
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hello"}]
});
let val = post_json(&format!("{base}/v1/chat/completions"), &body);
assert_eq!(val["object"], "chat.completion");
let choices = val["choices"].as_array().expect("choices array");
assert_eq!(choices.len(), 1);
assert_eq!(choices[0]["message"]["role"], "assistant");
assert!(choices[0]["message"]["content"].as_str().is_some());
assert!(choices[0]["logprobs"].is_null());
assert!(val["usage"]["total_tokens"].as_u64().unwrap_or(0) > 0);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_serves_openai_chat_logprobs() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), None).await;
});
wait_for_server(&format!("{base}/info"));
let body = serde_json::json!({
"model": "mini_gpt2",
"messages": [{"role": "user", "content": "hello"}],
"logprobs": true,
"top_logprobs": 3,
"max_tokens": 5
});
let val = post_json(&format!("{base}/v1/chat/completions"), &body);
let logprobs = &val["choices"][0]["logprobs"];
assert!(
logprobs.is_object(),
"logprobs should be an object when requested"
);
let content = logprobs["content"].as_array().expect("content array");
assert!(
!content.is_empty(),
"should produce at least one token's logprobs"
);
assert!(content.len() <= 5, "should respect max_tokens");
for entry in content {
assert!(entry["token"].is_string(), "token is a string");
assert!(entry["bytes"].is_array(), "bytes is an array");
let lp = entry["logprob"].as_f64().expect("logprob is a number");
assert!(
lp <= 0.0,
"logprob must be <= 0.0 (it is ln of a probability)"
);
let top = entry["top_logprobs"]
.as_array()
.expect("top_logprobs array");
assert!(
top.len() <= 3,
"top_logprobs must respect the requested count"
);
let mut prev = f64::INFINITY;
for alt in top {
assert!(alt["token"].is_string());
assert!(alt["bytes"].is_array());
let p = alt["logprob"].as_f64().unwrap();
assert!(p <= 0.0);
assert!(p <= prev + 1e-6, "top_logprobs must be sorted descending");
prev = p;
}
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_serves_openai_chat_logprobs_no_top() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), None).await;
});
wait_for_server(&format!("{base}/info"));
let body = serde_json::json!({
"messages": [{"role": "user", "content": "hi"}],
"logprobs": true,
"max_tokens": 2
});
let val = post_json(&format!("{base}/v1/chat/completions"), &body);
let content = val["choices"][0]["logprobs"]["content"]
.as_array()
.expect("content array");
for entry in content {
assert!(
entry["top_logprobs"]
.as_array()
.map(|a| a.is_empty())
.unwrap_or(true),
"top_logprobs must be empty when not requested"
);
assert!(entry["logprob"].as_f64().is_some());
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_lora_unload_restores_base() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), None).await;
});
wait_for_server(&format!("{base}/info"));
let body = serde_json::json!({});
let val = post_json(&format!("{base}/lora/unload"), &body);
assert!(val["message"].as_str().unwrap().contains("unloaded"));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_lora_load_bad_path_returns_error() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), None).await;
});
wait_for_server(&format!("{base}/info"));
let body = serde_json::json!({"path": "/nonexistent/lora.safetensors"});
let val = post_json(&format!("{base}/lora/load"), &body);
assert!(val["message"].as_str().unwrap().contains("Failed"));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_metrics_returns_prometheus_text() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), None).await;
});
wait_for_server(&format!("{base}/info"));
let chat_body = serde_json::json!({
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 2
});
let _ = post_json(&format!("{base}/chat"), &chat_body);
let mut resp = ureq::get(&format!("{base}/metrics"))
.call()
.expect("metrics request");
let text = resp.body_mut().read_to_string().expect("read metrics body");
assert!(text.contains("modelc_requests_total"), "should expose request counter");
assert!(text.contains("modelc_inference_duration_seconds_count"), "should expose histogram count");
assert!(text.contains("modelc_active_requests"), "should expose active request gauge");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_chat_accepts_json_schema() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), None).await;
});
wait_for_server(&format!("{base}/info"));
let body = serde_json::json!({
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 4,
"json_schema": {"type": "object"}
});
let val = post_json(&format!("{base}/chat"), &body);
assert!(val["message"].is_object(), "response should have message field");
assert!(
val["message"]["content"].as_str().is_some(),
"response should have content string"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_serves_openai_chat_stream() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), None).await;
});
wait_for_server(&format!("{base}/info"));
let body = serde_json::json!({
"model": "mini_gpt2",
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 3,
"stream": true
});
let body_str = serde_json::to_string(&body).unwrap();
let resp = ureq::post(&format!("{base}/v1/chat/completions"))
.content_type("application/json")
.send(&body_str)
.expect("request should succeed");
assert_eq!(resp.status(), 200);
let text = read_body(resp).map(|b| String::from_utf8_lossy(&b).into_owned()).unwrap_or_default();
assert!(text.contains("chat.completion.chunk"), "should emit OpenAI chunk objects");
assert!(text.contains("[DONE]"), "should end with [DONE]");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_serves_openai_completions() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), None).await;
});
wait_for_server(&format!("{base}/info"));
let body = serde_json::json!({
"model": "mini_gpt2",
"prompt": "hello",
"max_tokens": 4
});
let val = post_json(&format!("{base}/v1/completions"), &body);
assert_eq!(val["object"].as_str(), Some("text_completion"));
assert!(val["choices"].is_array(), "should have choices array");
assert_eq!(val["choices"][0]["index"].as_i64(), Some(0));
assert!(val["choices"][0]["text"].as_str().is_some(), "should have text");
assert_eq!(val["choices"][0]["finish_reason"].as_str(), Some("stop"));
assert!(val["usage"]["prompt_tokens"].is_number());
assert!(val["usage"]["completion_tokens"].is_number());
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_api_key_rejects_unauthenticated() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
let auth = modelc::serve::auth::AuthConfig::new(Some("secret".to_string()), None);
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), Some(auth)).await;
});
wait_for_server(&format!("{base}/info"));
let status: u16 = match ureq::get(&format!("{base}/metrics")).call() {
Ok(r) => r.status().as_u16(),
Err(ureq::Error::StatusCode(code)) => code,
Err(_) => 0,
};
assert_eq!(status, 401, "unauthenticated request should be rejected");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_api_key_accepts_authenticated() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
let auth = modelc::serve::auth::AuthConfig::new(Some("secret".to_string()), None);
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), Some(auth)).await;
});
wait_for_server(&format!("{base}/info"));
let resp = ureq::get(&format!("{base}/metrics"))
.header("Authorization", "Bearer secret")
.call()
.expect("request should succeed");
assert_eq!(resp.status(), 200, "authenticated request should pass through");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn run_server_rate_limit_rejects_over_limit() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
let auth = modelc::serve::auth::AuthConfig::new(None, Some(1));
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), Some(auth)).await;
});
wait_for_server(&format!("{base}/info"));
let resp1 = ureq::get(&format!("{base}/metrics"))
.call()
.expect("first request should succeed");
assert_eq!(resp1.status(), 200, "first request should pass");
let status: u16 = match ureq::get(&format!("{base}/metrics")).call() {
Ok(r) => r.status().as_u16(),
Err(ureq::Error::StatusCode(code)) => code,
Err(_) => 0,
};
assert_eq!(status, 429, "second request should be rate limited");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn run_server_handles_concurrent_transformer_requests() {
let model = common::create_gpt2_test_model();
let addr = ephemeral_addr();
let base = format!("http://{addr}");
tokio::spawn(async move {
let _ = run_server(model, addr, false, modelc::generate::GenerationConfig::default(), None).await;
});
wait_for_server(&format!("{base}/info"));
let body_a = serde_json::json!({
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 4
});
let body_b = serde_json::json!({
"messages": [{"role": "user", "content": "world"}],
"max_tokens": 4
});
let base_a = base.clone();
let base_b = base.clone();
let (resp_a, resp_b) = tokio::join!(
tokio::task::spawn_blocking(move || {
ureq::post(&format!("{base_a}/chat"))
.content_type("application/json")
.send(&serde_json::to_string(&body_a).unwrap())
.expect("concurrent request A should succeed")
}),
tokio::task::spawn_blocking(move || {
ureq::post(&format!("{base_b}/chat"))
.content_type("application/json")
.send(&serde_json::to_string(&body_b).unwrap())
.expect("concurrent request B should succeed")
}),
);
let resp_a = resp_a.expect("spawn A should not panic");
let resp_b = resp_b.expect("spawn B should not panic");
assert_eq!(resp_a.status(), 200, "concurrent request A should return 200");
assert_eq!(resp_b.status(), 200, "concurrent request B should return 200");
let val_a: serde_json::Value = serde_json::from_reader(resp_a.into_body().as_reader()).expect("A should be JSON");
let val_b: serde_json::Value = serde_json::from_reader(resp_b.into_body().as_reader()).expect("B should be JSON");
assert!(val_a["message"]["content"].as_str().is_some(), "A should have content");
assert!(val_b["message"]["content"].as_str().is_some(), "B should have content");
}