feat(api): add health check and models endpoints with documentation
- Add /health endpoint to check service status for orchestration systems - Add /models endpoint with OpenAI API compatible format - Document new endpoints in both English and Chinese API docs - Include example usage and response formats in documentation - Add comprehensive test coverage for health and models endpoints - Refactor model storage to include type information alongside instance - Move model ID and type methods to WhichModel implementation - Update API calls to access model instance through stored wrapper
This commit is contained in:
+234
-7
@@ -5,22 +5,32 @@ use aha::models::{GenerateModel, ModelInstance, WhichModel, load_model};
|
||||
use aha::utils::string_to_static_str;
|
||||
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||
use rocket::futures::StreamExt;
|
||||
use rocket::serde::json::Json;
|
||||
use rocket::serde::{json::Json, Serialize};
|
||||
use rocket::{
|
||||
Request,
|
||||
futures::Stream,
|
||||
get,
|
||||
http::{ContentType, Status},
|
||||
post,
|
||||
response::{Responder, stream::TextStream},
|
||||
};
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
static MODEL: OnceLock<Arc<RwLock<ModelInstance<'static>>>> = OnceLock::new();
|
||||
/// Wrapper to store model type together with the model instance
|
||||
struct StoredModel {
|
||||
which_model: WhichModel,
|
||||
instance: ModelInstance<'static>,
|
||||
}
|
||||
|
||||
static MODEL: OnceLock<Arc<RwLock<StoredModel>>> = OnceLock::new();
|
||||
|
||||
pub fn init(model_type: WhichModel, path: String) -> anyhow::Result<()> {
|
||||
let model_path = string_to_static_str(path);
|
||||
let model = load_model(model_type, model_path)?;
|
||||
MODEL.get_or_init(|| Arc::new(RwLock::new(model)));
|
||||
MODEL.get_or_init(|| Arc::new(RwLock::new(StoredModel {
|
||||
which_model: model_type,
|
||||
instance: model,
|
||||
})));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -62,7 +72,8 @@ pub(crate) async fn chat(
|
||||
.cloned()
|
||||
.ok_or_else(|| anyhow::anyhow!("model not init"))
|
||||
.unwrap();
|
||||
model_ref.write().await.generate(req.into_inner())
|
||||
let mut guard = model_ref.write().await;
|
||||
guard.instance.generate(req.into_inner())
|
||||
};
|
||||
match response {
|
||||
Ok(res) => {
|
||||
@@ -76,7 +87,7 @@ pub(crate) async fn chat(
|
||||
let text_stream = TextStream! {
|
||||
let model_ref = MODEL.get().cloned().ok_or_else(|| anyhow::anyhow!("model not init")).unwrap();
|
||||
let mut guard = model_ref.write().await;
|
||||
let stream_result = guard.generate_stream(req.into_inner());
|
||||
let stream_result = guard.instance.generate_stream(req.into_inner());
|
||||
match stream_result {
|
||||
Ok(stream) => {
|
||||
let mut stream = pin!(stream);
|
||||
@@ -113,7 +124,8 @@ pub(crate) async fn remove_background(req: Json<ChatCompletionParameters>) -> (S
|
||||
.cloned()
|
||||
.ok_or_else(|| anyhow::anyhow!("model not init"))
|
||||
.unwrap();
|
||||
model_ref.write().await.generate(req.into_inner())
|
||||
let mut guard = model_ref.write().await;
|
||||
guard.instance.generate(req.into_inner())
|
||||
};
|
||||
match response {
|
||||
Ok(res) => {
|
||||
@@ -132,7 +144,8 @@ pub(crate) async fn speech(req: Json<ChatCompletionParameters>) -> (Status, Stri
|
||||
.cloned()
|
||||
.ok_or_else(|| anyhow::anyhow!("model not init"))
|
||||
.unwrap();
|
||||
model_ref.write().await.generate(req.into_inner())
|
||||
let mut guard = model_ref.write().await;
|
||||
guard.instance.generate(req.into_inner())
|
||||
};
|
||||
match response {
|
||||
Ok(res) => {
|
||||
@@ -142,3 +155,217 @@ pub(crate) async fn speech(req: Json<ChatCompletionParameters>) -> (Status, Stri
|
||||
Err(e) => (Status::InternalServerError, e.to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
// Health check endpoint
|
||||
|
||||
#[derive(Serialize)]
|
||||
pub(crate) struct HealthResponse {
|
||||
status: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
error: Option<String>,
|
||||
}
|
||||
|
||||
#[get("/health")]
|
||||
pub(crate) async fn health() -> (Status, (ContentType, Json<HealthResponse>)) {
|
||||
if MODEL.get().is_some() {
|
||||
let response = HealthResponse {
|
||||
status: "ok".to_string(),
|
||||
error: None,
|
||||
};
|
||||
(Status::Ok, (ContentType::JSON, Json(response)))
|
||||
} else {
|
||||
let response = HealthResponse {
|
||||
status: "unhealthy".to_string(),
|
||||
error: Some("model not initialized".to_string()),
|
||||
};
|
||||
(Status::ServiceUnavailable, (ContentType::JSON, Json(response)))
|
||||
}
|
||||
}
|
||||
|
||||
// Models endpoint (OpenAI-compatible format)
|
||||
|
||||
/// OpenAI-compatible model object
|
||||
#[derive(Serialize)]
|
||||
struct ModelObject {
|
||||
id: String,
|
||||
object: String,
|
||||
created: Option<i64>,
|
||||
owned_by: String,
|
||||
}
|
||||
|
||||
/// OpenAI-compatible models list response
|
||||
#[derive(Serialize)]
|
||||
struct ModelsListResponse {
|
||||
object: String,
|
||||
data: Vec<ModelObject>,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct ErrorResponse {
|
||||
error: String,
|
||||
}
|
||||
|
||||
/// Convert WhichModel to a display-friendly model ID (kebab-case)
|
||||
fn which_model_to_id(which_model: WhichModel) -> &'static str {
|
||||
match which_model {
|
||||
WhichModel::MiniCPM4_0_5B => "minicpm4-0.5b",
|
||||
WhichModel::Qwen2_5vl3B => "qwen2.5vl-3b",
|
||||
WhichModel::Qwen2_5vl7B => "qwen2.5vl-7b",
|
||||
WhichModel::Qwen3_0_6B => "qwen3-0.6b",
|
||||
WhichModel::Qwen3ASR0_6B => "qwen3asr-0.6b",
|
||||
WhichModel::Qwen3ASR1_7B => "qwen3asr-1.7b",
|
||||
WhichModel::Qwen3vl2B => "qwen3vl-2b",
|
||||
WhichModel::Qwen3vl4B => "qwen3vl-4b",
|
||||
WhichModel::Qwen3vl8B => "qwen3vl-8b",
|
||||
WhichModel::Qwen3vl32B => "qwen3vl-32b",
|
||||
WhichModel::DeepSeekOCR => "deepseek-ocr",
|
||||
WhichModel::HunyuanOCR => "hunyuan-ocr",
|
||||
WhichModel::PaddleOCRVL => "paddleocr-vl",
|
||||
WhichModel::RMBG2_0 => "rmbg2.0",
|
||||
WhichModel::VoxCPM => "voxcpm",
|
||||
WhichModel::VoxCPM1_5 => "voxcpm1.5",
|
||||
WhichModel::GlmASRNano2512 => "glm-asr-nano-2512",
|
||||
WhichModel::FunASRNano2512 => "fun-asr-nano-2512",
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the owner/organization name for a model
|
||||
fn which_model_to_owner(which_model: WhichModel) -> &'static str {
|
||||
match which_model {
|
||||
WhichModel::MiniCPM4_0_5B => "OpenBMB",
|
||||
WhichModel::Qwen2_5vl3B | WhichModel::Qwen2_5vl7B => "Qwen",
|
||||
WhichModel::Qwen3_0_6B | WhichModel::Qwen3ASR0_6B | WhichModel::Qwen3ASR1_7B => "Qwen",
|
||||
WhichModel::Qwen3vl2B | WhichModel::Qwen3vl4B | WhichModel::Qwen3vl8B | WhichModel::Qwen3vl32B => "Qwen",
|
||||
WhichModel::DeepSeekOCR => "deepseek-ai",
|
||||
WhichModel::HunyuanOCR => "Tencent-Hunyuan",
|
||||
WhichModel::PaddleOCRVL => "PaddlePaddle",
|
||||
WhichModel::RMBG2_0 => "AI-ModelScope",
|
||||
WhichModel::VoxCPM | WhichModel::VoxCPM1_5 => "OpenBMB",
|
||||
WhichModel::GlmASRNano2512 => "ZhipuAI",
|
||||
WhichModel::FunASRNano2512 => "FunAudioLLM",
|
||||
}
|
||||
}
|
||||
|
||||
#[get("/models")]
|
||||
pub(crate) async fn models() -> (Status, (ContentType, Json<serde_json::Value>))
|
||||
{
|
||||
if let Some(model_ref) = MODEL.get() {
|
||||
let guard = model_ref.read().await;
|
||||
let which_model = guard.which_model;
|
||||
|
||||
let model_obj = ModelObject {
|
||||
id: which_model_to_id(which_model).to_string(),
|
||||
object: "model".to_string(),
|
||||
created: None, // We don't track creation time
|
||||
owned_by: which_model_to_owner(which_model).to_string(),
|
||||
};
|
||||
drop(guard);
|
||||
|
||||
let response = ModelsListResponse {
|
||||
object: "list".to_string(),
|
||||
data: vec![model_obj],
|
||||
};
|
||||
(Status::Ok, (ContentType::JSON, Json(serde_json::to_value(response).unwrap())))
|
||||
} else {
|
||||
let response = ErrorResponse {
|
||||
error: "model not initialized".to_string(),
|
||||
};
|
||||
(Status::ServiceUnavailable, (ContentType::JSON, Json(serde_json::to_value(response).unwrap())))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
// Test health endpoint when model is not initialized
|
||||
#[tokio::test]
|
||||
async fn test_health_endpoint_uninitialized() {
|
||||
let (status, (content_type, response)) = health().await;
|
||||
assert_eq!(status, Status::ServiceUnavailable);
|
||||
assert_eq!(content_type, ContentType::JSON);
|
||||
assert_eq!(response.status, "unhealthy");
|
||||
assert_eq!(response.error, Some("model not initialized".to_string()));
|
||||
}
|
||||
|
||||
// Test health endpoint when model is initialized
|
||||
// Note: This test requires a model to be initialized, which may not be feasible
|
||||
// in unit tests without access to model files. This is a placeholder for integration tests.
|
||||
//
|
||||
// #[tokio::test]
|
||||
// async fn test_health_endpoint_initialized() {
|
||||
// // This would require model initialization
|
||||
// // Consider moving to integration tests
|
||||
// }
|
||||
|
||||
// Test models endpoint when model is not initialized
|
||||
#[tokio::test]
|
||||
async fn test_models_endpoint_uninitialized() {
|
||||
let (status, (content_type, response)) = models().await;
|
||||
assert_eq!(status, Status::ServiceUnavailable);
|
||||
assert_eq!(content_type, ContentType::JSON);
|
||||
let error = response.get("error").and_then(|v| v.as_str());
|
||||
assert_eq!(error, Some("model not initialized"));
|
||||
}
|
||||
|
||||
// Test model type classification
|
||||
#[test]
|
||||
fn test_get_model_type_llm() {
|
||||
assert_eq!(WhichModel::Qwen3_0_6B.model_type(), "llm");
|
||||
assert_eq!(WhichModel::Qwen3vl2B.model_type(), "llm");
|
||||
assert_eq!(WhichModel::MiniCPM4_0_5B.model_type(), "llm");
|
||||
assert_eq!(WhichModel::Qwen2_5vl3B.model_type(), "llm");
|
||||
assert_eq!(WhichModel::Qwen2_5vl7B.model_type(), "llm");
|
||||
assert_eq!(WhichModel::Qwen3vl4B.model_type(), "llm");
|
||||
assert_eq!(WhichModel::Qwen3vl8B.model_type(), "llm");
|
||||
assert_eq!(WhichModel::Qwen3vl32B.model_type(), "llm");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_get_model_type_ocr() {
|
||||
assert_eq!(WhichModel::DeepSeekOCR.model_type(), "ocr");
|
||||
assert_eq!(WhichModel::HunyuanOCR.model_type(), "ocr");
|
||||
assert_eq!(WhichModel::PaddleOCRVL.model_type(), "ocr");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_get_model_type_asr() {
|
||||
assert_eq!(WhichModel::Qwen3ASR0_6B.model_type(), "asr");
|
||||
assert_eq!(WhichModel::Qwen3ASR1_7B.model_type(), "asr");
|
||||
assert_eq!(WhichModel::GlmASRNano2512.model_type(), "asr");
|
||||
assert_eq!(WhichModel::FunASRNano2512.model_type(), "asr");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_get_model_type_image() {
|
||||
assert_eq!(WhichModel::RMBG2_0.model_type(), "image");
|
||||
assert_eq!(WhichModel::VoxCPM.model_type(), "image");
|
||||
assert_eq!(WhichModel::VoxCPM1_5.model_type(), "image");
|
||||
}
|
||||
|
||||
// Test model_id retrieval
|
||||
#[test]
|
||||
fn test_get_model_id() {
|
||||
assert_eq!(WhichModel::Qwen3_0_6B.model_id(), "Qwen/Qwen3-0.6B");
|
||||
assert_eq!(WhichModel::DeepSeekOCR.model_id(), "deepseek-ai/DeepSeek-OCR");
|
||||
assert_eq!(WhichModel::VoxCPM1_5.model_id(), "OpenBMB/VoxCPM1.5");
|
||||
}
|
||||
|
||||
// Test OpenAI-compatible model ID conversion
|
||||
#[test]
|
||||
fn test_which_model_to_id() {
|
||||
assert_eq!(which_model_to_id(WhichModel::Qwen3_0_6B), "qwen3-0.6b");
|
||||
assert_eq!(which_model_to_id(WhichModel::DeepSeekOCR), "deepseek-ocr");
|
||||
assert_eq!(which_model_to_id(WhichModel::VoxCPM1_5), "voxcpm1.5");
|
||||
assert_eq!(which_model_to_id(WhichModel::MiniCPM4_0_5B), "minicpm4-0.5b");
|
||||
}
|
||||
|
||||
// Test owner/organization mapping
|
||||
#[test]
|
||||
fn test_which_model_to_owner() {
|
||||
assert_eq!(which_model_to_owner(WhichModel::Qwen3_0_6B), "Qwen");
|
||||
assert_eq!(which_model_to_owner(WhichModel::DeepSeekOCR), "deepseek-ai");
|
||||
assert_eq!(which_model_to_owner(WhichModel::VoxCPM1_5), "OpenBMB");
|
||||
assert_eq!(which_model_to_owner(WhichModel::HunyuanOCR), "Tencent-Hunyuan");
|
||||
}
|
||||
}
|
||||
|
||||
+7
-29
@@ -145,35 +145,11 @@ struct RunArgs {
|
||||
/// Get the default weight path for a given model
|
||||
/// Returns ~/.aha/{model_id} e.g., ~/.aha/OpenBMB/VoxCPM1.5
|
||||
fn get_default_weight_path(model: WhichModel) -> String {
|
||||
let model_id = get_model_id(model);
|
||||
let model_id = model.model_id();
|
||||
let save_dir = get_default_save_dir().expect("Failed to get home directory");
|
||||
format!("{}/{}", save_dir, model_id)
|
||||
}
|
||||
|
||||
/// Get the ModelScope model ID for a given WhichModel variant
|
||||
fn get_model_id(model: WhichModel) -> &'static str {
|
||||
match model {
|
||||
WhichModel::MiniCPM4_0_5B => "OpenBMB/MiniCPM4-0.5B",
|
||||
WhichModel::Qwen2_5vl3B => "Qwen/Qwen2.5-VL-3B-Instruct",
|
||||
WhichModel::Qwen2_5vl7B => "Qwen/Qwen2.5-VL-7B-Instruct",
|
||||
WhichModel::Qwen3_0_6B => "Qwen/Qwen3-0.6B",
|
||||
WhichModel::Qwen3ASR0_6B => "Qwen/Qwen3-ASR-0.6B",
|
||||
WhichModel::Qwen3ASR1_7B => "Qwen/Qwen3-ASR-1.7B",
|
||||
WhichModel::Qwen3vl2B => "Qwen/Qwen3-VL-2B-Instruct",
|
||||
WhichModel::Qwen3vl4B => "Qwen/Qwen3-VL-4B-Instruct",
|
||||
WhichModel::Qwen3vl8B => "Qwen/Qwen3-VL-8B-Instruct",
|
||||
WhichModel::Qwen3vl32B => "Qwen/Qwen3-VL-32B-Instruct",
|
||||
WhichModel::DeepSeekOCR => "deepseek-ai/DeepSeek-OCR",
|
||||
WhichModel::HunyuanOCR => "Tencent-Hunyuan/HunyuanOCR",
|
||||
WhichModel::PaddleOCRVL => "PaddlePaddle/PaddleOCR-VL",
|
||||
WhichModel::RMBG2_0 => "AI-ModelScope/RMBG-2.0",
|
||||
WhichModel::VoxCPM => "OpenBMB/VoxCPM-0.5B",
|
||||
WhichModel::VoxCPM1_5 => "OpenBMB/VoxCPM1.5",
|
||||
WhichModel::GlmASRNano2512 => "ZhipuAI/GLM-ASR-Nano-2512",
|
||||
WhichModel::FunASRNano2512 => "FunAudioLLM/Fun-ASR-Nano-2512",
|
||||
}
|
||||
}
|
||||
|
||||
/// List all supported models
|
||||
fn run_list() -> anyhow::Result<()> {
|
||||
let models = [
|
||||
@@ -204,7 +180,7 @@ fn run_list() -> anyhow::Result<()> {
|
||||
for model in models {
|
||||
let possible_value = model.to_possible_value().unwrap();
|
||||
let name = possible_value.get_name();
|
||||
let id = get_model_id(model);
|
||||
let id = model.model_id();
|
||||
println!("{:<30} {}", name, id);
|
||||
}
|
||||
|
||||
@@ -219,7 +195,7 @@ async fn run_cli(args: CliArgs) -> anyhow::Result<()> {
|
||||
save_dir,
|
||||
download_retries,
|
||||
} = args;
|
||||
let model_id = get_model_id(common.model);
|
||||
let model_id = common.model.model_id();
|
||||
|
||||
let model_path = match weight_path {
|
||||
Some(path) => path,
|
||||
@@ -265,7 +241,7 @@ async fn run_download(args: DownloadArgs) -> anyhow::Result<()> {
|
||||
save_dir,
|
||||
download_retries,
|
||||
} = args;
|
||||
let model_id = get_model_id(model);
|
||||
let model_id = model.model_id();
|
||||
|
||||
let save_dir = match save_dir {
|
||||
Some(dir) => dir,
|
||||
@@ -416,8 +392,10 @@ pub(crate) async fn start_http_server(address: String, port: u16) -> anyhow::Res
|
||||
builder = builder.mount("/chat", routes![api::chat]);
|
||||
// /images/remove_background
|
||||
builder = builder.mount("/images", routes![api::remove_background]);
|
||||
// /images/speech
|
||||
// /audio/speech
|
||||
builder = builder.mount("/audio", routes![api::speech]);
|
||||
// Health check and model info endpoints
|
||||
builder = builder.mount("/", routes![api::health, api::models]);
|
||||
|
||||
builder.launch().await?;
|
||||
Ok(())
|
||||
|
||||
@@ -74,6 +74,56 @@ pub enum WhichModel {
|
||||
FunASRNano2512,
|
||||
}
|
||||
|
||||
impl WhichModel {
|
||||
/// Get the ModelScope model ID for this model variant
|
||||
pub fn model_id(self) -> &'static str {
|
||||
match self {
|
||||
WhichModel::MiniCPM4_0_5B => "OpenBMB/MiniCPM4-0.5B",
|
||||
WhichModel::Qwen2_5vl3B => "Qwen/Qwen2.5-VL-3B-Instruct",
|
||||
WhichModel::Qwen2_5vl7B => "Qwen/Qwen2.5-VL-7B-Instruct",
|
||||
WhichModel::Qwen3_0_6B => "Qwen/Qwen3-0.6B",
|
||||
WhichModel::Qwen3ASR0_6B => "Qwen/Qwen3-ASR-0.6B",
|
||||
WhichModel::Qwen3ASR1_7B => "Qwen/Qwen3-ASR-1.7B",
|
||||
WhichModel::Qwen3vl2B => "Qwen/Qwen3-VL-2B-Instruct",
|
||||
WhichModel::Qwen3vl4B => "Qwen/Qwen3-VL-4B-Instruct",
|
||||
WhichModel::Qwen3vl8B => "Qwen/Qwen3-VL-8B-Instruct",
|
||||
WhichModel::Qwen3vl32B => "Qwen/Qwen3-VL-32B-Instruct",
|
||||
WhichModel::DeepSeekOCR => "deepseek-ai/DeepSeek-OCR",
|
||||
WhichModel::HunyuanOCR => "Tencent-Hunyuan/HunyuanOCR",
|
||||
WhichModel::PaddleOCRVL => "PaddlePaddle/PaddleOCR-VL",
|
||||
WhichModel::RMBG2_0 => "AI-ModelScope/RMBG-2.0",
|
||||
WhichModel::VoxCPM => "OpenBMB/VoxCPM-0.5B",
|
||||
WhichModel::VoxCPM1_5 => "OpenBMB/VoxCPM1.5",
|
||||
WhichModel::GlmASRNano2512 => "ZhipuAI/GLM-ASR-Nano-2512",
|
||||
WhichModel::FunASRNano2512 => "FunAudioLLM/Fun-ASR-Nano-2512",
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the model type category for this model variant
|
||||
pub fn model_type(self) -> &'static str {
|
||||
match self {
|
||||
// LLM models
|
||||
WhichModel::MiniCPM4_0_5B
|
||||
| WhichModel::Qwen2_5vl3B
|
||||
| WhichModel::Qwen2_5vl7B
|
||||
| WhichModel::Qwen3_0_6B
|
||||
| WhichModel::Qwen3vl2B
|
||||
| WhichModel::Qwen3vl4B
|
||||
| WhichModel::Qwen3vl8B
|
||||
| WhichModel::Qwen3vl32B => "llm",
|
||||
// OCR models
|
||||
WhichModel::DeepSeekOCR | WhichModel::HunyuanOCR | WhichModel::PaddleOCRVL => "ocr",
|
||||
// ASR models
|
||||
WhichModel::Qwen3ASR0_6B
|
||||
| WhichModel::Qwen3ASR1_7B
|
||||
| WhichModel::GlmASRNano2512
|
||||
| WhichModel::FunASRNano2512 => "asr",
|
||||
// Image models
|
||||
WhichModel::RMBG2_0 | WhichModel::VoxCPM | WhichModel::VoxCPM1_5 => "image",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub trait GenerateModel {
|
||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse>;
|
||||
fn generate_stream(
|
||||
|
||||
Reference in New Issue
Block a user