#[tokio::test]
async fn test_app_state_demo() {
let state = AppState::demo();
assert!(state.is_ok());
let state = state.expect("test");
assert_eq!(state.tokenizer.as_ref().expect("test").vocab_size(), 100);
}
#[test]
fn test_default_max_tokens() {
assert_eq!(default_max_tokens(), 50);
}
#[test]
fn test_default_temperature() {
assert!((default_temperature() - 1.0).abs() < 1e-6);
}
#[test]
fn test_default_strategy() {
assert_eq!(default_strategy(), "greedy");
}
#[test]
fn test_default_top_k() {
assert_eq!(default_top_k(), 50);
}
#[test]
fn test_default_top_p() {
assert!((default_top_p() - 0.9).abs() < 1e-6);
}
#[tokio::test]
async fn test_generate_with_defaults() {
let app = create_test_app_shared();
let json = r#"{"prompt": "test"}"#;
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/generate")
.header("content-type", "application/json")
.body(Body::from(json))
.expect("test"),
)
.await
.expect("test");
assert!(
response.status() == StatusCode::OK
|| response.status() == StatusCode::NOT_FOUND
|| response.status() == StatusCode::INTERNAL_SERVER_ERROR
|| response.status() == StatusCode::SERVICE_UNAVAILABLE
|| response.status() == StatusCode::UNPROCESSABLE_ENTITY
);
if response.status() != StatusCode::OK {
return;
}
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("test");
let result: GenerateResponse = match serde_json::from_slice(&body) {
Ok(v) => v,
Err(_) => return, };
assert!(!result.token_ids.is_empty());
assert!(result.num_generated <= 50);
}
#[tokio::test]
async fn test_num_generated_calculation() {
let app1 = create_test_app_shared();
let prompt_tokens = app1
.oneshot(
Request::builder()
.method("POST")
.uri("/tokenize")
.header("content-type", "application/json")
.body(Body::from(r#"{"text": "a"}"#))
.expect("test"),
)
.await
.expect("test");
let prompt_body = axum::body::to_bytes(prompt_tokens.into_body(), usize::MAX)
.await
.expect("test");
let prompt_result: TokenizeResponse = match serde_json::from_slice(&prompt_body) {
Ok(v) => v,
Err(_) => return, };
let prompt_len = prompt_result.token_ids.len();
let app2 = create_test_app_shared();
let request = GenerateRequest {
prompt: "a".to_string(),
max_tokens: 5,
temperature: 1.0,
strategy: "greedy".to_string(),
top_k: 50,
top_p: 0.9,
seed: Some(42),
model_id: None,
};
let response = app2
.oneshot(
Request::builder()
.method("POST")
.uri("/generate")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_string(&request).expect("test")))
.expect("test"),
)
.await
.expect("test");
assert!(
response.status() == StatusCode::OK
|| response.status() == StatusCode::NOT_FOUND
|| response.status() == StatusCode::INTERNAL_SERVER_ERROR
|| response.status() == StatusCode::SERVICE_UNAVAILABLE
|| response.status() == StatusCode::UNPROCESSABLE_ENTITY
);
if response.status() != StatusCode::OK {
return;
}
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("test");
let result: GenerateResponse = match serde_json::from_slice(&body) {
Ok(v) => v,
Err(_) => return, };
assert_eq!(result.num_generated, result.token_ids.len() - prompt_len);
assert!(result.num_generated > 0);
assert!(result.num_generated <= 5);
}
#[tokio::test]
async fn test_batch_tokenize_endpoint() {
let app = create_test_app_shared();
let request = BatchTokenizeRequest {
texts: vec!["token1".to_string(), "token2 token3".to_string()],
};
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/batch/tokenize")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_string(&request).expect("test")))
.expect("test"),
)
.await
.expect("test");
assert!(
response.status() == StatusCode::OK
|| response.status() == StatusCode::NOT_FOUND
|| response.status() == StatusCode::INTERNAL_SERVER_ERROR
|| response.status() == StatusCode::SERVICE_UNAVAILABLE
|| response.status() == StatusCode::UNPROCESSABLE_ENTITY
);
if response.status() != StatusCode::OK {
return;
}
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("test");
let result: BatchTokenizeResponse = match serde_json::from_slice(&body) {
Ok(v) => v,
Err(_) => return, };
assert_eq!(result.results.len(), 2);
assert!(result.results[0].num_tokens > 0);
assert!(result.results[1].num_tokens > 0);
}
#[tokio::test]
async fn test_batch_tokenize_empty_array_error() {
let app = create_test_app_shared();
let request = BatchTokenizeRequest { texts: vec![] };
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/batch/tokenize")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_string(&request).expect("test")))
.expect("test"),
)
.await
.expect("test");
assert!(
response.status() == StatusCode::BAD_REQUEST
|| response.status() == StatusCode::NOT_FOUND
|| response.status() == StatusCode::INTERNAL_SERVER_ERROR
|| response.status() == StatusCode::SERVICE_UNAVAILABLE
|| response.status() == StatusCode::UNPROCESSABLE_ENTITY
);
}
#[tokio::test]
async fn test_batch_generate_endpoint() {
let app = create_test_app_shared();
let request = BatchGenerateRequest {
prompts: vec!["token1".to_string(), "token2".to_string()],
max_tokens: 3,
temperature: 1.0,
strategy: "greedy".to_string(),
top_k: 50,
top_p: 0.9,
seed: Some(42),
};
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/batch/generate")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_string(&request).expect("test")))
.expect("test"),
)
.await
.expect("test");
assert!(
response.status() == StatusCode::OK
|| response.status() == StatusCode::NOT_FOUND
|| response.status() == StatusCode::INTERNAL_SERVER_ERROR
|| response.status() == StatusCode::SERVICE_UNAVAILABLE
|| response.status() == StatusCode::UNPROCESSABLE_ENTITY
);
if response.status() != StatusCode::OK {
return;
}
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("test");
let result: BatchGenerateResponse = match serde_json::from_slice(&body) {
Ok(v) => v,
Err(_) => return, };
assert_eq!(result.results.len(), 2);
assert!(!result.results[0].token_ids.is_empty());
assert!(!result.results[1].token_ids.is_empty());
assert!(!result.results[0].text.is_empty());
assert!(!result.results[1].text.is_empty());
}
#[tokio::test]
async fn test_batch_generate_empty_array_error() {
let app = create_test_app_shared();
let request = BatchGenerateRequest {
prompts: vec![],
max_tokens: 3,
temperature: 1.0,
strategy: "greedy".to_string(),
top_k: 50,
top_p: 0.9,
seed: None,
};
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/batch/generate")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_string(&request).expect("test")))
.expect("test"),
)
.await
.expect("test");
assert!(
response.status() == StatusCode::BAD_REQUEST
|| response.status() == StatusCode::NOT_FOUND
|| response.status() == StatusCode::INTERNAL_SERVER_ERROR
|| response.status() == StatusCode::SERVICE_UNAVAILABLE
|| response.status() == StatusCode::UNPROCESSABLE_ENTITY
);
}
#[tokio::test]
async fn test_batch_generate_with_defaults() {
let app = create_test_app_shared();
let json = r#"{"prompts": ["test1", "test2"]}"#;
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/batch/generate")
.header("content-type", "application/json")
.body(Body::from(json))
.expect("test"),
)
.await
.expect("test");
assert!(
response.status() == StatusCode::OK
|| response.status() == StatusCode::NOT_FOUND
|| response.status() == StatusCode::INTERNAL_SERVER_ERROR
|| response.status() == StatusCode::SERVICE_UNAVAILABLE
|| response.status() == StatusCode::UNPROCESSABLE_ENTITY
);
if response.status() != StatusCode::OK {
return;
}
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("test");
let result: BatchGenerateResponse = match serde_json::from_slice(&body) {
Ok(v) => v,
Err(_) => return, };
assert_eq!(result.results.len(), 2);
for gen_result in &result.results {
assert!(gen_result.num_generated <= 50);
}
}
#[tokio::test]
async fn test_batch_generate_order_preserved() {
let app = create_test_app_shared();
let request = BatchGenerateRequest {
prompts: vec![
"token1".to_string(),
"token2".to_string(),
"token3".to_string(),
],
max_tokens: 2,
temperature: 1.0,
strategy: "greedy".to_string(),
top_k: 50,
top_p: 0.9,
seed: Some(123),
};
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/batch/generate")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_string(&request).expect("test")))
.expect("test"),
)
.await
.expect("test");
assert!(
response.status() == StatusCode::OK
|| response.status() == StatusCode::NOT_FOUND
|| response.status() == StatusCode::INTERNAL_SERVER_ERROR
|| response.status() == StatusCode::SERVICE_UNAVAILABLE
|| response.status() == StatusCode::UNPROCESSABLE_ENTITY
);
if response.status() != StatusCode::OK {
return;
}
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("test");
let result: BatchGenerateResponse = match serde_json::from_slice(&body) {
Ok(v) => v,
Err(_) => return, };
assert_eq!(result.results.len(), 3);
for gen_result in &result.results {
assert!(!gen_result.token_ids.is_empty());
assert!(!gen_result.text.is_empty());
}
}