use crate::{openai::*, traits::HttpServer, types::*};
use async_trait::async_trait;
use axum::{
extract::State,
http::StatusCode as AxumStatusCode,
response::{sse::Event, IntoResponse, Response, Sse},
routing::{get, post},
Json, Router,
};
use ferrum_interfaces::engine::InferenceEngine;
use ferrum_types::{
FerrumError as Error, FinishReason, InferenceRequest, ModelId, Priority, RequestId,
SamplingParams,
};
use std::sync::Arc;
use tokio::sync::mpsc;
use tokio_stream::StreamExt;
use tower::ServiceBuilder;
use tower_http::{cors::CorsLayer, trace::TraceLayer};
use tracing::{debug, error, info, span, Level};
use uuid::Uuid;
static PROM_HANDLE: std::sync::OnceLock<metrics_exporter_prometheus::PrometheusHandle> =
std::sync::OnceLock::new();
pub fn init_prometheus_recorder() {
PROM_HANDLE.get_or_init(|| {
let builder = metrics_exporter_prometheus::PrometheusBuilder::new();
let handle = builder
.install_recorder()
.expect("Failed to install Prometheus recorder");
info!("Prometheus metrics recorder installed");
handle
});
}
pub struct AxumServer {
engine: Arc<dyn InferenceEngine + Send + Sync>,
config: ServerConfig,
}
impl AxumServer {
pub fn new(engine: Arc<dyn InferenceEngine + Send + Sync>) -> Self {
Self {
engine,
config: ServerConfig::default(),
}
}
fn build_router(&self) -> Router {
let app_state = AppState {
engine: self.engine.clone(),
};
Router::new()
.route("/v1/chat/completions", post(chat_completions_handler))
.route("/v1/completions", post(completions_handler))
.route("/v1/models", get(models_handler))
.route("/health", get(health_handler))
.route("/metrics", get(metrics_handler))
.route("/", get(root_handler))
.layer(
ServiceBuilder::new()
.layer(TraceLayer::new_for_http())
.layer(CorsLayer::permissive()), )
.with_state(app_state)
}
}
#[derive(Clone)]
struct AppState {
engine: Arc<dyn InferenceEngine + Send + Sync>,
}
#[async_trait]
impl HttpServer for AxumServer {
async fn start(&self, config: &ServerConfig) -> ferrum_types::Result<()> {
let addr = format!("{}:{}", config.host, config.port);
info!("Starting Axum server on {}", addr);
let app = self.build_router();
let listener = tokio::net::TcpListener::bind(&addr)
.await
.map_err(|e| Error::internal(format!("Failed to bind to {}: {}", addr, e)))?;
info!("Server listening on {}", addr);
axum::serve(listener, app)
.await
.map_err(|e| Error::internal(format!("Server error: {}", e)))?;
Ok(())
}
async fn stop(&self, _timeout: std::time::Duration) -> ferrum_types::Result<()> {
info!("Stopping Axum server");
Ok(())
}
fn is_running(&self) -> bool {
true
}
fn address(&self) -> Option<std::net::SocketAddr> {
format!("{}:{}", self.config.host, self.config.port)
.parse()
.ok()
}
fn register_handler(
&mut self,
_path: &str,
_method: HttpMethod,
_handler: Box<dyn crate::traits::RequestHandler>,
) {
unimplemented!("Dynamic handler registration not implemented in MVP")
}
fn register_middleware(&mut self, _middleware: Box<dyn crate::traits::Middleware>) {
unimplemented!("Dynamic middleware registration not implemented in MVP")
}
fn get_metrics(&self) -> ServerMetrics {
ServerMetrics {
total_requests: 0,
requests_by_endpoint: std::collections::HashMap::new(),
requests_by_status: std::collections::HashMap::new(),
avg_response_time_ms: 0.0,
p95_response_time_ms: 0.0,
p99_response_time_ms: 0.0,
active_connections: 0,
bytes_sent: 0,
bytes_received: 0,
error_rate: 0.0,
uptime_seconds: 0,
}
}
async fn health_check(&self) -> HealthStatus {
HealthStatus::Healthy
}
}
async fn chat_completions_handler(
State(state): State<AppState>,
Json(request): Json<ChatCompletionsRequest>,
) -> std::result::Result<Response, ServerError> {
let span = span!(Level::INFO, "chat_completions", model = %request.model);
let _enter = span.enter();
info!(
"Received chat completions request for model: {}",
request.model
);
debug!("Request: {:?}", request);
let inference_request =
convert_chat_request(&request).map_err(|e| ServerError::BadRequest(e.to_string()))?;
if request.stream.unwrap_or(false) {
handle_chat_completions_stream(state, request, inference_request).await
} else {
handle_chat_completions_sync(state, request, inference_request).await
}
}
async fn handle_chat_completions_stream(
state: AppState,
openai_request: ChatCompletionsRequest,
inference_request: InferenceRequest,
) -> std::result::Result<Response, ServerError> {
let (tx, rx) = mpsc::unbounded_channel::<std::result::Result<Event, axum::Error>>();
let engine = state.engine.clone();
let request_id = Uuid::new_v4().to_string();
let prompt_tokens = inference_request.prompt.split_whitespace().count() as u32;
tokio::spawn(async move {
let mut current_text = String::new();
let mut token_count = 0;
let max_tokens = openai_request.max_tokens.unwrap_or(100);
match engine.infer_stream(inference_request).await {
Ok(mut stream) => {
while let Some(result) = stream.next().await {
match result {
Ok(chunk) => {
current_text.push_str(&chunk.text);
token_count += 1;
let response_chunk = ChatCompletionsResponse {
id: request_id.clone(),
object: "chat.completion.chunk".to_string(),
created: chrono::Utc::now().timestamp() as u64,
model: openai_request.model.clone(),
choices: vec![ChatChoice {
index: 0,
message: None,
delta: Some(ChatMessage {
role: MessageRole::Assistant,
content: chunk.text.clone(),
name: None,
}),
finish_reason: None,
}],
usage: None,
};
let sse_event = Event::default()
.json_data(&response_chunk)
.unwrap_or_else(|_| Event::default().data("error"));
if tx.send(Ok(sse_event)).is_err() {
break;
}
if token_count >= max_tokens || chunk.finish_reason.is_some() {
let final_chunk = ChatCompletionsResponse {
id: request_id.clone(),
object: "chat.completion.chunk".to_string(),
created: chrono::Utc::now().timestamp() as u64,
model: openai_request.model.clone(),
choices: vec![ChatChoice {
index: 0,
message: None,
delta: None,
finish_reason: chunk
.finish_reason
.as_ref()
.map(finish_reason_to_string),
}],
usage: Some(Usage {
prompt_tokens: prompt_tokens,
completion_tokens: token_count,
total_tokens: prompt_tokens + token_count,
}),
};
let final_event = Event::default()
.json_data(&final_chunk)
.unwrap_or_else(|_| Event::default().data("error"));
let _ = tx.send(Ok(final_event));
let _ = tx.send(Ok(Event::default().data("[DONE]")));
break;
}
}
Err(e) => {
error!("Stream generation error: {}", e);
let _ = tx.send(Ok(
Event::default().data(&format!("{{\"error\": \"{}\"}}", e))
));
break;
}
}
}
}
Err(e) => {
error!("Failed to start streaming: {}", e);
let _ = tx.send(Ok(
Event::default().data(&format!("{{\"error\": \"{}\"}}", e))
));
}
}
});
let stream = tokio_stream::wrappers::UnboundedReceiverStream::new(rx);
let sse_stream = Sse::new(stream);
Ok(sse_stream.into_response())
}
async fn handle_chat_completions_sync(
state: AppState,
openai_request: ChatCompletionsRequest,
inference_request: InferenceRequest,
) -> std::result::Result<Response, ServerError> {
info!("Processing non-streaming chat completion");
let prompt_tokens = inference_request.prompt.split_whitespace().count() as u32;
match state.engine.infer(inference_request).await {
Ok(output) => {
let response = ChatCompletionsResponse {
id: Uuid::new_v4().to_string(),
object: "chat.completion".to_string(),
created: chrono::Utc::now().timestamp() as u64,
model: openai_request.model,
choices: vec![ChatChoice {
index: 0,
message: Some(ChatMessage {
role: MessageRole::Assistant,
content: output.text,
name: None,
}),
delta: None,
finish_reason: Some(finish_reason_to_string(&output.finish_reason)),
}],
usage: Some(Usage {
prompt_tokens: prompt_tokens,
completion_tokens: output.tokens.len() as u32,
total_tokens: prompt_tokens + output.tokens.len() as u32,
}),
};
Ok(Json(response).into_response())
}
Err(e) => {
error!("Generation failed: {}", e);
Err(ServerError::InternalError(e.to_string()))
}
}
}
fn convert_chat_request(
request: &ChatCompletionsRequest,
) -> ferrum_types::Result<InferenceRequest> {
let prompt = request
.messages
.iter()
.map(|msg| format!("{}: {}", msg.role.to_string(), msg.content))
.collect::<Vec<_>>()
.join("\n");
Ok(InferenceRequest {
id: RequestId(Uuid::new_v4()),
model_id: ModelId(request.model.clone()),
prompt,
sampling_params: SamplingParams {
max_tokens: request.max_tokens.unwrap_or(100) as usize,
temperature: request.temperature.unwrap_or(1.0),
top_p: request.top_p.unwrap_or(1.0),
top_k: None, repetition_penalty: 1.0,
presence_penalty: request.presence_penalty.unwrap_or(0.0),
frequency_penalty: request.frequency_penalty.unwrap_or(0.0),
stop_sequences: request.stop.clone().unwrap_or_default(),
seed: request.seed,
min_p: None,
tfs: None,
typical_p: None,
mirostat: None,
response_format: match &request.response_format {
Some(rf) if rf.format_type == "json_object" => {
ferrum_types::ResponseFormat::JsonObject
}
_ => ferrum_types::ResponseFormat::Text,
},
},
stream: request.stream.unwrap_or(false),
priority: Priority::Normal, client_id: None,
session_id: None,
created_at: chrono::Utc::now(),
metadata: std::collections::HashMap::new(),
})
}
async fn completions_handler(
State(_state): State<AppState>,
Json(_request): Json<CompletionsRequest>,
) -> std::result::Result<Response, ServerError> {
Err(ServerError::NotImplemented(
"Legacy completions not implemented in MVP".to_string(),
))
}
async fn models_handler(
State(state): State<AppState>,
) -> std::result::Result<Response, ServerError> {
let status = state.engine.status().await;
let now = chrono::Utc::now().timestamp() as u64;
let data = status
.loaded_models
.into_iter()
.map(|model_id| crate::openai::ModelInfo {
id: model_id.to_string(),
object: "model".to_string(),
created: now,
owned_by: "ferrum".to_string(),
permission: vec![],
root: None,
parent: None,
})
.collect();
let models = ModelListResponse {
object: "list".to_string(),
data,
};
Ok(Json(models).into_response())
}
async fn health_handler(
State(state): State<AppState>,
) -> std::result::Result<Response, ServerError> {
let engine_status = state.engine.status().await;
let scheduler_metrics = state.engine.metrics();
let health = serde_json::json!({
"status": "healthy",
"timestamp": chrono::Utc::now().to_rfc3339(),
"version": env!("CARGO_PKG_VERSION"),
"engine": {
"active_requests": engine_status.active_requests,
"queued_requests": engine_status.queued_requests,
},
"scheduler": {
"total_requests": scheduler_metrics.total_requests,
"successful_requests": scheduler_metrics.successful_requests,
"failed_requests": scheduler_metrics.failed_requests,
"throughput_rps": scheduler_metrics.throughput_rps,
}
});
Ok(Json(health).into_response())
}
async fn metrics_handler() -> std::result::Result<Response, ServerError> {
let body = match PROM_HANDLE.get() {
Some(handle) => handle.render(),
None => "# Prometheus recorder not initialized\n".to_string(),
};
Ok((
[(
axum::http::header::CONTENT_TYPE,
"text/plain; version=0.0.4; charset=utf-8",
)],
body,
)
.into_response())
}
async fn root_handler() -> std::result::Result<Response, ServerError> {
let info = serde_json::json!({
"name": "Ferrum Inference Server",
"version": env!("CARGO_PKG_VERSION"),
"api_version": "v1",
"status": "running"
});
Ok(Json(info).into_response())
}
#[derive(Debug)]
enum ServerError {
BadRequest(String),
InternalError(String),
NotImplemented(String),
}
impl IntoResponse for ServerError {
fn into_response(self) -> Response {
let (status, message) = match self {
ServerError::BadRequest(msg) => (AxumStatusCode::BAD_REQUEST, msg),
ServerError::InternalError(msg) => (AxumStatusCode::INTERNAL_SERVER_ERROR, msg),
ServerError::NotImplemented(msg) => (AxumStatusCode::NOT_IMPLEMENTED, msg),
};
let error = OpenAiError {
error: OpenAiErrorDetail {
message,
error_type: "server_error".to_string(),
param: None,
code: None,
},
};
(status, Json(error)).into_response()
}
}
impl std::fmt::Display for MessageRole {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
MessageRole::System => write!(f, "system"),
MessageRole::User => write!(f, "user"),
MessageRole::Assistant => write!(f, "assistant"),
MessageRole::Function => write!(f, "function"),
}
}
}
fn finish_reason_to_string(reason: &FinishReason) -> String {
match reason {
FinishReason::Length => "length".to_string(),
FinishReason::Stop => "stop".to_string(),
FinishReason::EOS => "stop".to_string(),
FinishReason::Cancelled => "cancelled".to_string(),
FinishReason::Error => "error".to_string(),
FinishReason::ContentFilter => "content_filter".to_string(),
}
}