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