Skip to main content

ferrum_server/
axum_server.rs

1//! Axum-based HTTP server implementation for Ferrum
2//!
3//! This module provides a concrete implementation of the HttpServer trait
4//! using the Axum web framework, with full OpenAI API compatibility.
5
6use 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
28/// Shared Prometheus recorder handle for rendering metrics.
29static PROM_HANDLE: std::sync::OnceLock<metrics_exporter_prometheus::PrometheusHandle> =
30    std::sync::OnceLock::new();
31
32/// Initialize the Prometheus metrics recorder.
33///
34/// Must be called once before any `metrics::counter!()` / `histogram!()` calls.
35/// Safe to call multiple times — subsequent calls are no-ops.
36pub 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
47/// Axum-based server implementation
48pub struct AxumServer {
49    engine: Arc<dyn InferenceEngine + Send + Sync>,
50    config: ServerConfig,
51}
52
53impl AxumServer {
54    /// Create a new Axum server
55    pub fn new(engine: Arc<dyn InferenceEngine + Send + Sync>) -> Self {
56        Self {
57            engine,
58            config: ServerConfig::default(),
59        }
60    }
61
62    /// Build the router with all routes
63    fn build_router(&self) -> Router {
64        let app_state = AppState {
65            engine: self.engine.clone(),
66        };
67
68        Router::new()
69            // OpenAI API routes
70            .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/models", get(models_handler))
74            // Health & observability
75            .route("/health", get(health_handler))
76            .route("/metrics", get(metrics_handler))
77            .route("/", get(root_handler))
78            // Apply middleware
79            .layer(
80                ServiceBuilder::new()
81                    .layer(TraceLayer::new_for_http())
82                    .layer(CorsLayer::permissive()), // For MVP, allow all origins
83            )
84            .with_state(app_state)
85    }
86}
87
88/// Application state shared across handlers
89#[derive(Clone)]
90struct AppState {
91    engine: Arc<dyn InferenceEngine + Send + Sync>,
92}
93
94#[async_trait]
95impl HttpServer for AxumServer {
96    async fn start(&self, config: &ServerConfig) -> ferrum_types::Result<()> {
97        let addr = format!("{}:{}", config.host, config.port);
98        info!("Starting Axum server on {}", addr);
99
100        let app = self.build_router();
101        let listener = tokio::net::TcpListener::bind(&addr)
102            .await
103            .map_err(|e| Error::internal(format!("Failed to bind to {}: {}", addr, e)))?;
104
105        info!("Server listening on {}", addr);
106
107        axum::serve(listener, app)
108            .await
109            .map_err(|e| Error::internal(format!("Server error: {}", e)))?;
110
111        Ok(())
112    }
113
114    async fn stop(&self, _timeout: std::time::Duration) -> ferrum_types::Result<()> {
115        info!("Stopping Axum server");
116        // Axum doesn't have explicit stop - server stops when task is cancelled
117        Ok(())
118    }
119
120    fn is_running(&self) -> bool {
121        // For MVP, always return true when server object exists
122        true
123    }
124
125    fn address(&self) -> Option<std::net::SocketAddr> {
126        // For MVP, return configured address
127        format!("{}:{}", self.config.host, self.config.port)
128            .parse()
129            .ok()
130    }
131
132    fn register_handler(
133        &mut self,
134        _path: &str,
135        _method: HttpMethod,
136        _handler: Box<dyn crate::traits::RequestHandler>,
137    ) {
138        // For MVP, routes are static
139        unimplemented!("Dynamic handler registration not implemented in MVP")
140    }
141
142    fn register_middleware(&mut self, _middleware: Box<dyn crate::traits::Middleware>) {
143        // For MVP, middleware is static
144        unimplemented!("Dynamic middleware registration not implemented in MVP")
145    }
146
147    fn get_metrics(&self) -> ServerMetrics {
148        // Return empty metrics for MVP
149        ServerMetrics {
150            total_requests: 0,
151            requests_by_endpoint: std::collections::HashMap::new(),
152            requests_by_status: std::collections::HashMap::new(),
153            avg_response_time_ms: 0.0,
154            p95_response_time_ms: 0.0,
155            p99_response_time_ms: 0.0,
156            active_connections: 0,
157            bytes_sent: 0,
158            bytes_received: 0,
159            error_rate: 0.0,
160            uptime_seconds: 0,
161        }
162    }
163
164    async fn health_check(&self) -> HealthStatus {
165        HealthStatus::Healthy
166    }
167}
168
169/// Main chat completions handler
170async fn chat_completions_handler(
171    State(state): State<AppState>,
172    Json(request): Json<ChatCompletionsRequest>,
173) -> std::result::Result<Response, ServerError> {
174    let span = span!(Level::INFO, "chat_completions", model = %request.model);
175    let _enter = span.enter();
176
177    info!(
178        "Received chat completions request for model: {}",
179        request.model
180    );
181    debug!("Request: {:?}", request);
182
183    // Convert OpenAI request to internal format
184    let inference_request =
185        convert_chat_request(&request).map_err(|e| ServerError::BadRequest(e.to_string()))?;
186
187    // Check if streaming is requested
188    if request.stream.unwrap_or(false) {
189        handle_chat_completions_stream(state, request, inference_request).await
190    } else {
191        handle_chat_completions_sync(state, request, inference_request).await
192    }
193}
194
195/// Handle streaming chat completions
196async fn handle_chat_completions_stream(
197    state: AppState,
198    openai_request: ChatCompletionsRequest,
199    inference_request: InferenceRequest,
200) -> std::result::Result<Response, ServerError> {
201    let (tx, rx) = mpsc::unbounded_channel::<std::result::Result<Event, axum::Error>>();
202
203    // Spawn task to generate tokens
204    let engine = state.engine.clone();
205    let request_id = Uuid::new_v4().to_string();
206    let prompt_tokens = inference_request.prompt.split_whitespace().count() as u32;
207
208    tokio::spawn(async move {
209        let mut current_text = String::new();
210        let mut token_count = 0;
211        let max_tokens = openai_request.max_tokens.unwrap_or(100);
212
213        match engine.infer_stream(inference_request).await {
214            Ok(mut stream) => {
215                while let Some(result) = stream.next().await {
216                    match result {
217                        Ok(chunk) => {
218                            current_text.push_str(&chunk.text);
219                            token_count += 1;
220
221                            // Create streaming response chunk
222                            let response_chunk = ChatCompletionsResponse {
223                                id: request_id.clone(),
224                                object: "chat.completion.chunk".to_string(),
225                                created: chrono::Utc::now().timestamp() as u64,
226                                model: openai_request.model.clone(),
227                                choices: vec![ChatChoice {
228                                    index: 0,
229                                    message: None,
230                                    delta: Some(ChatMessage {
231                                        role: MessageRole::Assistant,
232                                        content: chunk.text.clone(),
233                                        name: None,
234                                    }),
235                                    finish_reason: None,
236                                }],
237                                usage: None,
238                            };
239
240                            let sse_event = Event::default()
241                                .json_data(&response_chunk)
242                                .unwrap_or_else(|_| Event::default().data("error"));
243                            if tx.send(Ok(sse_event)).is_err() {
244                                break;
245                            }
246
247                            // Check stopping conditions
248                            if token_count >= max_tokens || chunk.finish_reason.is_some() {
249                                // Send final chunk
250                                let final_chunk = ChatCompletionsResponse {
251                                    id: request_id.clone(),
252                                    object: "chat.completion.chunk".to_string(),
253                                    created: chrono::Utc::now().timestamp() as u64,
254                                    model: openai_request.model.clone(),
255                                    choices: vec![ChatChoice {
256                                        index: 0,
257                                        message: None,
258                                        delta: None,
259                                        finish_reason: chunk
260                                            .finish_reason
261                                            .as_ref()
262                                            .map(finish_reason_to_string),
263                                    }],
264                                    usage: Some(Usage {
265                                        prompt_tokens: prompt_tokens,
266                                        completion_tokens: token_count,
267                                        total_tokens: prompt_tokens + token_count,
268                                    }),
269                                };
270
271                                let final_event = Event::default()
272                                    .json_data(&final_chunk)
273                                    .unwrap_or_else(|_| Event::default().data("error"));
274                                let _ = tx.send(Ok(final_event));
275                                let _ = tx.send(Ok(Event::default().data("[DONE]")));
276                                break;
277                            }
278                        }
279                        Err(e) => {
280                            error!("Stream generation error: {}", e);
281                            let _ = tx.send(Ok(
282                                Event::default().data(&format!("{{\"error\": \"{}\"}}", e))
283                            ));
284                            break;
285                        }
286                    }
287                }
288            }
289            Err(e) => {
290                error!("Failed to start streaming: {}", e);
291                let _ = tx.send(Ok(
292                    Event::default().data(&format!("{{\"error\": \"{}\"}}", e))
293                ));
294            }
295        }
296    });
297
298    let stream = tokio_stream::wrappers::UnboundedReceiverStream::new(rx);
299    let sse_stream = Sse::new(stream);
300
301    Ok(sse_stream.into_response())
302}
303
304/// Handle non-streaming chat completions
305async fn handle_chat_completions_sync(
306    state: AppState,
307    openai_request: ChatCompletionsRequest,
308    inference_request: InferenceRequest,
309) -> std::result::Result<Response, ServerError> {
310    info!("Processing non-streaming chat completion");
311
312    let prompt_tokens = inference_request.prompt.split_whitespace().count() as u32;
313    match state.engine.infer(inference_request).await {
314        Ok(output) => {
315            let response = ChatCompletionsResponse {
316                id: Uuid::new_v4().to_string(),
317                object: "chat.completion".to_string(),
318                created: chrono::Utc::now().timestamp() as u64,
319                model: openai_request.model,
320                choices: vec![ChatChoice {
321                    index: 0,
322                    message: Some(ChatMessage {
323                        role: MessageRole::Assistant,
324                        content: output.text,
325                        name: None,
326                    }),
327                    delta: None,
328                    finish_reason: Some(finish_reason_to_string(&output.finish_reason)),
329                }],
330                usage: Some(Usage {
331                    prompt_tokens: prompt_tokens,
332                    completion_tokens: output.tokens.len() as u32,
333                    total_tokens: prompt_tokens + output.tokens.len() as u32,
334                }),
335            };
336
337            Ok(Json(response).into_response())
338        }
339        Err(e) => {
340            error!("Generation failed: {}", e);
341            Err(ServerError::InternalError(e.to_string()))
342        }
343    }
344}
345
346/// Convert OpenAI chat request to internal inference request
347fn convert_chat_request(
348    request: &ChatCompletionsRequest,
349) -> ferrum_types::Result<InferenceRequest> {
350    // Combine all messages into a single prompt for MVP
351    let prompt = request
352        .messages
353        .iter()
354        .map(|msg| format!("{}: {}", msg.role.to_string(), msg.content))
355        .collect::<Vec<_>>()
356        .join("\n");
357
358    Ok(InferenceRequest {
359        id: RequestId(Uuid::new_v4()),
360        model_id: ModelId(request.model.clone()),
361        prompt,
362        sampling_params: SamplingParams {
363            max_tokens: request.max_tokens.unwrap_or(100) as usize,
364            temperature: request.temperature.unwrap_or(1.0),
365            top_p: request.top_p.unwrap_or(1.0),
366            top_k: None, // OpenAI doesn't use top-k
367            repetition_penalty: 1.0,
368            presence_penalty: request.presence_penalty.unwrap_or(0.0),
369            frequency_penalty: request.frequency_penalty.unwrap_or(0.0),
370            stop_sequences: request.stop.clone().unwrap_or_default(),
371            seed: request.seed,
372            min_p: None,
373            tfs: None,
374            typical_p: None,
375            mirostat: None,
376            response_format: match &request.response_format {
377                Some(rf) if rf.format_type == "json_object" => {
378                    ferrum_types::ResponseFormat::JsonObject
379                }
380                _ => ferrum_types::ResponseFormat::Text,
381            },
382        },
383        stream: request.stream.unwrap_or(false),
384        priority: Priority::Normal, // Default priority
385        client_id: None,
386        session_id: None,
387        created_at: chrono::Utc::now(),
388        metadata: std::collections::HashMap::new(),
389    })
390}
391
392/// Other handlers
393async fn completions_handler(
394    State(_state): State<AppState>,
395    Json(_request): Json<CompletionsRequest>,
396) -> std::result::Result<Response, ServerError> {
397    // TODO: Implement legacy completions endpoint
398    Err(ServerError::NotImplemented(
399        "Legacy completions not implemented in MVP".to_string(),
400    ))
401}
402
403/// Embeddings handler — text and image embedding via OpenAI-compatible API.
404async fn embeddings_handler(
405    State(state): State<AppState>,
406    Json(request): Json<EmbeddingsRequest>,
407) -> std::result::Result<Response, ServerError> {
408    let span = span!(Level::INFO, "embeddings", model = %request.model);
409    let _enter = span.enter();
410
411    // Flatten input into individual items
412    let items: Vec<EmbeddingItem> = match request.input {
413        EmbeddingInput::Single(text) => vec![EmbeddingItem {
414            text: Some(text),
415            image: None,
416        }],
417        EmbeddingInput::Batch(texts) => texts
418            .into_iter()
419            .map(|t| EmbeddingItem {
420                text: Some(t),
421                image: None,
422            })
423            .collect(),
424        EmbeddingInput::SingleObject(item) => vec![item],
425        EmbeddingInput::BatchObjects(items) => items,
426    };
427
428    if items.is_empty() {
429        return Err(ServerError::BadRequest("Empty input".to_string()));
430    }
431
432    let mut data = Vec::with_capacity(items.len());
433    let mut total_tokens = 0u32;
434
435    for (idx, item) in items.iter().enumerate() {
436        let embedding = if let Some(ref image) = item.image {
437            state
438                .engine
439                .embed_image(image)
440                .await
441                .map_err(|e| ServerError::InternalError(format!("embed_image: {e}")))?
442        } else if let Some(ref text) = item.text {
443            total_tokens += text.len() as u32;
444            state
445                .engine
446                .embed_text(text)
447                .await
448                .map_err(|e| ServerError::InternalError(format!("embed_text: {e}")))?
449        } else {
450            return Err(ServerError::BadRequest(
451                "Each input must have either 'text' or 'image'".to_string(),
452            ));
453        };
454
455        data.push(EmbeddingData {
456            object: "embedding".to_string(),
457            embedding,
458            index: idx,
459        });
460    }
461
462    let response = EmbeddingsResponse {
463        object: "list".to_string(),
464        data,
465        model: request.model,
466        usage: EmbeddingUsage {
467            prompt_tokens: total_tokens,
468            total_tokens,
469        },
470    };
471
472    Ok(Json(response).into_response())
473}
474
475async fn models_handler(
476    State(state): State<AppState>,
477) -> std::result::Result<Response, ServerError> {
478    let status = state.engine.status().await;
479    let now = chrono::Utc::now().timestamp() as u64;
480    let data = status
481        .loaded_models
482        .into_iter()
483        .map(|model_id| crate::openai::ModelInfo {
484            id: model_id.to_string(),
485            object: "model".to_string(),
486            created: now,
487            owned_by: "ferrum".to_string(),
488            permission: vec![],
489            root: None,
490            parent: None,
491        })
492        .collect();
493
494    let models = ModelListResponse {
495        object: "list".to_string(),
496        data,
497    };
498
499    Ok(Json(models).into_response())
500}
501
502async fn health_handler(
503    State(state): State<AppState>,
504) -> std::result::Result<Response, ServerError> {
505    let engine_status = state.engine.status().await;
506    let scheduler_metrics = state.engine.metrics();
507
508    let health = serde_json::json!({
509        "status": "healthy",
510        "timestamp": chrono::Utc::now().to_rfc3339(),
511        "version": env!("CARGO_PKG_VERSION"),
512        "engine": {
513            "active_requests": engine_status.active_requests,
514            "queued_requests": engine_status.queued_requests,
515        },
516        "scheduler": {
517            "total_requests": scheduler_metrics.total_requests,
518            "successful_requests": scheduler_metrics.successful_requests,
519            "failed_requests": scheduler_metrics.failed_requests,
520            "throughput_rps": scheduler_metrics.throughput_rps,
521        }
522    });
523
524    Ok(Json(health).into_response())
525}
526
527/// Prometheus metrics endpoint — returns metrics in Prometheus text format.
528async fn metrics_handler() -> std::result::Result<Response, ServerError> {
529    let body = match PROM_HANDLE.get() {
530        Some(handle) => handle.render(),
531        None => "# Prometheus recorder not initialized\n".to_string(),
532    };
533
534    Ok((
535        [(
536            axum::http::header::CONTENT_TYPE,
537            "text/plain; version=0.0.4; charset=utf-8",
538        )],
539        body,
540    )
541        .into_response())
542}
543
544async fn root_handler() -> std::result::Result<Response, ServerError> {
545    let info = serde_json::json!({
546        "name": "Ferrum Inference Server",
547        "version": env!("CARGO_PKG_VERSION"),
548        "api_version": "v1",
549        "status": "running"
550    });
551
552    Ok(Json(info).into_response())
553}
554
555/// Server error type for HTTP responses
556#[derive(Debug)]
557enum ServerError {
558    BadRequest(String),
559    InternalError(String),
560    NotImplemented(String),
561}
562
563impl IntoResponse for ServerError {
564    fn into_response(self) -> Response {
565        let (status, message) = match self {
566            ServerError::BadRequest(msg) => (AxumStatusCode::BAD_REQUEST, msg),
567            ServerError::InternalError(msg) => (AxumStatusCode::INTERNAL_SERVER_ERROR, msg),
568            ServerError::NotImplemented(msg) => (AxumStatusCode::NOT_IMPLEMENTED, msg),
569        };
570
571        let error = OpenAiError {
572            error: OpenAiErrorDetail {
573                message,
574                error_type: "server_error".to_string(),
575                param: None,
576                code: None,
577            },
578        };
579
580        (status, Json(error)).into_response()
581    }
582}
583
584impl std::fmt::Display for MessageRole {
585    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
586        match self {
587            MessageRole::System => write!(f, "system"),
588            MessageRole::User => write!(f, "user"),
589            MessageRole::Assistant => write!(f, "assistant"),
590            MessageRole::Function => write!(f, "function"),
591        }
592    }
593}
594
595/// Convert FinishReason to OpenAI API string
596fn finish_reason_to_string(reason: &FinishReason) -> String {
597    match reason {
598        FinishReason::Length => "length".to_string(),
599        FinishReason::Stop => "stop".to_string(),
600        FinishReason::EOS => "stop".to_string(),
601        FinishReason::Cancelled => "cancelled".to_string(),
602        FinishReason::Error => "error".to_string(),
603        FinishReason::ContentFilter => "content_filter".to_string(),
604    }
605}