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/audio/transcriptions", post(transcriptions_handler))
74            .route("/v1/audio/speech", post(speech_handler))
75            .route("/v1/models", get(models_handler))
76            // Health & observability
77            .route("/health", get(health_handler))
78            .route("/metrics", get(metrics_handler))
79            .route("/", get(root_handler))
80            // Apply middleware
81            .layer(
82                ServiceBuilder::new()
83                    .layer(TraceLayer::new_for_http())
84                    .layer(CorsLayer::permissive()), // For MVP, allow all origins
85            )
86            .with_state(app_state)
87    }
88}
89
90/// Application state shared across handlers
91#[derive(Clone)]
92struct AppState {
93    engine: Arc<dyn InferenceEngine + Send + Sync>,
94}
95
96#[async_trait]
97impl HttpServer for AxumServer {
98    async fn start(&self, config: &ServerConfig) -> ferrum_types::Result<()> {
99        let addr = format!("{}:{}", config.host, config.port);
100        info!("Starting Axum server on {}", addr);
101
102        let app = self.build_router();
103        let listener = tokio::net::TcpListener::bind(&addr)
104            .await
105            .map_err(|e| Error::internal(format!("Failed to bind to {}: {}", addr, e)))?;
106
107        info!("Server listening on {}", addr);
108
109        axum::serve(listener, app)
110            .await
111            .map_err(|e| Error::internal(format!("Server error: {}", e)))?;
112
113        Ok(())
114    }
115
116    async fn stop(&self, _timeout: std::time::Duration) -> ferrum_types::Result<()> {
117        info!("Stopping Axum server");
118        // Axum doesn't have explicit stop - server stops when task is cancelled
119        Ok(())
120    }
121
122    fn is_running(&self) -> bool {
123        // For MVP, always return true when server object exists
124        true
125    }
126
127    fn address(&self) -> Option<std::net::SocketAddr> {
128        // For MVP, return configured address
129        format!("{}:{}", self.config.host, self.config.port)
130            .parse()
131            .ok()
132    }
133
134    fn register_handler(
135        &mut self,
136        _path: &str,
137        _method: HttpMethod,
138        _handler: Box<dyn crate::traits::RequestHandler>,
139    ) {
140        // For MVP, routes are static
141        unimplemented!("Dynamic handler registration not implemented in MVP")
142    }
143
144    fn register_middleware(&mut self, _middleware: Box<dyn crate::traits::Middleware>) {
145        // For MVP, middleware is static
146        unimplemented!("Dynamic middleware registration not implemented in MVP")
147    }
148
149    fn get_metrics(&self) -> ServerMetrics {
150        // Return empty metrics for MVP
151        ServerMetrics {
152            total_requests: 0,
153            requests_by_endpoint: std::collections::HashMap::new(),
154            requests_by_status: std::collections::HashMap::new(),
155            avg_response_time_ms: 0.0,
156            p95_response_time_ms: 0.0,
157            p99_response_time_ms: 0.0,
158            active_connections: 0,
159            bytes_sent: 0,
160            bytes_received: 0,
161            error_rate: 0.0,
162            uptime_seconds: 0,
163        }
164    }
165
166    async fn health_check(&self) -> HealthStatus {
167        HealthStatus::Healthy
168    }
169}
170
171/// Main chat completions handler
172async fn chat_completions_handler(
173    State(state): State<AppState>,
174    Json(request): Json<ChatCompletionsRequest>,
175) -> std::result::Result<Response, ServerError> {
176    let span = span!(Level::INFO, "chat_completions", model = %request.model);
177    let _enter = span.enter();
178
179    info!(
180        "Received chat completions request for model: {}",
181        request.model
182    );
183    debug!("Request: {:?}", request);
184
185    // Convert OpenAI request to internal format
186    let inference_request =
187        convert_chat_request(&request).map_err(|e| ServerError::BadRequest(e.to_string()))?;
188
189    // Check if streaming is requested
190    if request.stream.unwrap_or(false) {
191        handle_chat_completions_stream(state, request, inference_request).await
192    } else {
193        handle_chat_completions_sync(state, request, inference_request).await
194    }
195}
196
197/// Handle streaming chat completions
198async fn handle_chat_completions_stream(
199    state: AppState,
200    openai_request: ChatCompletionsRequest,
201    inference_request: InferenceRequest,
202) -> std::result::Result<Response, ServerError> {
203    let (tx, rx) = mpsc::unbounded_channel::<std::result::Result<Event, axum::Error>>();
204
205    // Spawn task to generate tokens
206    let engine = state.engine.clone();
207    let request_id = Uuid::new_v4().to_string();
208    let prompt_tokens = inference_request.prompt.split_whitespace().count() as u32;
209
210    tokio::spawn(async move {
211        let mut current_text = String::new();
212        let mut token_count = 0;
213        let max_tokens = openai_request.max_tokens.unwrap_or(100);
214
215        match engine.infer_stream(inference_request).await {
216            Ok(mut stream) => {
217                while let Some(result) = stream.next().await {
218                    match result {
219                        Ok(chunk) => {
220                            current_text.push_str(&chunk.text);
221                            token_count += 1;
222
223                            // Create streaming response chunk
224                            let response_chunk = ChatCompletionsResponse {
225                                id: request_id.clone(),
226                                object: "chat.completion.chunk".to_string(),
227                                created: chrono::Utc::now().timestamp() as u64,
228                                model: openai_request.model.clone(),
229                                choices: vec![ChatChoice {
230                                    index: 0,
231                                    message: None,
232                                    delta: Some(ChatMessage {
233                                        role: MessageRole::Assistant,
234                                        content: chunk.text.clone(),
235                                        name: None,
236                                    }),
237                                    finish_reason: None,
238                                }],
239                                usage: None,
240                            };
241
242                            let sse_event = Event::default()
243                                .json_data(&response_chunk)
244                                .unwrap_or_else(|_| Event::default().data("error"));
245                            if tx.send(Ok(sse_event)).is_err() {
246                                break;
247                            }
248
249                            // Check stopping conditions
250                            if token_count >= max_tokens || chunk.finish_reason.is_some() {
251                                // Send final chunk
252                                let final_chunk = ChatCompletionsResponse {
253                                    id: request_id.clone(),
254                                    object: "chat.completion.chunk".to_string(),
255                                    created: chrono::Utc::now().timestamp() as u64,
256                                    model: openai_request.model.clone(),
257                                    choices: vec![ChatChoice {
258                                        index: 0,
259                                        message: None,
260                                        delta: None,
261                                        finish_reason: chunk
262                                            .finish_reason
263                                            .as_ref()
264                                            .map(finish_reason_to_string),
265                                    }],
266                                    usage: Some(Usage {
267                                        prompt_tokens,
268                                        completion_tokens: token_count,
269                                        total_tokens: prompt_tokens + token_count,
270                                    }),
271                                };
272
273                                let final_event = Event::default()
274                                    .json_data(&final_chunk)
275                                    .unwrap_or_else(|_| Event::default().data("error"));
276                                let _ = tx.send(Ok(final_event));
277                                let _ = tx.send(Ok(Event::default().data("[DONE]")));
278                                break;
279                            }
280                        }
281                        Err(e) => {
282                            error!("Stream generation error: {}", e);
283                            let _ = tx.send(Ok(
284                                Event::default().data(format!("{{\"error\": \"{}\"}}", e))
285                            ));
286                            break;
287                        }
288                    }
289                }
290            }
291            Err(e) => {
292                error!("Failed to start streaming: {}", e);
293                let _ = tx.send(Ok(
294                    Event::default().data(format!("{{\"error\": \"{}\"}}", e))
295                ));
296            }
297        }
298    });
299
300    let stream = tokio_stream::wrappers::UnboundedReceiverStream::new(rx);
301    let sse_stream = Sse::new(stream);
302
303    Ok(sse_stream.into_response())
304}
305
306/// Handle non-streaming chat completions
307async fn handle_chat_completions_sync(
308    state: AppState,
309    openai_request: ChatCompletionsRequest,
310    inference_request: InferenceRequest,
311) -> std::result::Result<Response, ServerError> {
312    info!("Processing non-streaming chat completion");
313
314    let prompt_tokens = inference_request.prompt.split_whitespace().count() as u32;
315    match state.engine.infer(inference_request).await {
316        Ok(output) => {
317            let response = ChatCompletionsResponse {
318                id: Uuid::new_v4().to_string(),
319                object: "chat.completion".to_string(),
320                created: chrono::Utc::now().timestamp() as u64,
321                model: openai_request.model,
322                choices: vec![ChatChoice {
323                    index: 0,
324                    message: Some(ChatMessage {
325                        role: MessageRole::Assistant,
326                        content: output.text,
327                        name: None,
328                    }),
329                    delta: None,
330                    finish_reason: Some(finish_reason_to_string(&output.finish_reason)),
331                }],
332                usage: Some(Usage {
333                    prompt_tokens,
334                    completion_tokens: output.tokens.len() as u32,
335                    total_tokens: prompt_tokens + output.tokens.len() as u32,
336                }),
337            };
338
339            Ok(Json(response).into_response())
340        }
341        Err(e) => {
342            error!("Generation failed: {}", e);
343            Err(ServerError::InternalError(e.to_string()))
344        }
345    }
346}
347
348/// Convert OpenAI chat request to internal inference request
349fn convert_chat_request(
350    request: &ChatCompletionsRequest,
351) -> ferrum_types::Result<InferenceRequest> {
352    let prompt = apply_chat_template(&request.messages, &request.model);
353
354    Ok(InferenceRequest {
355        id: RequestId(Uuid::new_v4()),
356        model_id: ModelId(request.model.clone()),
357        prompt,
358        sampling_params: SamplingParams {
359            max_tokens: request.max_tokens.unwrap_or(100) as usize,
360            temperature: request.temperature.unwrap_or(1.0),
361            top_p: request.top_p.unwrap_or(1.0),
362            top_k: None, // OpenAI doesn't use top-k
363            repetition_penalty: 1.0,
364            presence_penalty: request.presence_penalty.unwrap_or(0.0),
365            frequency_penalty: request.frequency_penalty.unwrap_or(0.0),
366            stop_sequences: request.stop.clone().unwrap_or_default(),
367            seed: request.seed,
368            min_p: None,
369            tfs: None,
370            typical_p: None,
371            mirostat: None,
372            response_format: match &request.response_format {
373                Some(rf) if rf.format_type == "json_object" => {
374                    ferrum_types::ResponseFormat::JsonObject
375                }
376                Some(rf) if rf.format_type == "json_schema" => {
377                    // OpenAI wraps the actual schema one level deeper.
378                    match rf.json_schema.as_ref().map(|js| &js.schema) {
379                        Some(schema) => match serde_json::to_string(schema) {
380                            Ok(s) => ferrum_types::ResponseFormat::JsonSchema(s),
381                            Err(_) => ferrum_types::ResponseFormat::Text,
382                        },
383                        None => ferrum_types::ResponseFormat::JsonObject, // fallback
384                    }
385                }
386                _ => ferrum_types::ResponseFormat::Text,
387            },
388        },
389        stream: request.stream.unwrap_or(false),
390        priority: Priority::Normal, // Default priority
391        client_id: None,
392        session_id: None,
393        created_at: chrono::Utc::now(),
394        metadata: std::collections::HashMap::new(),
395    })
396}
397
398/// Other handlers
399async fn completions_handler(
400    State(_state): State<AppState>,
401    Json(_request): Json<CompletionsRequest>,
402) -> std::result::Result<Response, ServerError> {
403    // TODO: Implement legacy completions endpoint
404    Err(ServerError::NotImplemented(
405        "Legacy completions not implemented in MVP".to_string(),
406    ))
407}
408
409/// Embeddings handler — text and image embedding via OpenAI-compatible API.
410async fn embeddings_handler(
411    State(state): State<AppState>,
412    Json(request): Json<EmbeddingsRequest>,
413) -> std::result::Result<Response, ServerError> {
414    let span = span!(Level::INFO, "embeddings", model = %request.model);
415    let _enter = span.enter();
416
417    // Flatten input into individual items
418    let items: Vec<EmbeddingItem> = match request.input {
419        EmbeddingInput::Single(text) => vec![EmbeddingItem {
420            text: Some(text),
421            image: None,
422        }],
423        EmbeddingInput::Batch(texts) => texts
424            .into_iter()
425            .map(|t| EmbeddingItem {
426                text: Some(t),
427                image: None,
428            })
429            .collect(),
430        EmbeddingInput::SingleObject(item) => vec![item],
431        EmbeddingInput::BatchObjects(items) => items,
432    };
433
434    if items.is_empty() {
435        return Err(ServerError::BadRequest("Empty input".to_string()));
436    }
437
438    let mut data = Vec::with_capacity(items.len());
439    let mut total_tokens = 0u32;
440
441    for (idx, item) in items.iter().enumerate() {
442        let embedding = if let Some(ref image) = item.image {
443            state
444                .engine
445                .embed_image(image)
446                .await
447                .map_err(|e| ServerError::InternalError(format!("embed_image: {e}")))?
448        } else if let Some(ref text) = item.text {
449            total_tokens += text.len() as u32;
450            state
451                .engine
452                .embed_text(text)
453                .await
454                .map_err(|e| ServerError::InternalError(format!("embed_text: {e}")))?
455        } else {
456            return Err(ServerError::BadRequest(
457                "Each input must have either 'text' or 'image'".to_string(),
458            ));
459        };
460
461        data.push(EmbeddingData {
462            object: "embedding".to_string(),
463            embedding,
464            index: idx,
465        });
466    }
467
468    let response = EmbeddingsResponse {
469        object: "list".to_string(),
470        data,
471        model: request.model,
472        usage: EmbeddingUsage {
473            prompt_tokens: total_tokens,
474            total_tokens,
475        },
476    };
477
478    Ok(Json(response).into_response())
479}
480
481/// Audio transcription handler (OpenAI-compatible multipart form).
482async fn transcriptions_handler(
483    State(state): State<AppState>,
484    mut multipart: axum::extract::Multipart,
485) -> std::result::Result<Response, ServerError> {
486    let span = span!(Level::INFO, "transcription");
487    let _enter = span.enter();
488
489    let mut file_data: Option<Vec<u8>> = None;
490    let mut language: Option<String> = None;
491
492    while let Some(field) = multipart
493        .next_field()
494        .await
495        .map_err(|e| ServerError::BadRequest(format!("multipart: {e}")))?
496    {
497        let name = field.name().unwrap_or("").to_string();
498        match name.as_str() {
499            "file" => {
500                file_data = Some(
501                    field
502                        .bytes()
503                        .await
504                        .map_err(|e| ServerError::BadRequest(format!("read file: {e}")))?
505                        .to_vec(),
506                );
507            }
508            "language" => {
509                language = field.text().await.ok().filter(|s| !s.is_empty());
510            }
511            _ => {} // ignore model, response_format, etc. for now
512        }
513    }
514
515    let data = file_data.ok_or_else(|| ServerError::BadRequest("missing 'file' field".into()))?;
516
517    let text = state
518        .engine
519        .transcribe_bytes(&data, language.as_deref())
520        .await
521        .map_err(|e| ServerError::InternalError(format!("transcribe: {e}")))?;
522
523    Ok(Json(TranscriptionResponse { text }).into_response())
524}
525
526/// TTS speech synthesis handler (OpenAI-compatible /v1/audio/speech)
527async fn speech_handler(
528    State(state): State<AppState>,
529    Json(request): Json<SpeechRequest>,
530) -> std::result::Result<Response, ServerError> {
531    let span = span!(Level::INFO, "speech");
532    let _guard = span.enter();
533
534    let language = if request.language.is_empty() || request.language == "auto" {
535        None
536    } else {
537        Some(request.language.as_str())
538    };
539
540    let chunk_frames = 10usize;
541    let sample_rate = state.engine.tts_sample_rate();
542
543    if request.stream {
544        // Streaming: chunked transfer encoding with WAV audio
545        let (tx, rx) =
546            mpsc::unbounded_channel::<std::result::Result<axum::body::Bytes, std::io::Error>>();
547
548        let engine = state.engine.clone();
549        let text = request.input.clone();
550        let lang = request.language.clone();
551
552        tokio::task::spawn_blocking(move || {
553            let lang_opt = if lang.is_empty() || lang == "auto" {
554                None
555            } else {
556                Some(lang.as_str())
557            };
558            let rt = tokio::runtime::Handle::current();
559
560            match rt.block_on(engine.synthesize_speech(&text, lang_opt, chunk_frames)) {
561                Ok(chunks) => {
562                    for chunk in &chunks {
563                        let wav_bytes = pcm_to_wav_bytes(chunk, sample_rate);
564                        let _ = tx.send(Ok(axum::body::Bytes::from(wav_bytes)));
565                    }
566                }
567                Err(e) => {
568                    error!("TTS error: {e}");
569                }
570            }
571        });
572
573        let stream = tokio_stream::wrappers::UnboundedReceiverStream::new(rx);
574        let body = axum::body::Body::from_stream(stream);
575        Ok(Response::builder()
576            .status(200)
577            .header("content-type", "audio/wav")
578            .header("transfer-encoding", "chunked")
579            .body(body)
580            .unwrap())
581    } else {
582        // Non-streaming: return complete WAV
583        let chunks = state
584            .engine
585            .synthesize_speech(&request.input, language, chunk_frames)
586            .await
587            .map_err(|e| ServerError::InternalError(format!("TTS: {e}")))?;
588
589        let all_samples: Vec<f32> = chunks.into_iter().flatten().collect();
590        let wav_bytes = pcm_to_wav_bytes(&all_samples, sample_rate);
591
592        Ok(Response::builder()
593            .status(200)
594            .header("content-type", "audio/wav")
595            .header("content-length", wav_bytes.len().to_string())
596            .body(axum::body::Body::from(wav_bytes))
597            .unwrap())
598    }
599}
600
601/// Convert PCM f32 samples to WAV bytes (16-bit, mono).
602fn pcm_to_wav_bytes(samples: &[f32], sample_rate: u32) -> Vec<u8> {
603    let num_samples = samples.len();
604    let data_size = num_samples * 2; // 16-bit = 2 bytes per sample
605    let file_size = 44 + data_size;
606
607    let mut buf = Vec::with_capacity(file_size);
608    // RIFF header
609    buf.extend_from_slice(b"RIFF");
610    buf.extend_from_slice(&((file_size - 8) as u32).to_le_bytes());
611    buf.extend_from_slice(b"WAVE");
612    // fmt chunk
613    buf.extend_from_slice(b"fmt ");
614    buf.extend_from_slice(&16u32.to_le_bytes()); // chunk size
615    buf.extend_from_slice(&1u16.to_le_bytes()); // PCM
616    buf.extend_from_slice(&1u16.to_le_bytes()); // mono
617    buf.extend_from_slice(&sample_rate.to_le_bytes());
618    buf.extend_from_slice(&(sample_rate * 2).to_le_bytes()); // byte rate
619    buf.extend_from_slice(&2u16.to_le_bytes()); // block align
620    buf.extend_from_slice(&16u16.to_le_bytes()); // bits per sample
621                                                 // data chunk
622    buf.extend_from_slice(b"data");
623    buf.extend_from_slice(&(data_size as u32).to_le_bytes());
624    for &s in samples {
625        let i16_val = (s.clamp(-1.0, 1.0) * 32767.0) as i16;
626        buf.extend_from_slice(&i16_val.to_le_bytes());
627    }
628    buf
629}
630
631async fn models_handler(
632    State(state): State<AppState>,
633) -> std::result::Result<Response, ServerError> {
634    let status = state.engine.status().await;
635    let now = chrono::Utc::now().timestamp() as u64;
636    let data = status
637        .loaded_models
638        .into_iter()
639        .map(|model_id| crate::openai::ModelInfo {
640            id: model_id.to_string(),
641            object: "model".to_string(),
642            created: now,
643            owned_by: "ferrum".to_string(),
644            permission: vec![],
645            root: None,
646            parent: None,
647        })
648        .collect();
649
650    let models = ModelListResponse {
651        object: "list".to_string(),
652        data,
653    };
654
655    Ok(Json(models).into_response())
656}
657
658async fn health_handler(
659    State(state): State<AppState>,
660) -> std::result::Result<Response, ServerError> {
661    let engine_status = state.engine.status().await;
662    let scheduler_metrics = state.engine.metrics();
663
664    let health = serde_json::json!({
665        "status": "healthy",
666        "timestamp": chrono::Utc::now().to_rfc3339(),
667        "version": env!("CARGO_PKG_VERSION"),
668        "engine": {
669            "active_requests": engine_status.active_requests,
670            "queued_requests": engine_status.queued_requests,
671        },
672        "scheduler": {
673            "total_requests": scheduler_metrics.total_requests,
674            "successful_requests": scheduler_metrics.successful_requests,
675            "failed_requests": scheduler_metrics.failed_requests,
676            "throughput_rps": scheduler_metrics.throughput_rps,
677        }
678    });
679
680    Ok(Json(health).into_response())
681}
682
683/// Prometheus metrics endpoint — returns metrics in Prometheus text format.
684async fn metrics_handler() -> std::result::Result<Response, ServerError> {
685    let body = match PROM_HANDLE.get() {
686        Some(handle) => handle.render(),
687        None => "# Prometheus recorder not initialized\n".to_string(),
688    };
689
690    Ok((
691        [(
692            axum::http::header::CONTENT_TYPE,
693            "text/plain; version=0.0.4; charset=utf-8",
694        )],
695        body,
696    )
697        .into_response())
698}
699
700async fn root_handler() -> std::result::Result<Response, ServerError> {
701    let info = serde_json::json!({
702        "name": "Ferrum Inference Server",
703        "version": env!("CARGO_PKG_VERSION"),
704        "api_version": "v1",
705        "status": "running"
706    });
707
708    Ok(Json(info).into_response())
709}
710
711/// Server error type for HTTP responses
712#[derive(Debug)]
713enum ServerError {
714    BadRequest(String),
715    InternalError(String),
716    NotImplemented(String),
717}
718
719impl IntoResponse for ServerError {
720    fn into_response(self) -> Response {
721        let (status, message) = match self {
722            ServerError::BadRequest(msg) => (AxumStatusCode::BAD_REQUEST, msg),
723            ServerError::InternalError(msg) => (AxumStatusCode::INTERNAL_SERVER_ERROR, msg),
724            ServerError::NotImplemented(msg) => (AxumStatusCode::NOT_IMPLEMENTED, msg),
725        };
726
727        let error = OpenAiError {
728            error: OpenAiErrorDetail {
729                message,
730                error_type: "server_error".to_string(),
731                param: None,
732                code: None,
733            },
734        };
735
736        (status, Json(error)).into_response()
737    }
738}
739
740impl std::fmt::Display for MessageRole {
741    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
742        match self {
743            MessageRole::System => write!(f, "system"),
744            MessageRole::User => write!(f, "user"),
745            MessageRole::Assistant => write!(f, "assistant"),
746            MessageRole::Function => write!(f, "function"),
747        }
748    }
749}
750
751/// Render OpenAI-style chat messages into the prompt string the model was
752/// trained on. Mirrors the templates `ferrum-cli::commands::run` uses so
753/// `/v1/chat/completions` produces the same behaviour as the interactive CLI.
754///
755/// Detects model family from the request's `model` field:
756///   - qwen (Qwen2 / Qwen2.5 / Qwen3): ChatML with `<|im_start|>` / `<|im_end|>`
757///     (Qwen3 adds the empty `<think></think>` marker to disable reasoning)
758///   - llama 3: `<|start_header_id|>…<|end_header_id|>` + `<|eot_id|>`
759///   - fallback: TinyLlama-style `<|system|>` / `<|user|>` / `<|assistant|>`
760///     with `</s>` separators
761///
762/// All templates end with the assistant header so the first generated token
763/// becomes the reply content (no extra role prefix).
764fn apply_chat_template(messages: &[ChatMessage], model_id: &str) -> String {
765    let model_lower = model_id.to_lowercase();
766
767    if model_lower.contains("qwen") {
768        let mut prompt = String::new();
769        for msg in messages {
770            let role = match msg.role {
771                MessageRole::System => "system",
772                MessageRole::User => "user",
773                MessageRole::Assistant => "assistant",
774                MessageRole::Function => "function",
775            };
776            prompt.push_str(&format!(
777                "<|im_start|>{}\n{}<|im_end|>\n",
778                role, msg.content
779            ));
780        }
781        prompt.push_str("<|im_start|>assistant\n");
782        if model_lower.contains("qwen3") {
783            // Qwen3: disable thinking mode by inserting an empty think block.
784            prompt.push_str("<think>\n\n</think>\n\n");
785        }
786        prompt
787    } else if model_lower.contains("llama") && model_lower.contains("3") {
788        let mut prompt = String::from("<|begin_of_text|>");
789        for msg in messages {
790            let role = match msg.role {
791                MessageRole::System => "system",
792                MessageRole::User => "user",
793                MessageRole::Assistant => "assistant",
794                MessageRole::Function => "function",
795            };
796            prompt.push_str(&format!(
797                "<|start_header_id|>{}<|end_header_id|>\n\n{}<|eot_id|>",
798                role, msg.content
799            ));
800        }
801        prompt.push_str("<|start_header_id|>assistant<|end_header_id|>\n\n");
802        prompt
803    } else {
804        // TinyLlama / generic chat format. Promote the first system message
805        // to the top; subsequent ones (rare) are emitted inline.
806        let has_system = messages
807            .iter()
808            .any(|m| matches!(m.role, MessageRole::System));
809        let mut prompt = String::new();
810        if !has_system {
811            prompt.push_str("<|system|>\nYou are a helpful assistant.</s>\n");
812        }
813        for msg in messages {
814            let tag = match msg.role {
815                MessageRole::System => "system",
816                MessageRole::User => "user",
817                MessageRole::Assistant => "assistant",
818                MessageRole::Function => "assistant",
819            };
820            prompt.push_str(&format!("<|{}|>\n{}</s>\n", tag, msg.content));
821        }
822        prompt.push_str("<|assistant|>\n");
823        prompt
824    }
825}
826
827/// Convert FinishReason to OpenAI API string
828fn finish_reason_to_string(reason: &FinishReason) -> String {
829    match reason {
830        FinishReason::Length => "length".to_string(),
831        FinishReason::Stop => "stop".to_string(),
832        FinishReason::EOS => "stop".to_string(),
833        FinishReason::Cancelled => "cancelled".to_string(),
834        FinishReason::Error => "error".to_string(),
835        FinishReason::ContentFilter => "content_filter".to_string(),
836    }
837}
838
839#[cfg(test)]
840mod chat_template_tests {
841    use super::*;
842
843    fn msg(role: MessageRole, content: &str) -> ChatMessage {
844        ChatMessage {
845            role,
846            content: content.to_string(),
847            name: None,
848        }
849    }
850
851    #[test]
852    fn qwen3_renders_chatml_with_think_marker() {
853        let out = apply_chat_template(
854            &[
855                msg(MessageRole::System, "You are helpful."),
856                msg(MessageRole::User, "Hi"),
857            ],
858            "qwen3:0.6b",
859        );
860        assert!(out.contains("<|im_start|>system\nYou are helpful.<|im_end|>"));
861        assert!(out.contains("<|im_start|>user\nHi<|im_end|>"));
862        assert!(out.ends_with("<|im_start|>assistant\n<think>\n\n</think>\n\n"));
863    }
864
865    #[test]
866    fn qwen2_renders_chatml_without_think() {
867        let out = apply_chat_template(&[msg(MessageRole::User, "Hi")], "Qwen/Qwen2.5-7B-Instruct");
868        assert!(out.ends_with("<|im_start|>assistant\n"));
869        assert!(!out.contains("<think>"));
870    }
871
872    #[test]
873    fn multi_turn_preserves_order() {
874        let out = apply_chat_template(
875            &[
876                msg(MessageRole::User, "A"),
877                msg(MessageRole::Assistant, "B"),
878                msg(MessageRole::User, "C"),
879            ],
880            "qwen3",
881        );
882        let a_idx = out.find("A").unwrap();
883        let b_idx = out.find("B").unwrap();
884        let c_idx = out.find("C").unwrap();
885        assert!(a_idx < b_idx && b_idx < c_idx);
886    }
887
888    #[test]
889    fn llama3_renders_header_format() {
890        let out = apply_chat_template(
891            &[
892                msg(MessageRole::System, "sys"),
893                msg(MessageRole::User, "hi"),
894            ],
895            "meta-llama/Llama-3.2-1B-Instruct",
896        );
897        assert!(out.starts_with("<|begin_of_text|>"));
898        assert!(out.contains("<|start_header_id|>system<|end_header_id|>\n\nsys<|eot_id|>"));
899        assert!(out.contains("<|start_header_id|>user<|end_header_id|>\n\nhi<|eot_id|>"));
900        assert!(out.ends_with("<|start_header_id|>assistant<|end_header_id|>\n\n"));
901    }
902
903    #[test]
904    fn unknown_model_uses_tinyllama_fallback() {
905        let out = apply_chat_template(&[msg(MessageRole::User, "hi")], "mystery-model");
906        assert!(out.contains("<|system|>"));
907        assert!(out.contains("<|user|>\nhi</s>"));
908        assert!(out.ends_with("<|assistant|>\n"));
909    }
910}