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/models", get(models_handler))
74 .route("/health", get(health_handler))
76 .route("/metrics", get(metrics_handler))
77 .route("/", get(root_handler))
78 .layer(
80 ServiceBuilder::new()
81 .layer(TraceLayer::new_for_http())
82 .layer(CorsLayer::permissive()), )
84 .with_state(app_state)
85 }
86}
87
88#[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 Ok(())
118 }
119
120 fn is_running(&self) -> bool {
121 true
123 }
124
125 fn address(&self) -> Option<std::net::SocketAddr> {
126 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 unimplemented!("Dynamic handler registration not implemented in MVP")
140 }
141
142 fn register_middleware(&mut self, _middleware: Box<dyn crate::traits::Middleware>) {
143 unimplemented!("Dynamic middleware registration not implemented in MVP")
145 }
146
147 fn get_metrics(&self) -> ServerMetrics {
148 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
169async 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 let inference_request =
185 convert_chat_request(&request).map_err(|e| ServerError::BadRequest(e.to_string()))?;
186
187 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
195async 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 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 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 if token_count >= max_tokens || chunk.finish_reason.is_some() {
249 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
304async 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
346fn convert_chat_request(
348 request: &ChatCompletionsRequest,
349) -> ferrum_types::Result<InferenceRequest> {
350 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, 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, client_id: None,
386 session_id: None,
387 created_at: chrono::Utc::now(),
388 metadata: std::collections::HashMap::new(),
389 })
390}
391
392async fn completions_handler(
394 State(_state): State<AppState>,
395 Json(_request): Json<CompletionsRequest>,
396) -> std::result::Result<Response, ServerError> {
397 Err(ServerError::NotImplemented(
399 "Legacy completions not implemented in MVP".to_string(),
400 ))
401}
402
403async 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 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
527async 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#[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
595fn 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}