feat: strong typing for complete endpoints

This commit is contained in:
2026-04-10 18:27:02 +02:00
parent 562d154480
commit d5856557b4
9 changed files with 609 additions and 158 deletions
+79 -15
View File
@@ -3,21 +3,6 @@ use utoipa::ToSchema;
use crate::errors::OllamaError;
#[derive(Serialize, Deserialize, ToSchema)]
pub struct ChatRequest {
pub model: String,
pub prompt: Option<String>,
pub messages: Option<Vec<Message>>,
#[serde(default)]
pub stream: bool,
pub temperature: Option<f32>,
pub top_p: Option<f32>,
pub max_tokens: Option<u32>,
pub stop: Option<Vec<String>>,
pub system: Option<String>,
pub keep_alive: Option<String>,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
#[serde(rename_all = "lowercase")]
pub enum Role {
@@ -37,6 +22,8 @@ pub struct Message {
pub content: String,
}
// ---------------------------
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ModelsResponse {
pub models: Vec<ModelInfo>,
@@ -91,3 +78,80 @@ pub struct UnloadModelResponse {
pub model: String,
pub status: String,
}
#[derive(Debug, Deserialize, Serialize, ToSchema)]
pub struct BaseLLMRequest {
pub model: String,
#[serde(default)]
pub stream: bool,
pub temperature: Option<f32>,
pub top_p: Option<f32>,
// Ollama-native
pub top_k: Option<u32>,
pub repeat_penalty: Option<f32>,
pub seed: Option<i64>,
pub num_ctx: Option<u32>,
pub num_predict: Option<u32>,
pub stop: Option<Vec<String>>,
pub keep_alive: Option<String>,
}
#[derive(Debug, Deserialize, Serialize, ToSchema)]
pub struct CompletionRequest {
#[serde(flatten)]
pub base: BaseLLMRequest,
pub prompt: String,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
#[serde(rename_all = "snake_case")]
pub enum CompletionObject {
TextCompletion,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
#[serde(rename_all = "snake_case")]
pub enum FinishReason {
Stop,
Length,
ContentFilter,
ToolCalls,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct CompletionResponse {
pub id: String,
pub object: CompletionObject,
pub created: u64,
pub model: String,
pub choices: Vec<Choice>,
pub usage: Usage,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct Choice {
pub text: String,
pub index: u32,
pub finish_reason: FinishReason,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct Usage {
pub prompt_tokens: u32,
pub completion_tokens: u32,
pub total_tokens: u32,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct CompletionChunk {
pub id: String,
pub object: String,
pub choices: Vec<Choice>,
}
+36 -2
View File
@@ -1,8 +1,6 @@
use serde::{Deserialize, Serialize};
use utoipa::ToSchema;
use crate::errors::OllamaError;
#[derive(Debug, Serialize, Deserialize)]
pub struct OllamaModels {
pub models: Vec<OllamaModel>,
@@ -25,3 +23,39 @@ pub struct OllamaModelDetails {
pub parameter_size: Option<String>,
pub quantization_level: Option<String>,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct OllamaOptions {
pub temperature: Option<f32>,
pub top_p: Option<f32>,
pub top_k: Option<u32>,
pub repeat_penalty: Option<f32>,
pub seed: Option<i64>,
pub num_ctx: Option<u32>,
pub num_predict: Option<u32>,
}
#[derive(Debug, Serialize)]
pub struct OllamaGenerateRequest<'a> {
pub model: &'a str,
pub prompt: &'a str,
pub stream: bool,
pub options: OllamaOptions,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct OllamaGenerateResponse {
pub model: String,
pub created_at: Option<String>,
pub response: String,
pub done: bool,
#[serde(default)]
pub context: Option<Vec<u64>>,
pub total_duration: Option<u64>,
pub load_duration: Option<u64>,
pub prompt_eval_count: Option<u32>,
pub eval_count: Option<u32>,
}
+122 -110
View File
@@ -1,5 +1,6 @@
use crate::dto::{api, ollama};
use crate::errors::OllamaError;
use axum::Json;
use axum::response::sse::Event;
use futures::StreamExt;
use reqwest::Client;
@@ -36,48 +37,16 @@ impl OllamaProvider {
Ok(res.models.iter().any(|m| m.name == model))
}
// fn build_options(body: &Value) -> Value {
// json!({
// "temperature": body.get("temperature"),
// "top_p": body.get("top_p"),
// "num_predict": body.get("max_tokens"),
// })
// }
fn extract_completion_params<'a>(
&self,
body: &'a api::CompletionRequest,
) -> Result<(&'a str, &'a str), OllamaError> {
let prompt = body.prompt.trim();
// async fn validate_model(&self, model: &str) -> Result<(), OllamaError> {
// let available = self.list_models().await?;
// let exists = available
// .get("models")
// .and_then(|m| m.as_array())
// .map(|arr| {
// arr.iter()
// .any(|m| m.get("name").and_then(|n| n.as_str()) == Some(model))
// })
// .unwrap_or(false);
let model = body.base.model.as_str();
// if !exists {
// return Err(OllamaError::ModelNotFound(model.to_string()));
// }
// Ok(())
// }
// fn extract_completion_params<'a>(
// &self,
// body: &'a Value,
// ) -> Result<(&'a str, &'a str), OllamaError> {
// let prompt = body
// .get("prompt")
// .and_then(|v| v.as_str())
// .filter(|s| !s.trim().is_empty())
// .ok_or(OllamaError::MissingPrompt)?;
// let model = body
// .get("model")
// .and_then(|v| v.as_str())
// .ok_or(OllamaError::MissingModel)?;
// Ok((prompt, model))
// }
Ok((prompt, model))
}
// fn format_completion_response(&self, res: &Value) -> Value {
// json!({
@@ -209,7 +178,7 @@ impl OllamaProvider {
pub async fn unload_model(&self, model: &str) -> Result<api::UnloadModelResponse, OllamaError> {
let url = format!("{}/api/generate", self.base_url);
let exists = self.model_exists(model).await?;
if !exists {
return Err(OllamaError::ModelNotFound(model.to_string()));
@@ -237,88 +206,131 @@ impl OllamaProvider {
})
}
// pub async fn completions(&self, body: Value) -> Result<Value, OllamaError> {
// let (prompt, model) = self.extract_completion_params(&body)?;
// self.validate_model(model).await?;
pub async fn completions(
&self,
body: &api::CompletionRequest,
) -> Result<api::CompletionResponse, OllamaError> {
let url = format!("{}/api/generate", self.base_url);
// let payload = json!({
// "model": model,
// "prompt": prompt,
// "stream": false,
// "options": Self::build_options(&body),
// });
let (prompt, model) = self.extract_completion_params(body)?;
// let res = self
// .client
// .post(format!("{}/api/generate", self.base_url))
// .json(&payload)
// .send()
// .await?
// .json::<Value>()
// .await?;
if prompt.is_empty() {
return Err(OllamaError::MissingPrompt);
}
// Ok(self.format_completion_response(&res))
// }
if model.is_empty() {
return Err(OllamaError::MissingModel);
}
// pub async fn completions_stream(
// &self,
// body: Value,
// ) -> Result<ReceiverStream<Result<Event, OllamaError>>, OllamaError> {
// let (prompt, model) = self.extract_completion_params(&body)?;
// self.validate_model(model).await?;
let exists = self.model_exists(model).await?;
if !exists {
return Err(OllamaError::ModelNotFound(model.to_string()));
}
// let payload = json!({
// "model": model,
// "prompt": prompt,
// "stream": true,
// "options": Self::build_options(&body),
// });
let options = ollama::OllamaOptions::from(body);
// let mut byte_stream = self
// .client
// .post(format!("{}/api/generate", self.base_url))
// .json(&payload)
// .send()
// .await?
// .bytes_stream();
let payload = ollama::OllamaGenerateRequest {
model,
prompt,
stream: false,
options,
};
// let (tx, rx) = tokio::sync::mpsc::channel(32);
let res = self
.client
.post(url)
.json(&payload)
.send()
.await?
.json::<ollama::OllamaGenerateResponse>()
.await?;
// tokio::spawn(async move {
// while let Some(chunk) = byte_stream.next().await {
// let chunk = match chunk {
// Ok(b) => b,
// Err(e) => {
// let _ = tx.send(Err(OllamaError::Http(e))).await;
// break;
// }
// };
Ok(api::CompletionResponse::from(res))
}
// if let Ok(json) = serde_json::from_slice::<Value>(&chunk) {
// let token = json.get("response").and_then(|v| v.as_str()).unwrap_or("");
// let done = json.get("done").and_then(|v| v.as_bool()).unwrap_or(false);
pub async fn completions_stream(
&self,
body: &api::CompletionRequest,
) -> Result<ReceiverStream<Result<Event, OllamaError>>, OllamaError> {
let url = format!("{}/api/generate", self.base_url);
// // OpenAI-compatible SSE chunk
// let event_data = serde_json::to_string(&json!({
// "id": "cmpl-ollama",
// "object": "text_completion",
// "choices": [{ "text": token, "index": 0, "finish_reason": null }],
// }))
// .unwrap_or_default();
let (prompt, model) = self.extract_completion_params(body)?;
// let _ = tx.send(Ok(Event::default().data(event_data))).await;
if prompt.is_empty() {
return Err(OllamaError::MissingPrompt);
}
// if done {
// // Final [DONE] sentinel — matches OpenAI streaming protocol
// let _ = tx.send(Ok(Event::default().data("[DONE]"))).await;
// break;
// }
// }
// }
// });
if model.is_empty() {
return Err(OllamaError::MissingModel);
}
// Ok(ReceiverStream::new(rx))
// }
let exists = self.model_exists(model).await?;
if !exists {
return Err(OllamaError::ModelNotFound(model.to_string()));
}
let options = ollama::OllamaOptions::from(body);
let payload = ollama::OllamaGenerateRequest {
model,
prompt,
stream: true,
options,
};
let mut byte_stream = self
.client
.post(url)
.json(&payload)
.send()
.await?
.bytes_stream();
let (tx, rx) = tokio::sync::mpsc::channel(32);
tokio::spawn(async move {
while let Some(chunk) = byte_stream.next().await {
let chunk = match chunk {
Ok(b) => b,
Err(e) => {
let _ = tx.send(Err(OllamaError::Http(e))).await;
break;
}
};
// 🔥 IMPORTANT: typed deserialization
let parsed: ollama::OllamaGenerateResponse = match serde_json::from_slice(&chunk) {
Ok(v) => v,
Err(_) => continue,
};
// map → OpenAI chunk
let event_data = serde_json::to_string(&api::CompletionChunk {
id: "cmpl-ollama".to_string(),
object: "text_completion".to_string(),
choices: vec![api::Choice {
text: parsed.response,
index: 0,
finish_reason: if parsed.done {
api::FinishReason::Stop
} else {
api::FinishReason::Length
},
}],
})
.unwrap_or_default();
let _ = tx.send(Ok(Event::default().data(event_data))).await;
if parsed.done {
let _ = tx.send(Ok(Event::default().data("[DONE]"))).await;
break;
}
}
});
Ok(ReceiverStream::new(rx))
}
// pub async fn chat_completions(&self, body: Value) -> Result<Value, OllamaError> {
// let (model, messages) = self.extract_chat_params(&body)?;
+40
View File
@@ -1,5 +1,8 @@
use crate::dto::{api, ollama};
use chrono::Utc;
use uuid::Uuid;
impl From<ollama::OllamaModel> for api::ModelInfo {
fn from(m: ollama::OllamaModel) -> Self {
Self {
@@ -14,3 +17,40 @@ impl From<ollama::OllamaModel> for api::ModelInfo {
}
}
}
impl From<&api::CompletionRequest> for ollama::OllamaOptions {
fn from(req: &api::CompletionRequest) -> Self {
Self {
temperature: req.base.temperature,
top_p: req.base.top_p,
top_k: req.base.top_k,
repeat_penalty: req.base.repeat_penalty,
seed: req.base.seed,
num_ctx: req.base.num_ctx,
num_predict: req.base.num_predict,
}
}
}
impl From<ollama::OllamaGenerateResponse> for api::CompletionResponse {
fn from(res: ollama::OllamaGenerateResponse) -> Self {
Self {
id: Uuid::new_v4().to_string(),
object: api::CompletionObject::TextCompletion,
model: res.model,
created: Utc::now().timestamp() as u64,
choices: vec![api::Choice {
text: res.response,
index: 0,
finish_reason: api::FinishReason::Stop,
}],
usage: api::Usage {
prompt_tokens: res.prompt_eval_count.unwrap_or(0),
completion_tokens: res.eval_count.unwrap_or(0),
total_tokens: res.prompt_eval_count.unwrap_or(0) + res.eval_count.unwrap_or(0),
},
}
}
}
+20 -22
View File
@@ -8,6 +8,7 @@ use axum::{
};
use serde_json::Value;
use crate::dto::api;
use crate::errors::OllamaError;
use crate::state::app_state::AppState;
@@ -18,31 +19,28 @@ use crate::state::app_state::AppState;
(status = 200, description = "Chat completion", body = Value),
)
)]
// pub async fn completions(
// State(state): State<AppState>,
// Json(body): Json<Value>,
// ) -> Result<Response, (axum::http::StatusCode, String)> {
// let wants_stream = body
// .get("stream")
// .and_then(|v| v.as_bool())
// .unwrap_or(false);
pub async fn completions(
State(state): State<AppState>,
Json(body): Json<api::CompletionRequest>,
) -> Result<Response, (axum::http::StatusCode, String)> {
let wants_stream = body.base.stream;
// if wants_stream {
// let stream = state
// .ollama
// .completions_stream(body)
// .await
// .map_err(ollama_err)?;
if wants_stream {
let stream = state
.ollama
.completions_stream(&body)
.await
.map_err(ollama_err)?;
// Ok(Sse::new(stream)
// .keep_alive(KeepAlive::default())
// .into_response())
// } else {
// let response = state.ollama.completions(body).await.map_err(ollama_err)?;
Ok(Sse::new(stream)
.keep_alive(KeepAlive::default())
.into_response())
} else {
let response = state.ollama.completions(&body).await.map_err(ollama_err)?;
// Ok(Json(response).into_response())
// }
// }
Ok(Json(response).into_response())
}
}
// pub async fn chat_completions(
// State(state): State<AppState>,
+6 -5
View File
@@ -11,11 +11,12 @@ use axum::{Router, middleware, routing::get, routing::post};
// }
pub fn protected_router() -> Router<AppState> {
Router::new().route("/models", get(models::list_models))
// .route("/completions", post(chat::completions))
// .route("/chat/completions", post(chat::chat_completions))
.route("/models/{model}/load", post(models::load_model))
.route("/models/{model}/unload", post(models::unload_model))
Router::new()
.route("/models", get(models::list_models))
.route("/completions", post(chat::completions))
// .route("/chat/completions", post(chat::chat_completions))
.route("/models/{model}/load", post(models::load_model))
.route("/models/{model}/unload", post(models::unload_model))
}
pub fn router() -> Router<AppState> {
-1
View File
@@ -37,7 +37,6 @@ pub async fn unload_model(
State(state): State<AppState>,
Path(model): Path<String>,
) -> Result<Json<api::UnloadModelResponse>, (axum::http::StatusCode, String)> {
let response = state
.ollama
.unload_model(&model)