//! Inference endpoints with streaming support use std::sync::Arc; use axum::extract::State; use rtx_inference::{InferenceEngine, ServingTokenizer}; use tokio::sync::RwLock; use crate::{ApiError, ApiResult}; /// Shared application state carrying the inference engine. /// /// `engine` is `None` until a model has been loaded (e.g. by `main.rs` /// or an admin endpoint); handlers respond 503 rather than fabricating /// output when no engine is available. /// /// `tokenizer` defaults to byte-level (each UTF-8 byte is one token id); /// load a real vocabulary tokenizer with `ServingTokenizer::from_file` and /// attach it via `ServingServer::with_engine_and_tokenizer` to use it /// instead. #[derive(Clone)] pub struct AppState { pub engine: Option>>, pub tokenizer: ServingTokenizer, } impl Default for AppState { fn default() -> Self { Self { engine: None, tokenizer: ServingTokenizer::default(), } } } /// Inference request #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct InferenceRequest { /// Model ID to use for inference pub model: String, /// Input prompt pub prompt: String, /// Maximum tokens to generate pub max_tokens: Option, /// Temperature for sampling pub temperature: Option, /// Whether to stream the response pub stream: Option, } /// Inference response #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct InferenceResponse { /// Generated text pub text: String, /// Finish reason pub finish_reason: String, /// Usage statistics pub usage: TokenUsage, } /// Token usage statistics #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct TokenUsage { /// Prompt tokens pub prompt_tokens: u32, /// Completion tokens pub completion_tokens: u32, /// Total tokens pub total_tokens: u32, } /// Handle inference request by dispatching to the rtx-inference engine. /// /// Tokenization is pluggable via `state.tokenizer` (`ServingTokenizer`), /// defaulting to byte-level (each UTF-8 byte is a token id) for backward /// compatibility. Attach a real vocabulary tokenizer with /// `ServingTokenizer::from_file` for vocabulary-aware tokenization. Returns /// 503 when no engine/model is loaded. pub async fn inference( State(state): State, request: axum::Json, ) -> ApiResult> { let req = request.0; let engine = state.engine.as_ref().ok_or_else(|| { ApiError::service_unavailable("no inference engine configured — load a model first") })?; let input_tokens: Vec = state.tokenizer.encode(&req.prompt); let prompt_tokens = input_tokens.len() as u32; let max_new_tokens = req.max_tokens.unwrap_or(128) as usize; let mut infer_req = rtx_inference::InferenceRequest::new(req.model.clone(), input_tokens, max_new_tokens); if let Some(t) = req.temperature { infer_req.temperature = t; } let result = { let engine = engine.read().await; engine .infer(infer_req) .await .map_err(|e| ApiError::internal(format!("inference failed: {e}")))? }; let completion_tokens = result.output_tokens.len() as u32; let response_text = state.tokenizer.decode(&result.output_tokens); let response = InferenceResponse { text: response_text, finish_reason: format!("{:?}", result.finish_reason).to_lowercase(), usage: TokenUsage { prompt_tokens, completion_tokens, total_tokens: prompt_tokens + completion_tokens, }, }; Ok(axum::Json(response)) } #[cfg(test)] mod tests { use super::*; use axum::{Router, http::StatusCode, routing::post}; use axum_test::TestServer; /// Without an engine configured, the endpoint must refuse (503), /// never fabricate a response. #[tokio::test] async fn test_inference_endpoint_no_engine_returns_503() { let app = Router::new() .route("/v1/completions", post(inference)) .with_state(AppState::default()); let server = TestServer::new(app).unwrap(); let request = InferenceRequest { model: "test-model".to_string(), prompt: "Hello, world!".to_string(), max_tokens: Some(100), temperature: Some(0.7), stream: Some(false), }; let response = server.post("/v1/completions").json(&request).await; response.assert_status(StatusCode::SERVICE_UNAVAILABLE); } #[test] fn test_inference_request_serialization() { let request = InferenceRequest { model: "test-model".to_string(), prompt: "Hello, world!".to_string(), max_tokens: Some(100), temperature: Some(0.7), stream: Some(false), }; let json = serde_json::to_string(&request).unwrap(); assert!(json.contains("test-model")); assert!(json.contains("Hello, world!")); } #[test] fn test_app_state_default_uses_byte_level_tokenizer() { let state = AppState::default(); assert!(matches!(state.tokenizer, ServingTokenizer::ByteLevel)); let ids = state.tokenizer.encode("hi"); assert_eq!(state.tokenizer.decode(&ids), "hi"); } #[test] fn test_inference_response_serialization() { let response = InferenceResponse { text: "Hello back!".to_string(), finish_reason: "stop".to_string(), usage: TokenUsage { prompt_tokens: 3, completion_tokens: 2, total_tokens: 5, }, }; let json = serde_json::to_string(&response).unwrap(); assert!(json.contains("Hello back!")); assert!(json.contains("stop")); } }