1use crate::{openai::*, traits::HttpServer, types::*};
7use async_trait::async_trait;
8use axum::{
9 extract::State,
10 http::StatusCode as AxumStatusCode,
11 response::{sse::Event, IntoResponse, Response, Sse},
12 routing::{get, post},
13 Json, Router,
14};
15use ferrum_interfaces::engine::InferenceEngine;
16use ferrum_types::{
17 FerrumError as Error, FinishReason, InferenceRequest, ModelId, Priority, RequestId,
18 SamplingParams,
19};
20use std::sync::Arc;
21use tokio::sync::mpsc;
22use tokio_stream::StreamExt;
23use tower::ServiceBuilder;
24use tower_http::{cors::CorsLayer, trace::TraceLayer};
25use tracing::{debug, error, info, span, Level};
26use uuid::Uuid;
27
28static PROM_HANDLE: std::sync::OnceLock<metrics_exporter_prometheus::PrometheusHandle> =
30 std::sync::OnceLock::new();
31
32pub fn init_prometheus_recorder() {
37 PROM_HANDLE.get_or_init(|| {
38 let builder = metrics_exporter_prometheus::PrometheusBuilder::new();
39 let handle = builder
40 .install_recorder()
41 .expect("Failed to install Prometheus recorder");
42 info!("Prometheus metrics recorder installed");
43 handle
44 });
45}
46
47pub struct AxumServer {
49 engine: Arc<dyn InferenceEngine + Send + Sync>,
50 config: ServerConfig,
51}
52
53impl AxumServer {
54 pub fn new(engine: Arc<dyn InferenceEngine + Send + Sync>) -> Self {
56 Self {
57 engine,
58 config: ServerConfig::default(),
59 }
60 }
61
62 fn build_router(&self) -> Router {
64 let app_state = AppState {
65 engine: self.engine.clone(),
66 };
67
68 Router::new()
69 .route("/v1/chat/completions", post(chat_completions_handler))
71 .route("/v1/completions", post(completions_handler))
72 .route("/v1/embeddings", post(embeddings_handler))
73 .route("/v1/audio/transcriptions", post(transcriptions_handler))
74 .route("/v1/audio/speech", post(speech_handler))
75 .route("/v1/models", get(models_handler))
76 .route("/health", get(health_handler))
78 .route("/metrics", get(metrics_handler))
79 .route("/", get(root_handler))
80 .layer(
82 ServiceBuilder::new()
83 .layer(TraceLayer::new_for_http())
84 .layer(CorsLayer::permissive()), )
86 .with_state(app_state)
87 }
88}
89
90#[derive(Clone)]
92struct AppState {
93 engine: Arc<dyn InferenceEngine + Send + Sync>,
94}
95
96#[async_trait]
97impl HttpServer for AxumServer {
98 async fn start(&self, config: &ServerConfig) -> ferrum_types::Result<()> {
99 let addr = format!("{}:{}", config.host, config.port);
100 info!("Starting Axum server on {}", addr);
101
102 let app = self.build_router();
103 let listener = tokio::net::TcpListener::bind(&addr)
104 .await
105 .map_err(|e| Error::internal(format!("Failed to bind to {}: {}", addr, e)))?;
106
107 info!("Server listening on {}", addr);
108
109 axum::serve(listener, app)
110 .await
111 .map_err(|e| Error::internal(format!("Server error: {}", e)))?;
112
113 Ok(())
114 }
115
116 async fn stop(&self, _timeout: std::time::Duration) -> ferrum_types::Result<()> {
117 info!("Stopping Axum server");
118 Ok(())
120 }
121
122 fn is_running(&self) -> bool {
123 true
125 }
126
127 fn address(&self) -> Option<std::net::SocketAddr> {
128 format!("{}:{}", self.config.host, self.config.port)
130 .parse()
131 .ok()
132 }
133
134 fn register_handler(
135 &mut self,
136 _path: &str,
137 _method: HttpMethod,
138 _handler: Box<dyn crate::traits::RequestHandler>,
139 ) {
140 unimplemented!("Dynamic handler registration not implemented in MVP")
142 }
143
144 fn register_middleware(&mut self, _middleware: Box<dyn crate::traits::Middleware>) {
145 unimplemented!("Dynamic middleware registration not implemented in MVP")
147 }
148
149 fn get_metrics(&self) -> ServerMetrics {
150 ServerMetrics {
152 total_requests: 0,
153 requests_by_endpoint: std::collections::HashMap::new(),
154 requests_by_status: std::collections::HashMap::new(),
155 avg_response_time_ms: 0.0,
156 p95_response_time_ms: 0.0,
157 p99_response_time_ms: 0.0,
158 active_connections: 0,
159 bytes_sent: 0,
160 bytes_received: 0,
161 error_rate: 0.0,
162 uptime_seconds: 0,
163 }
164 }
165
166 async fn health_check(&self) -> HealthStatus {
167 HealthStatus::Healthy
168 }
169}
170
171async fn chat_completions_handler(
173 State(state): State<AppState>,
174 Json(request): Json<ChatCompletionsRequest>,
175) -> std::result::Result<Response, ServerError> {
176 let span = span!(Level::INFO, "chat_completions", model = %request.model);
177 let _enter = span.enter();
178
179 info!(
180 "Received chat completions request for model: {}",
181 request.model
182 );
183 debug!("Request: {:?}", request);
184
185 let inference_request =
187 convert_chat_request(&request).map_err(|e| ServerError::BadRequest(e.to_string()))?;
188
189 if request.stream.unwrap_or(false) {
191 handle_chat_completions_stream(state, request, inference_request).await
192 } else {
193 handle_chat_completions_sync(state, request, inference_request).await
194 }
195}
196
197async fn handle_chat_completions_stream(
199 state: AppState,
200 openai_request: ChatCompletionsRequest,
201 inference_request: InferenceRequest,
202) -> std::result::Result<Response, ServerError> {
203 let (tx, rx) = mpsc::unbounded_channel::<std::result::Result<Event, axum::Error>>();
204
205 let engine = state.engine.clone();
207 let request_id = Uuid::new_v4().to_string();
208 let prompt_tokens = inference_request.prompt.split_whitespace().count() as u32;
209
210 tokio::spawn(async move {
211 let mut current_text = String::new();
212 let mut token_count = 0;
213 let max_tokens = openai_request.max_tokens.unwrap_or(100);
214
215 match engine.infer_stream(inference_request).await {
216 Ok(mut stream) => {
217 while let Some(result) = stream.next().await {
218 match result {
219 Ok(chunk) => {
220 current_text.push_str(&chunk.text);
221 token_count += 1;
222
223 let response_chunk = ChatCompletionsResponse {
225 id: request_id.clone(),
226 object: "chat.completion.chunk".to_string(),
227 created: chrono::Utc::now().timestamp() as u64,
228 model: openai_request.model.clone(),
229 choices: vec![ChatChoice {
230 index: 0,
231 message: None,
232 delta: Some(ChatMessage {
233 role: MessageRole::Assistant,
234 content: chunk.text.clone(),
235 name: None,
236 }),
237 finish_reason: None,
238 }],
239 usage: None,
240 };
241
242 let sse_event = Event::default()
243 .json_data(&response_chunk)
244 .unwrap_or_else(|_| Event::default().data("error"));
245 if tx.send(Ok(sse_event)).is_err() {
246 break;
247 }
248
249 if token_count >= max_tokens || chunk.finish_reason.is_some() {
251 let final_chunk = ChatCompletionsResponse {
253 id: request_id.clone(),
254 object: "chat.completion.chunk".to_string(),
255 created: chrono::Utc::now().timestamp() as u64,
256 model: openai_request.model.clone(),
257 choices: vec![ChatChoice {
258 index: 0,
259 message: None,
260 delta: None,
261 finish_reason: chunk
262 .finish_reason
263 .as_ref()
264 .map(finish_reason_to_string),
265 }],
266 usage: Some(Usage {
267 prompt_tokens: prompt_tokens,
268 completion_tokens: token_count,
269 total_tokens: prompt_tokens + token_count,
270 }),
271 };
272
273 let final_event = Event::default()
274 .json_data(&final_chunk)
275 .unwrap_or_else(|_| Event::default().data("error"));
276 let _ = tx.send(Ok(final_event));
277 let _ = tx.send(Ok(Event::default().data("[DONE]")));
278 break;
279 }
280 }
281 Err(e) => {
282 error!("Stream generation error: {}", e);
283 let _ = tx.send(Ok(
284 Event::default().data(&format!("{{\"error\": \"{}\"}}", e))
285 ));
286 break;
287 }
288 }
289 }
290 }
291 Err(e) => {
292 error!("Failed to start streaming: {}", e);
293 let _ = tx.send(Ok(
294 Event::default().data(&format!("{{\"error\": \"{}\"}}", e))
295 ));
296 }
297 }
298 });
299
300 let stream = tokio_stream::wrappers::UnboundedReceiverStream::new(rx);
301 let sse_stream = Sse::new(stream);
302
303 Ok(sse_stream.into_response())
304}
305
306async fn handle_chat_completions_sync(
308 state: AppState,
309 openai_request: ChatCompletionsRequest,
310 inference_request: InferenceRequest,
311) -> std::result::Result<Response, ServerError> {
312 info!("Processing non-streaming chat completion");
313
314 let prompt_tokens = inference_request.prompt.split_whitespace().count() as u32;
315 match state.engine.infer(inference_request).await {
316 Ok(output) => {
317 let response = ChatCompletionsResponse {
318 id: Uuid::new_v4().to_string(),
319 object: "chat.completion".to_string(),
320 created: chrono::Utc::now().timestamp() as u64,
321 model: openai_request.model,
322 choices: vec![ChatChoice {
323 index: 0,
324 message: Some(ChatMessage {
325 role: MessageRole::Assistant,
326 content: output.text,
327 name: None,
328 }),
329 delta: None,
330 finish_reason: Some(finish_reason_to_string(&output.finish_reason)),
331 }],
332 usage: Some(Usage {
333 prompt_tokens: prompt_tokens,
334 completion_tokens: output.tokens.len() as u32,
335 total_tokens: prompt_tokens + output.tokens.len() as u32,
336 }),
337 };
338
339 Ok(Json(response).into_response())
340 }
341 Err(e) => {
342 error!("Generation failed: {}", e);
343 Err(ServerError::InternalError(e.to_string()))
344 }
345 }
346}
347
348fn convert_chat_request(
350 request: &ChatCompletionsRequest,
351) -> ferrum_types::Result<InferenceRequest> {
352 let prompt = request
354 .messages
355 .iter()
356 .map(|msg| format!("{}: {}", msg.role.to_string(), msg.content))
357 .collect::<Vec<_>>()
358 .join("\n");
359
360 Ok(InferenceRequest {
361 id: RequestId(Uuid::new_v4()),
362 model_id: ModelId(request.model.clone()),
363 prompt,
364 sampling_params: SamplingParams {
365 max_tokens: request.max_tokens.unwrap_or(100) as usize,
366 temperature: request.temperature.unwrap_or(1.0),
367 top_p: request.top_p.unwrap_or(1.0),
368 top_k: None, repetition_penalty: 1.0,
370 presence_penalty: request.presence_penalty.unwrap_or(0.0),
371 frequency_penalty: request.frequency_penalty.unwrap_or(0.0),
372 stop_sequences: request.stop.clone().unwrap_or_default(),
373 seed: request.seed,
374 min_p: None,
375 tfs: None,
376 typical_p: None,
377 mirostat: None,
378 response_format: match &request.response_format {
379 Some(rf) if rf.format_type == "json_object" => {
380 ferrum_types::ResponseFormat::JsonObject
381 }
382 _ => ferrum_types::ResponseFormat::Text,
383 },
384 },
385 stream: request.stream.unwrap_or(false),
386 priority: Priority::Normal, client_id: None,
388 session_id: None,
389 created_at: chrono::Utc::now(),
390 metadata: std::collections::HashMap::new(),
391 })
392}
393
394async fn completions_handler(
396 State(_state): State<AppState>,
397 Json(_request): Json<CompletionsRequest>,
398) -> std::result::Result<Response, ServerError> {
399 Err(ServerError::NotImplemented(
401 "Legacy completions not implemented in MVP".to_string(),
402 ))
403}
404
405async fn embeddings_handler(
407 State(state): State<AppState>,
408 Json(request): Json<EmbeddingsRequest>,
409) -> std::result::Result<Response, ServerError> {
410 let span = span!(Level::INFO, "embeddings", model = %request.model);
411 let _enter = span.enter();
412
413 let items: Vec<EmbeddingItem> = match request.input {
415 EmbeddingInput::Single(text) => vec![EmbeddingItem {
416 text: Some(text),
417 image: None,
418 }],
419 EmbeddingInput::Batch(texts) => texts
420 .into_iter()
421 .map(|t| EmbeddingItem {
422 text: Some(t),
423 image: None,
424 })
425 .collect(),
426 EmbeddingInput::SingleObject(item) => vec![item],
427 EmbeddingInput::BatchObjects(items) => items,
428 };
429
430 if items.is_empty() {
431 return Err(ServerError::BadRequest("Empty input".to_string()));
432 }
433
434 let mut data = Vec::with_capacity(items.len());
435 let mut total_tokens = 0u32;
436
437 for (idx, item) in items.iter().enumerate() {
438 let embedding = if let Some(ref image) = item.image {
439 state
440 .engine
441 .embed_image(image)
442 .await
443 .map_err(|e| ServerError::InternalError(format!("embed_image: {e}")))?
444 } else if let Some(ref text) = item.text {
445 total_tokens += text.len() as u32;
446 state
447 .engine
448 .embed_text(text)
449 .await
450 .map_err(|e| ServerError::InternalError(format!("embed_text: {e}")))?
451 } else {
452 return Err(ServerError::BadRequest(
453 "Each input must have either 'text' or 'image'".to_string(),
454 ));
455 };
456
457 data.push(EmbeddingData {
458 object: "embedding".to_string(),
459 embedding,
460 index: idx,
461 });
462 }
463
464 let response = EmbeddingsResponse {
465 object: "list".to_string(),
466 data,
467 model: request.model,
468 usage: EmbeddingUsage {
469 prompt_tokens: total_tokens,
470 total_tokens,
471 },
472 };
473
474 Ok(Json(response).into_response())
475}
476
477async fn transcriptions_handler(
479 State(state): State<AppState>,
480 mut multipart: axum::extract::Multipart,
481) -> std::result::Result<Response, ServerError> {
482 let span = span!(Level::INFO, "transcription");
483 let _enter = span.enter();
484
485 let mut file_data: Option<Vec<u8>> = None;
486 let mut language: Option<String> = None;
487
488 while let Some(field) = multipart
489 .next_field()
490 .await
491 .map_err(|e| ServerError::BadRequest(format!("multipart: {e}")))?
492 {
493 let name = field.name().unwrap_or("").to_string();
494 match name.as_str() {
495 "file" => {
496 file_data = Some(
497 field
498 .bytes()
499 .await
500 .map_err(|e| ServerError::BadRequest(format!("read file: {e}")))?
501 .to_vec(),
502 );
503 }
504 "language" => {
505 language = field.text().await.ok().filter(|s| !s.is_empty());
506 }
507 _ => {} }
509 }
510
511 let data = file_data.ok_or_else(|| ServerError::BadRequest("missing 'file' field".into()))?;
512
513 let text = state
514 .engine
515 .transcribe_bytes(&data, language.as_deref())
516 .await
517 .map_err(|e| ServerError::InternalError(format!("transcribe: {e}")))?;
518
519 Ok(Json(TranscriptionResponse { text }).into_response())
520}
521
522async fn speech_handler(
524 State(state): State<AppState>,
525 Json(request): Json<SpeechRequest>,
526) -> std::result::Result<Response, ServerError> {
527 let span = span!(Level::INFO, "speech");
528 let _guard = span.enter();
529
530 let language = if request.language.is_empty() || request.language == "auto" {
531 None
532 } else {
533 Some(request.language.as_str())
534 };
535
536 let chunk_frames = 10usize;
537 let sample_rate = state.engine.tts_sample_rate();
538
539 if request.stream {
540 let (tx, rx) =
542 mpsc::unbounded_channel::<std::result::Result<axum::body::Bytes, std::io::Error>>();
543
544 let engine = state.engine.clone();
545 let text = request.input.clone();
546 let lang = request.language.clone();
547
548 tokio::task::spawn_blocking(move || {
549 let lang_opt = if lang.is_empty() || lang == "auto" {
550 None
551 } else {
552 Some(lang.as_str())
553 };
554 let rt = tokio::runtime::Handle::current();
555
556 match rt.block_on(engine.synthesize_speech(&text, lang_opt, chunk_frames)) {
557 Ok(chunks) => {
558 for chunk in &chunks {
559 let wav_bytes = pcm_to_wav_bytes(chunk, sample_rate);
560 let _ = tx.send(Ok(axum::body::Bytes::from(wav_bytes)));
561 }
562 }
563 Err(e) => {
564 error!("TTS error: {e}");
565 }
566 }
567 });
568
569 let stream = tokio_stream::wrappers::UnboundedReceiverStream::new(rx);
570 let body = axum::body::Body::from_stream(stream);
571 Ok(Response::builder()
572 .status(200)
573 .header("content-type", "audio/wav")
574 .header("transfer-encoding", "chunked")
575 .body(body)
576 .unwrap())
577 } else {
578 let chunks = state
580 .engine
581 .synthesize_speech(&request.input, language, chunk_frames)
582 .await
583 .map_err(|e| ServerError::InternalError(format!("TTS: {e}")))?;
584
585 let all_samples: Vec<f32> = chunks.into_iter().flatten().collect();
586 let wav_bytes = pcm_to_wav_bytes(&all_samples, sample_rate);
587
588 Ok(Response::builder()
589 .status(200)
590 .header("content-type", "audio/wav")
591 .header("content-length", wav_bytes.len().to_string())
592 .body(axum::body::Body::from(wav_bytes))
593 .unwrap())
594 }
595}
596
597fn pcm_to_wav_bytes(samples: &[f32], sample_rate: u32) -> Vec<u8> {
599 let num_samples = samples.len();
600 let data_size = num_samples * 2; let file_size = 44 + data_size;
602
603 let mut buf = Vec::with_capacity(file_size);
604 buf.extend_from_slice(b"RIFF");
606 buf.extend_from_slice(&((file_size - 8) as u32).to_le_bytes());
607 buf.extend_from_slice(b"WAVE");
608 buf.extend_from_slice(b"fmt ");
610 buf.extend_from_slice(&16u32.to_le_bytes()); buf.extend_from_slice(&1u16.to_le_bytes()); buf.extend_from_slice(&1u16.to_le_bytes()); buf.extend_from_slice(&sample_rate.to_le_bytes());
614 buf.extend_from_slice(&(sample_rate * 2).to_le_bytes()); buf.extend_from_slice(&2u16.to_le_bytes()); buf.extend_from_slice(&16u16.to_le_bytes()); buf.extend_from_slice(b"data");
619 buf.extend_from_slice(&(data_size as u32).to_le_bytes());
620 for &s in samples {
621 let i16_val = (s.clamp(-1.0, 1.0) * 32767.0) as i16;
622 buf.extend_from_slice(&i16_val.to_le_bytes());
623 }
624 buf
625}
626
627async fn models_handler(
628 State(state): State<AppState>,
629) -> std::result::Result<Response, ServerError> {
630 let status = state.engine.status().await;
631 let now = chrono::Utc::now().timestamp() as u64;
632 let data = status
633 .loaded_models
634 .into_iter()
635 .map(|model_id| crate::openai::ModelInfo {
636 id: model_id.to_string(),
637 object: "model".to_string(),
638 created: now,
639 owned_by: "ferrum".to_string(),
640 permission: vec![],
641 root: None,
642 parent: None,
643 })
644 .collect();
645
646 let models = ModelListResponse {
647 object: "list".to_string(),
648 data,
649 };
650
651 Ok(Json(models).into_response())
652}
653
654async fn health_handler(
655 State(state): State<AppState>,
656) -> std::result::Result<Response, ServerError> {
657 let engine_status = state.engine.status().await;
658 let scheduler_metrics = state.engine.metrics();
659
660 let health = serde_json::json!({
661 "status": "healthy",
662 "timestamp": chrono::Utc::now().to_rfc3339(),
663 "version": env!("CARGO_PKG_VERSION"),
664 "engine": {
665 "active_requests": engine_status.active_requests,
666 "queued_requests": engine_status.queued_requests,
667 },
668 "scheduler": {
669 "total_requests": scheduler_metrics.total_requests,
670 "successful_requests": scheduler_metrics.successful_requests,
671 "failed_requests": scheduler_metrics.failed_requests,
672 "throughput_rps": scheduler_metrics.throughput_rps,
673 }
674 });
675
676 Ok(Json(health).into_response())
677}
678
679async fn metrics_handler() -> std::result::Result<Response, ServerError> {
681 let body = match PROM_HANDLE.get() {
682 Some(handle) => handle.render(),
683 None => "# Prometheus recorder not initialized\n".to_string(),
684 };
685
686 Ok((
687 [(
688 axum::http::header::CONTENT_TYPE,
689 "text/plain; version=0.0.4; charset=utf-8",
690 )],
691 body,
692 )
693 .into_response())
694}
695
696async fn root_handler() -> std::result::Result<Response, ServerError> {
697 let info = serde_json::json!({
698 "name": "Ferrum Inference Server",
699 "version": env!("CARGO_PKG_VERSION"),
700 "api_version": "v1",
701 "status": "running"
702 });
703
704 Ok(Json(info).into_response())
705}
706
707#[derive(Debug)]
709enum ServerError {
710 BadRequest(String),
711 InternalError(String),
712 NotImplemented(String),
713}
714
715impl IntoResponse for ServerError {
716 fn into_response(self) -> Response {
717 let (status, message) = match self {
718 ServerError::BadRequest(msg) => (AxumStatusCode::BAD_REQUEST, msg),
719 ServerError::InternalError(msg) => (AxumStatusCode::INTERNAL_SERVER_ERROR, msg),
720 ServerError::NotImplemented(msg) => (AxumStatusCode::NOT_IMPLEMENTED, msg),
721 };
722
723 let error = OpenAiError {
724 error: OpenAiErrorDetail {
725 message,
726 error_type: "server_error".to_string(),
727 param: None,
728 code: None,
729 },
730 };
731
732 (status, Json(error)).into_response()
733 }
734}
735
736impl std::fmt::Display for MessageRole {
737 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
738 match self {
739 MessageRole::System => write!(f, "system"),
740 MessageRole::User => write!(f, "user"),
741 MessageRole::Assistant => write!(f, "assistant"),
742 MessageRole::Function => write!(f, "function"),
743 }
744 }
745}
746
747fn finish_reason_to_string(reason: &FinishReason) -> String {
749 match reason {
750 FinishReason::Length => "length".to_string(),
751 FinishReason::Stop => "stop".to_string(),
752 FinishReason::EOS => "stop".to_string(),
753 FinishReason::Cancelled => "cancelled".to_string(),
754 FinishReason::Error => "error".to_string(),
755 FinishReason::ContentFilter => "content_filter".to_string(),
756 }
757}