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/models", get(models_handler))
73 .route("/health", get(health_handler))
75 .route("/metrics", get(metrics_handler))
76 .route("/", get(root_handler))
77 .layer(
79 ServiceBuilder::new()
80 .layer(TraceLayer::new_for_http())
81 .layer(CorsLayer::permissive()), )
83 .with_state(app_state)
84 }
85}
86
87#[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 Ok(())
117 }
118
119 fn is_running(&self) -> bool {
120 true
122 }
123
124 fn address(&self) -> Option<std::net::SocketAddr> {
125 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 unimplemented!("Dynamic handler registration not implemented in MVP")
139 }
140
141 fn register_middleware(&mut self, _middleware: Box<dyn crate::traits::Middleware>) {
142 unimplemented!("Dynamic middleware registration not implemented in MVP")
144 }
145
146 fn get_metrics(&self) -> ServerMetrics {
147 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
168async 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 let inference_request =
184 convert_chat_request(&request).map_err(|e| ServerError::BadRequest(e.to_string()))?;
185
186 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
194async 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 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 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 if token_count >= max_tokens || chunk.finish_reason.is_some() {
248 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
303async 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
345fn convert_chat_request(
347 request: &ChatCompletionsRequest,
348) -> ferrum_types::Result<InferenceRequest> {
349 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, 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, client_id: None,
385 session_id: None,
386 created_at: chrono::Utc::now(),
387 metadata: std::collections::HashMap::new(),
388 })
389}
390
391async fn completions_handler(
393 State(_state): State<AppState>,
394 Json(_request): Json<CompletionsRequest>,
395) -> std::result::Result<Response, ServerError> {
396 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
454async 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#[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
522fn 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}