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: 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: 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    // Combine all messages into a single prompt for MVP
353    let prompt = request
354        .messages
355        .iter()
356        .map(|msg| format!("{}: {}", msg.role.to_string(), msg.content))
357        .collect::<Vec<_>>()
358        .join("\n");
359
360    Ok(InferenceRequest {
361        id: RequestId(Uuid::new_v4()),
362        model_id: ModelId(request.model.clone()),
363        prompt,
364        sampling_params: SamplingParams {
365            max_tokens: request.max_tokens.unwrap_or(100) as usize,
366            temperature: request.temperature.unwrap_or(1.0),
367            top_p: request.top_p.unwrap_or(1.0),
368            top_k: None, // OpenAI doesn't use top-k
369            repetition_penalty: 1.0,
370            presence_penalty: request.presence_penalty.unwrap_or(0.0),
371            frequency_penalty: request.frequency_penalty.unwrap_or(0.0),
372            stop_sequences: request.stop.clone().unwrap_or_default(),
373            seed: request.seed,
374            min_p: None,
375            tfs: None,
376            typical_p: None,
377            mirostat: None,
378            response_format: match &request.response_format {
379                Some(rf) if rf.format_type == "json_object" => {
380                    ferrum_types::ResponseFormat::JsonObject
381                }
382                _ => ferrum_types::ResponseFormat::Text,
383            },
384        },
385        stream: request.stream.unwrap_or(false),
386        priority: Priority::Normal, // Default priority
387        client_id: None,
388        session_id: None,
389        created_at: chrono::Utc::now(),
390        metadata: std::collections::HashMap::new(),
391    })
392}
393
394/// Other handlers
395async fn completions_handler(
396    State(_state): State<AppState>,
397    Json(_request): Json<CompletionsRequest>,
398) -> std::result::Result<Response, ServerError> {
399    // TODO: Implement legacy completions endpoint
400    Err(ServerError::NotImplemented(
401        "Legacy completions not implemented in MVP".to_string(),
402    ))
403}
404
405/// Embeddings handler — text and image embedding via OpenAI-compatible API.
406async fn embeddings_handler(
407    State(state): State<AppState>,
408    Json(request): Json<EmbeddingsRequest>,
409) -> std::result::Result<Response, ServerError> {
410    let span = span!(Level::INFO, "embeddings", model = %request.model);
411    let _enter = span.enter();
412
413    // Flatten input into individual items
414    let items: Vec<EmbeddingItem> = match request.input {
415        EmbeddingInput::Single(text) => vec![EmbeddingItem {
416            text: Some(text),
417            image: None,
418        }],
419        EmbeddingInput::Batch(texts) => texts
420            .into_iter()
421            .map(|t| EmbeddingItem {
422                text: Some(t),
423                image: None,
424            })
425            .collect(),
426        EmbeddingInput::SingleObject(item) => vec![item],
427        EmbeddingInput::BatchObjects(items) => items,
428    };
429
430    if items.is_empty() {
431        return Err(ServerError::BadRequest("Empty input".to_string()));
432    }
433
434    let mut data = Vec::with_capacity(items.len());
435    let mut total_tokens = 0u32;
436
437    for (idx, item) in items.iter().enumerate() {
438        let embedding = if let Some(ref image) = item.image {
439            state
440                .engine
441                .embed_image(image)
442                .await
443                .map_err(|e| ServerError::InternalError(format!("embed_image: {e}")))?
444        } else if let Some(ref text) = item.text {
445            total_tokens += text.len() as u32;
446            state
447                .engine
448                .embed_text(text)
449                .await
450                .map_err(|e| ServerError::InternalError(format!("embed_text: {e}")))?
451        } else {
452            return Err(ServerError::BadRequest(
453                "Each input must have either 'text' or 'image'".to_string(),
454            ));
455        };
456
457        data.push(EmbeddingData {
458            object: "embedding".to_string(),
459            embedding,
460            index: idx,
461        });
462    }
463
464    let response = EmbeddingsResponse {
465        object: "list".to_string(),
466        data,
467        model: request.model,
468        usage: EmbeddingUsage {
469            prompt_tokens: total_tokens,
470            total_tokens,
471        },
472    };
473
474    Ok(Json(response).into_response())
475}
476
477/// Audio transcription handler (OpenAI-compatible multipart form).
478async fn transcriptions_handler(
479    State(state): State<AppState>,
480    mut multipart: axum::extract::Multipart,
481) -> std::result::Result<Response, ServerError> {
482    let span = span!(Level::INFO, "transcription");
483    let _enter = span.enter();
484
485    let mut file_data: Option<Vec<u8>> = None;
486    let mut language: Option<String> = None;
487
488    while let Some(field) = multipart
489        .next_field()
490        .await
491        .map_err(|e| ServerError::BadRequest(format!("multipart: {e}")))?
492    {
493        let name = field.name().unwrap_or("").to_string();
494        match name.as_str() {
495            "file" => {
496                file_data = Some(
497                    field
498                        .bytes()
499                        .await
500                        .map_err(|e| ServerError::BadRequest(format!("read file: {e}")))?
501                        .to_vec(),
502                );
503            }
504            "language" => {
505                language = field.text().await.ok().filter(|s| !s.is_empty());
506            }
507            _ => {} // ignore model, response_format, etc. for now
508        }
509    }
510
511    let data = file_data.ok_or_else(|| ServerError::BadRequest("missing 'file' field".into()))?;
512
513    let text = state
514        .engine
515        .transcribe_bytes(&data, language.as_deref())
516        .await
517        .map_err(|e| ServerError::InternalError(format!("transcribe: {e}")))?;
518
519    Ok(Json(TranscriptionResponse { text }).into_response())
520}
521
522/// TTS speech synthesis handler (OpenAI-compatible /v1/audio/speech)
523async fn speech_handler(
524    State(state): State<AppState>,
525    Json(request): Json<SpeechRequest>,
526) -> std::result::Result<Response, ServerError> {
527    let span = span!(Level::INFO, "speech");
528    let _guard = span.enter();
529
530    let language = if request.language.is_empty() || request.language == "auto" {
531        None
532    } else {
533        Some(request.language.as_str())
534    };
535
536    let chunk_frames = 10usize;
537    let sample_rate = state.engine.tts_sample_rate();
538
539    if request.stream {
540        // Streaming: chunked transfer encoding with WAV audio
541        let (tx, rx) =
542            mpsc::unbounded_channel::<std::result::Result<axum::body::Bytes, std::io::Error>>();
543
544        let engine = state.engine.clone();
545        let text = request.input.clone();
546        let lang = request.language.clone();
547
548        tokio::task::spawn_blocking(move || {
549            let lang_opt = if lang.is_empty() || lang == "auto" {
550                None
551            } else {
552                Some(lang.as_str())
553            };
554            let rt = tokio::runtime::Handle::current();
555
556            match rt.block_on(engine.synthesize_speech(&text, lang_opt, chunk_frames)) {
557                Ok(chunks) => {
558                    for chunk in &chunks {
559                        let wav_bytes = pcm_to_wav_bytes(chunk, sample_rate);
560                        let _ = tx.send(Ok(axum::body::Bytes::from(wav_bytes)));
561                    }
562                }
563                Err(e) => {
564                    error!("TTS error: {e}");
565                }
566            }
567        });
568
569        let stream = tokio_stream::wrappers::UnboundedReceiverStream::new(rx);
570        let body = axum::body::Body::from_stream(stream);
571        Ok(Response::builder()
572            .status(200)
573            .header("content-type", "audio/wav")
574            .header("transfer-encoding", "chunked")
575            .body(body)
576            .unwrap())
577    } else {
578        // Non-streaming: return complete WAV
579        let chunks = state
580            .engine
581            .synthesize_speech(&request.input, language, chunk_frames)
582            .await
583            .map_err(|e| ServerError::InternalError(format!("TTS: {e}")))?;
584
585        let all_samples: Vec<f32> = chunks.into_iter().flatten().collect();
586        let wav_bytes = pcm_to_wav_bytes(&all_samples, sample_rate);
587
588        Ok(Response::builder()
589            .status(200)
590            .header("content-type", "audio/wav")
591            .header("content-length", wav_bytes.len().to_string())
592            .body(axum::body::Body::from(wav_bytes))
593            .unwrap())
594    }
595}
596
597/// Convert PCM f32 samples to WAV bytes (16-bit, mono).
598fn pcm_to_wav_bytes(samples: &[f32], sample_rate: u32) -> Vec<u8> {
599    let num_samples = samples.len();
600    let data_size = num_samples * 2; // 16-bit = 2 bytes per sample
601    let file_size = 44 + data_size;
602
603    let mut buf = Vec::with_capacity(file_size);
604    // RIFF header
605    buf.extend_from_slice(b"RIFF");
606    buf.extend_from_slice(&((file_size - 8) as u32).to_le_bytes());
607    buf.extend_from_slice(b"WAVE");
608    // fmt chunk
609    buf.extend_from_slice(b"fmt ");
610    buf.extend_from_slice(&16u32.to_le_bytes()); // chunk size
611    buf.extend_from_slice(&1u16.to_le_bytes()); // PCM
612    buf.extend_from_slice(&1u16.to_le_bytes()); // mono
613    buf.extend_from_slice(&sample_rate.to_le_bytes());
614    buf.extend_from_slice(&(sample_rate * 2).to_le_bytes()); // byte rate
615    buf.extend_from_slice(&2u16.to_le_bytes()); // block align
616    buf.extend_from_slice(&16u16.to_le_bytes()); // bits per sample
617                                                 // data chunk
618    buf.extend_from_slice(b"data");
619    buf.extend_from_slice(&(data_size as u32).to_le_bytes());
620    for &s in samples {
621        let i16_val = (s.clamp(-1.0, 1.0) * 32767.0) as i16;
622        buf.extend_from_slice(&i16_val.to_le_bytes());
623    }
624    buf
625}
626
627async fn models_handler(
628    State(state): State<AppState>,
629) -> std::result::Result<Response, ServerError> {
630    let status = state.engine.status().await;
631    let now = chrono::Utc::now().timestamp() as u64;
632    let data = status
633        .loaded_models
634        .into_iter()
635        .map(|model_id| crate::openai::ModelInfo {
636            id: model_id.to_string(),
637            object: "model".to_string(),
638            created: now,
639            owned_by: "ferrum".to_string(),
640            permission: vec![],
641            root: None,
642            parent: None,
643        })
644        .collect();
645
646    let models = ModelListResponse {
647        object: "list".to_string(),
648        data,
649    };
650
651    Ok(Json(models).into_response())
652}
653
654async fn health_handler(
655    State(state): State<AppState>,
656) -> std::result::Result<Response, ServerError> {
657    let engine_status = state.engine.status().await;
658    let scheduler_metrics = state.engine.metrics();
659
660    let health = serde_json::json!({
661        "status": "healthy",
662        "timestamp": chrono::Utc::now().to_rfc3339(),
663        "version": env!("CARGO_PKG_VERSION"),
664        "engine": {
665            "active_requests": engine_status.active_requests,
666            "queued_requests": engine_status.queued_requests,
667        },
668        "scheduler": {
669            "total_requests": scheduler_metrics.total_requests,
670            "successful_requests": scheduler_metrics.successful_requests,
671            "failed_requests": scheduler_metrics.failed_requests,
672            "throughput_rps": scheduler_metrics.throughput_rps,
673        }
674    });
675
676    Ok(Json(health).into_response())
677}
678
679/// Prometheus metrics endpoint — returns metrics in Prometheus text format.
680async fn metrics_handler() -> std::result::Result<Response, ServerError> {
681    let body = match PROM_HANDLE.get() {
682        Some(handle) => handle.render(),
683        None => "# Prometheus recorder not initialized\n".to_string(),
684    };
685
686    Ok((
687        [(
688            axum::http::header::CONTENT_TYPE,
689            "text/plain; version=0.0.4; charset=utf-8",
690        )],
691        body,
692    )
693        .into_response())
694}
695
696async fn root_handler() -> std::result::Result<Response, ServerError> {
697    let info = serde_json::json!({
698        "name": "Ferrum Inference Server",
699        "version": env!("CARGO_PKG_VERSION"),
700        "api_version": "v1",
701        "status": "running"
702    });
703
704    Ok(Json(info).into_response())
705}
706
707/// Server error type for HTTP responses
708#[derive(Debug)]
709enum ServerError {
710    BadRequest(String),
711    InternalError(String),
712    NotImplemented(String),
713}
714
715impl IntoResponse for ServerError {
716    fn into_response(self) -> Response {
717        let (status, message) = match self {
718            ServerError::BadRequest(msg) => (AxumStatusCode::BAD_REQUEST, msg),
719            ServerError::InternalError(msg) => (AxumStatusCode::INTERNAL_SERVER_ERROR, msg),
720            ServerError::NotImplemented(msg) => (AxumStatusCode::NOT_IMPLEMENTED, msg),
721        };
722
723        let error = OpenAiError {
724            error: OpenAiErrorDetail {
725                message,
726                error_type: "server_error".to_string(),
727                param: None,
728                code: None,
729            },
730        };
731
732        (status, Json(error)).into_response()
733    }
734}
735
736impl std::fmt::Display for MessageRole {
737    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
738        match self {
739            MessageRole::System => write!(f, "system"),
740            MessageRole::User => write!(f, "user"),
741            MessageRole::Assistant => write!(f, "assistant"),
742            MessageRole::Function => write!(f, "function"),
743        }
744    }
745}
746
747/// Convert FinishReason to OpenAI API string
748fn finish_reason_to_string(reason: &FinishReason) -> String {
749    match reason {
750        FinishReason::Length => "length".to_string(),
751        FinishReason::Stop => "stop".to_string(),
752        FinishReason::EOS => "stop".to_string(),
753        FinishReason::Cancelled => "cancelled".to_string(),
754        FinishReason::Error => "error".to_string(),
755        FinishReason::ContentFilter => "content_filter".to_string(),
756    }
757}