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/audio/transcriptions", post(transcriptions_handler))
74 .route("/v1/audio/speech", post(speech_handler))
75 .route("/v1/models", get(models_handler))
76 .route("/health", get(health_handler))
78 .route("/metrics", get(metrics_handler))
79 .route("/", get(root_handler))
80 .layer(
82 ServiceBuilder::new()
83 .layer(TraceLayer::new_for_http())
84 .layer(CorsLayer::permissive()), )
86 .with_state(app_state)
87 }
88}
89
90#[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 Ok(())
120 }
121
122 fn is_running(&self) -> bool {
123 true
125 }
126
127 fn address(&self) -> Option<std::net::SocketAddr> {
128 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 unimplemented!("Dynamic handler registration not implemented in MVP")
142 }
143
144 fn register_middleware(&mut self, _middleware: Box<dyn crate::traits::Middleware>) {
145 unimplemented!("Dynamic middleware registration not implemented in MVP")
147 }
148
149 fn get_metrics(&self) -> ServerMetrics {
150 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
171async 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 let inference_request =
187 convert_chat_request(&request).map_err(|e| ServerError::BadRequest(e.to_string()))?;
188
189 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
197async 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 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 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 if token_count >= max_tokens || chunk.finish_reason.is_some() {
251 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
306async 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
348fn 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, 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 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, }
385 }
386 _ => ferrum_types::ResponseFormat::Text,
387 },
388 },
389 stream: request.stream.unwrap_or(false),
390 priority: Priority::Normal, client_id: None,
392 session_id: None,
393 created_at: chrono::Utc::now(),
394 metadata: std::collections::HashMap::new(),
395 })
396}
397
398async fn completions_handler(
400 State(_state): State<AppState>,
401 Json(_request): Json<CompletionsRequest>,
402) -> std::result::Result<Response, ServerError> {
403 Err(ServerError::NotImplemented(
405 "Legacy completions not implemented in MVP".to_string(),
406 ))
407}
408
409async 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 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
481async 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 _ => {} }
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
526async 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 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 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
601fn pcm_to_wav_bytes(samples: &[f32], sample_rate: u32) -> Vec<u8> {
603 let num_samples = samples.len();
604 let data_size = num_samples * 2; let file_size = 44 + data_size;
606
607 let mut buf = Vec::with_capacity(file_size);
608 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 buf.extend_from_slice(b"fmt ");
614 buf.extend_from_slice(&16u32.to_le_bytes()); buf.extend_from_slice(&1u16.to_le_bytes()); buf.extend_from_slice(&1u16.to_le_bytes()); buf.extend_from_slice(&sample_rate.to_le_bytes());
618 buf.extend_from_slice(&(sample_rate * 2).to_le_bytes()); buf.extend_from_slice(&2u16.to_le_bytes()); buf.extend_from_slice(&16u16.to_le_bytes()); 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
683async 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#[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
751fn 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 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 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
827fn 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}