From 63a89b0c9656d8de045ae89d300019240711cb90 Mon Sep 17 00:00:00 2001 From: XiaoYang Date: Fri, 6 Mar 2026 11:38:24 +0800 Subject: [PATCH] feat(api): add OpenAI-compatible ASR transcription endpoint - Implement POST /audio/transcriptions and /v1/audio/transcriptions endpoints for automatic speech recognition - Add support for multipart/form-data audio file uploads - Support multiple audio formats (wav, mp3, m4a, etc.) - Implement language detection with 29 supported languages - Return transcription text in OpenAI-compatible JSON format - Add proper error handling and validation - Include comprehensive tests for ASR functionality refactor(api): restructure API modules and export MODEL - Move API declarations to src/api/mod.rs - Add ASR module and type definitions - Export MODEL static for use in ASR module - Mount ASR routes at both /audio and /v1/audio endpoints --- Makefile | 8 ++ src/api/asr.rs | 192 +++++++++++++++++++++++++++++++++++++ src/api/asr_types.rs | 78 +++++++++++++++ src/{api.rs => api/mod.rs} | 12 ++- src/main.rs | 6 +- src/utils/mod.rs | 106 ++++++++++++++++++++ 6 files changed, 398 insertions(+), 4 deletions(-) create mode 100644 src/api/asr.rs create mode 100644 src/api/asr_types.rs rename src/{api.rs => api/mod.rs} (98%) diff --git a/Makefile b/Makefile index 3fe8d57..1504d86 100644 --- a/Makefile +++ b/Makefile @@ -14,6 +14,14 @@ build_mac: @echo "Building project for macOS..." @cargo build --features metal --release +build_mac_universal: + @echo "Building universal binary for macOS..." + @cargo build --features metal --release --target aarch64-apple-darwin + @cargo build --features metal --release --target x86_64-apple-darwin + @mkdir -p target/universal/release + @lipo -create target/aarch64-apple-darwin/release/aha target/x86_64-apple-darwin/release/aha -output target/universal/release/aha + + test: @echo "Running tests..." @cargo test diff --git a/src/api/asr.rs b/src/api/asr.rs new file mode 100644 index 0000000..9aeb0d6 --- /dev/null +++ b/src/api/asr.rs @@ -0,0 +1,192 @@ +// OpenAI-compatible ASR (Automatic Speech Recognition) API endpoint +// Implements POST /audio/transcriptions and /v1/audio/transcriptions + +use aha::models::GenerateModel; +use aha::utils::{clean_asr_response, map_language_code}; +use aha_openai_dive::v1::resources::chat::{ + ChatCompletionParameters, ChatMessage, ChatMessageAudioContentPart, ChatMessageContent, + ChatMessageContentPart, AudioUrlType, +}; +use rocket::http::Status; +use rocket::serde::json::Json; +use rocket::{form::Form, post}; + +use super::asr_types::{ErrorResponse, ErrorDetail, TranscriptionRequest, TranscriptionResponse}; +use super::MODEL; + +/// Handle audio transcription requests +/// +/// This endpoint accepts multipart/form-data with an audio file and returns +/// the transcription text in OpenAI-compatible format. +/// +/// # Supported Parameters +/// - `file`: Audio file (required) - wav, mp3, m4a, etc. +/// - `model`: Model name (optional, ignored) +/// - `language`: Language code (optional) - zh, en, yue, ar, de, fr, es, pt, id, it, ko, ru, th, vi, ja, tr, hi, ms, nl, sv, da, fi, pl, cs, fil, fa, el, ro, hu, mk +/// - `prompt`: Optional prompt text (ignored in this implementation) +/// - `response_format`: Response format (only "json" supported) +/// - `temperature`: Sampling temperature (0.0 to 1.0, default 0.0) +/// +/// # Returns +/// JSON response with format: `{"text": "transcribed text"}` +#[post("/transcriptions", data = "")] +pub(crate) async fn transcriptions(req: Form>) -> (Status, Json) { + // Validate response_format (only JSON supported) + if let Some(ref format) = req.response_format { + if format != "json" && format != "text" { + return error_response( + Status::BadRequest, + "invalid_request_error", + "Only 'json' response format is supported", + Some("unsupported_format".to_string()), + ); + } + } + + // Get the audio file path + let file_path = match req.file.path() { + Some(path) => path, + None => { + return error_response( + Status::BadRequest, + "invalid_request_error", + "Audio file is required", + Some("missing_file".to_string()), + ); + } + }; + + // Build file:// URL for the model + let file_url = format!("file://{}", file_path.display()); + + // Map language code to full language name + let language_name = req.language.as_ref().and_then(|code| map_language_code(code)); + + // Build ChatCompletionParameters for the ASR model + let audio_part = ChatMessageContentPart::Audio(ChatMessageAudioContentPart { + r#type: "audio".to_string(), + audio_url: AudioUrlType { url: file_url }, + }); + + let params = ChatCompletionParameters { + messages: vec![ChatMessage::User { + content: ChatMessageContent::ContentPart(vec![audio_part]), + name: None, + }], + model: req.model.clone().unwrap_or_else(|| "asr".to_string()), + temperature: req.temperature.or(Some(0.0)), + max_tokens: None, + stream: None, + top_p: None, + frequency_penalty: None, + presence_penalty: None, + stop: None, + n: None, + tools: None, + tool_choice: None, + response_format: None, + metadata: language_name.map(|lang| { + let mut map = std::collections::HashMap::new(); + map.insert("language".to_string(), lang); + map + }), + ..Default::default() + }; + + // Get the model and generate transcription + let model_ref = match MODEL.get() { + Some(m) => m, + None => { + return error_response( + Status::ServiceUnavailable, + "service_unavailable", + "Model not initialized", + Some("model_not_loaded".to_string()), + ); + } + }; + + let response = { + let mut guard = model_ref.write().await; + guard.instance.generate(params) + }; + + match response { + Ok(chat_response) => { + // Extract the transcription text from the response + let raw_text = chat_response + .choices + .first() + .and_then(|choice| { + if let ChatMessage::Assistant { content, .. } = &choice.message { + content.as_ref().and_then(|c| { + if let ChatMessageContent::Text(text) = c { + Some(text.clone()) + } else { + None + } + }) + } else { + None + } + }) + .unwrap_or_else(|| String::new()); + + // Clean the response (remove "language English" prefix) + let cleaned_text = clean_asr_response(&raw_text); + + // Return OpenAI-compatible transcription response + let transcription = TranscriptionResponse { text: cleaned_text }; + (Status::Ok, Json(serde_json::to_value(transcription).unwrap())) + } + Err(e) => { + // Determine appropriate error status based on error message + let error_msg = e.to_string(); + let (status, error_type, code) = if error_msg.contains("audio") || error_msg.contains("decode") { + (Status::BadRequest, "invalid_request_error", Some("audio_decode_error".to_string())) + } else { + (Status::InternalServerError, "server_error", Some("inference_error".to_string())) + }; + + error_response(status, error_type, &error_msg, code) + } + } +} + +/// Helper function to create error responses in OpenAI format +fn error_response( + status: Status, + error_type: &str, + message: &str, + code: Option, +) -> (Status, Json) { + let error_response = ErrorResponse { + error: ErrorDetail { + message: message.to_string(), + error_type: error_type.to_string(), + code, + }, + }; + (status, Json(serde_json::to_value(error_response).unwrap())) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_error_response_serialization() { + let (status, json) = error_response( + Status::BadRequest, + "invalid_request_error", + "Test error message", + Some("test_code".to_string()), + ); + + assert_eq!(status, Status::BadRequest); + let parsed: serde_json::Value = serde_json::from_str(&json.to_string()).unwrap(); + assert_eq!(parsed["error"]["message"], "Test error message"); + assert_eq!(parsed["error"]["type"], "invalid_request_error"); + assert_eq!(parsed["error"]["code"], "test_code"); + } +} diff --git a/src/api/asr_types.rs b/src/api/asr_types.rs new file mode 100644 index 0000000..4ef26df --- /dev/null +++ b/src/api/asr_types.rs @@ -0,0 +1,78 @@ +// ASR API data types for OpenAI-compatible transcription endpoint + +use rocket::form::FromForm; +use serde::Serialize; + +/// Request parameters for audio transcription +#[derive(Debug, FromForm)] +pub(crate) struct TranscriptionRequest<'r> { + /// The audio file to transcribe + pub(crate) file: rocket::fs::TempFile<'r>, + + /// ID of the model to use (ignored, always uses loaded model) + pub(crate) model: Option, + + /// Language code (e.g., "zh", "en") + pub(crate) language: Option, + + /// Optional text to guide the transcription (not implemented, ignored) + #[allow(dead_code)] + pub(crate) prompt: Option, + + /// Response format (only "json" supported) + pub(crate) response_format: Option, + + /// Sampling temperature (0.0 to 1.0) + pub(crate) temperature: Option, +} + +/// Standard transcription response +#[derive(Debug, Serialize)] +pub(crate) struct TranscriptionResponse { + pub(crate) text: String, +} + +/// Error response following OpenAI format +#[derive(Debug, Serialize)] +pub(crate) struct ErrorResponse { + pub(crate) error: ErrorDetail, +} + +#[derive(Debug, Serialize)] +pub(crate) struct ErrorDetail { + pub(crate) message: String, + #[serde(rename = "type")] + pub(crate) error_type: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub(crate) code: Option, +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_transcription_response_serialization() { + let response = TranscriptionResponse { + text: "Hello, world!".to_string(), + }; + let json = serde_json::to_string(&response).unwrap(); + assert_eq!(json, r#"{"text":"Hello, world!"}"#); + } + + #[test] + fn test_error_response_serialization() { + let error = ErrorResponse { + error: ErrorDetail { + message: "Invalid audio file".to_string(), + error_type: "invalid_request_error".to_string(), + code: Some("invalid_audio".to_string()), + }, + }; + let json = serde_json::to_string(&error).unwrap(); + let parsed: serde_json::Value = serde_json::from_str(&json).unwrap(); + assert_eq!(parsed["error"]["message"], "Invalid audio file"); + assert_eq!(parsed["error"]["type"], "invalid_request_error"); + assert_eq!(parsed["error"]["code"], "invalid_audio"); + } +} diff --git a/src/api.rs b/src/api/mod.rs similarity index 98% rename from src/api.rs rename to src/api/mod.rs index 060324a..c5b30ce 100644 --- a/src/api.rs +++ b/src/api/mod.rs @@ -18,13 +18,21 @@ use rocket::{ }; use tokio::sync::RwLock; +// ASR (Automatic Speech Recognition) API module +pub(crate) mod asr; +pub(crate) mod asr_types; + +// Re-export ASR routes +pub(crate) use asr::transcriptions; + /// Wrapper to store model type together with the model instance -struct StoredModel { +pub(crate) struct StoredModel { which_model: WhichModel, instance: ModelInstance<'static>, } -static MODEL: OnceLock>> = OnceLock::new(); +// Export MODEL for use in ASR module +pub(crate) static MODEL: OnceLock>> = OnceLock::new(); static SHUTDOWN_FLAG: OnceLock> = OnceLock::new(); static SERVER_PORT: OnceLock = OnceLock::new(); static ALLOW_REMOTE_SHUTDOWN: OnceLock = OnceLock::new(); diff --git a/src/main.rs b/src/main.rs index 95d8179..4fa6fe4 100644 --- a/src/main.rs +++ b/src/main.rs @@ -632,8 +632,10 @@ pub(crate) async fn start_http_server( builder = builder.mount("/chat", routes![api::chat]); // /images/remove_background builder = builder.mount("/images", routes![api::remove_background]); - // /audio/speech - builder = builder.mount("/audio", routes![api::speech]); + // /audio/speech and /audio/transcriptions (ASR transcription endpoint) + builder = builder.mount("/audio", routes![api::speech, api::transcriptions]); + // /v1/audio/transcriptions (OpenAI standard ASR transcription endpoint) + builder = builder.mount("/v1/audio", routes![api::transcriptions]); // Health check and model info endpoints builder = builder.mount("/", routes![api::health, api::models]); // Shutdown endpoint diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 031950a..a5f7ab1 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -850,3 +850,109 @@ pub fn load_tensor_from_pt( let t = Tensor::from_vec(data, shape, device)?; Ok(t) } + +/// Map OpenAI language code to full language name for ASR models +/// +/// Supports 29 languages as per Qwen3ASR specification +pub fn map_language_code(code: &str) -> Option { + match code.to_lowercase().as_str() { + "zh" => Some("Chinese".to_string()), + "en" => Some("English".to_string()), + "yue" => Some("Cantonese".to_string()), + "ar" => Some("Arabic".to_string()), + "de" => Some("German".to_string()), + "fr" => Some("French".to_string()), + "es" => Some("Spanish".to_string()), + "pt" => Some("Portuguese".to_string()), + "id" => Some("Indonesian".to_string()), + "it" => Some("Italian".to_string()), + "ko" => Some("Korean".to_string()), + "ru" => Some("Russian".to_string()), + "th" => Some("Thai".to_string()), + "vi" => Some("Vietnamese".to_string()), + "ja" => Some("Japanese".to_string()), + "tr" => Some("Turkish".to_string()), + "hi" => Some("Hindi".to_string()), + "ms" => Some("Malay".to_string()), + "nl" => Some("Dutch".to_string()), + "sv" => Some("Swedish".to_string()), + "da" => Some("Danish".to_string()), + "fi" => Some("Finnish".to_string()), + "pl" => Some("Polish".to_string()), + "cs" => Some("Czech".to_string()), + "fil" => Some("Filipino".to_string()), + "fa" => Some("Persian".to_string()), + "el" => Some("Greek".to_string()), + "ro" => Some("Romanian".to_string()), + "hu" => Some("Hungarian".to_string()), + "mk" => Some("Macedonian".to_string()), + _ => None, + } +} + +/// Clean ASR model output by extracting pure text from model-specific format +/// +/// Qwen3ASR outputs format: "language EnglishThe morning sun..." +/// This function extracts the text after "" marker. +/// If no marker is found, returns the original text trimmed (for compatibility). +pub fn clean_asr_response(raw: &str) -> String { + if let Some(start) = raw.find("") { + raw[start + "".len()..].trim().to_string() + } else { + raw.trim().to_string() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_map_language_code_chinese() { + assert_eq!(map_language_code("zh"), Some("Chinese".to_string())); + } + + #[test] + fn test_map_language_code_english() { + assert_eq!(map_language_code("en"), Some("English".to_string())); + } + + #[test] + fn test_map_language_code_case_insensitive() { + assert_eq!(map_language_code("ZH"), Some("Chinese".to_string())); + assert_eq!(map_language_code("EN"), Some("English".to_string())); + } + + #[test] + fn test_map_language_code_invalid() { + assert_eq!(map_language_code("xx"), None); + } + + #[test] + fn test_clean_asr_response_standard_format() { + let raw = "language EnglishThe morning sun cast golden light"; + let cleaned = clean_asr_response(raw); + assert_eq!(cleaned, "The morning sun cast golden light"); + } + + #[test] + fn test_clean_asr_response_chinese_format() { + let raw = "language Chinese科技不断改变着我们的生活"; + let cleaned = clean_asr_response(raw); + assert_eq!(cleaned, "科技不断改变着我们的生活"); + } + + #[test] + fn test_clean_asr_response_with_newlines() { + let raw = "language English\n\n Hello world\n "; + let cleaned = clean_asr_response(raw); + assert_eq!(cleaned, "Hello world"); + } + + #[test] + fn test_clean_asr_response_no_marker() { + let raw = " Plain text without marker "; + let cleaned = clean_asr_response(raw); + assert_eq!(cleaned, "Plain text without marker"); + } +}