From 562d154480cbe5f22ae59636c7c4f83a7003ed07 Mon Sep 17 00:00:00 2001 From: LucasX Ubuntu Date: Fri, 10 Apr 2026 17:41:25 +0200 Subject: [PATCH] feat: add proper typing for api --- src/dto/api.rs | 93 ++++++++ src/dto/mod.rs | 2 + src/dto/ollama.rs | 27 +++ src/lib.rs | 1 + src/main.rs | 3 +- src/openapi.rs | 34 +-- src/providers/ollama.rs | 400 -------------------------------- src/providers/ollama/client.rs | 409 +++++++++++++++++++++++++++++++++ src/providers/ollama/mapper.rs | 16 ++ src/providers/ollama/mod.rs | 2 + src/routes/v1/chat.rs | 96 ++++---- src/routes/v1/mod.rs | 21 +- src/routes/v1/models.rs | 44 ++-- src/routes/v1/openapi.rs | 12 +- src/state/app_state.rs | 2 +- 15 files changed, 656 insertions(+), 506 deletions(-) create mode 100644 src/dto/api.rs create mode 100644 src/dto/mod.rs create mode 100644 src/dto/ollama.rs delete mode 100644 src/providers/ollama.rs create mode 100644 src/providers/ollama/client.rs create mode 100644 src/providers/ollama/mapper.rs create mode 100644 src/providers/ollama/mod.rs diff --git a/src/dto/api.rs b/src/dto/api.rs new file mode 100644 index 0000000..e3c0fea --- /dev/null +++ b/src/dto/api.rs @@ -0,0 +1,93 @@ +use serde::{Deserialize, Serialize}; +use utoipa::ToSchema; + +use crate::errors::OllamaError; + +#[derive(Serialize, Deserialize, ToSchema)] +pub struct ChatRequest { + pub model: String, + pub prompt: Option, + pub messages: Option>, + #[serde(default)] + pub stream: bool, + pub temperature: Option, + pub top_p: Option, + pub max_tokens: Option, + pub stop: Option>, + pub system: Option, + pub keep_alive: Option, +} + +#[derive(Debug, Serialize, Deserialize, ToSchema)] +#[serde(rename_all = "lowercase")] +pub enum Role { + System, + User, + Assistant, +} + +#[derive(Serialize, Deserialize, ToSchema)] +pub struct ChatResponse { + pub response: String, +} + +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct Message { + pub role: Role, + pub content: String, +} + +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct ModelsResponse { + pub models: Vec, +} + +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct ModelInfo { + pub name: String, + pub family: Option, + pub parameter_size: Option, + pub quantization: Option, +} + +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct LoadModelResponse { + pub model: String, + pub status: String, + pub keep_alive: String, +} + +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct LoadModelBody { + pub keep_alive: Option, +} + +impl LoadModelBody { + pub fn parse_keep_alive(s: &str) -> Result<(), OllamaError> { + let s = s.trim(); + + if s == "-1" || s.parse::().is_ok() { + return Ok(()); + } + + let split = s + .find(|c: char| c.is_alphabetic()) + .ok_or_else(|| OllamaError::InvalidKeepAlive(s.to_string()))?; + + let (num, unit) = s.split_at(split); + + num.parse::() + .map_err(|_| OllamaError::InvalidKeepAlive(s.to_string()))?; + + match unit { + "s" | "m" | "h" => Ok(()), + _ => Err(OllamaError::InvalidKeepAlive(s.to_string())), + } + } +} + +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct UnloadModelResponse { + pub model: String, + pub status: String, +} diff --git a/src/dto/mod.rs b/src/dto/mod.rs new file mode 100644 index 0000000..b6fe8b2 --- /dev/null +++ b/src/dto/mod.rs @@ -0,0 +1,2 @@ +pub mod api; +pub mod ollama; diff --git a/src/dto/ollama.rs b/src/dto/ollama.rs new file mode 100644 index 0000000..06b69a5 --- /dev/null +++ b/src/dto/ollama.rs @@ -0,0 +1,27 @@ +use serde::{Deserialize, Serialize}; +use utoipa::ToSchema; + +use crate::errors::OllamaError; + +#[derive(Debug, Serialize, Deserialize)] +pub struct OllamaModels { + pub models: Vec, +} + +#[derive(Debug, Serialize, Deserialize)] +pub struct OllamaModel { + pub name: String, + + pub details: Option, + + pub size: Option, + pub digest: Option, + pub modified_at: Option, +} + +#[derive(Debug, Serialize, Deserialize)] +pub struct OllamaModelDetails { + pub family: Option, + pub parameter_size: Option, + pub quantization_level: Option, +} diff --git a/src/lib.rs b/src/lib.rs index bd9bc1a..f99fbf7 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,2 +1,3 @@ +pub mod dto; pub mod errors; pub mod providers; diff --git a/src/main.rs b/src/main.rs index 135725d..125bbed 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,11 +1,12 @@ mod auth; +mod dto; mod errors; mod openapi; mod providers; mod routes; mod state; -use crate::providers::ollama::OllamaProvider; +use crate::providers::ollama::client::OllamaProvider; use crate::state::app_state::AppState; use axum::Router; diff --git a/src/openapi.rs b/src/openapi.rs index ed06820..f510f44 100644 --- a/src/openapi.rs +++ b/src/openapi.rs @@ -1,18 +1,18 @@ -use utoipa::OpenApi; +// use utoipa::OpenApi; -#[derive(OpenApi)] -#[openapi( - paths( - crate::routes::v1::chat::completions - ), - components( - schemas( - // add your request/response structs here later - ) - ), - tags( - (name = "chat", description = "Chat endpoints"), - (name = "models", description = "Model management") - ) -)] -pub struct V1ApiDoc; +// #[derive(OpenApi)] +// #[openapi( +// paths( +// crate::routes::v1::chat::completions +// ), +// components( +// schemas( +// // add your request/response structs here later +// ) +// ), +// tags( +// (name = "chat", description = "Chat endpoints"), +// (name = "models", description = "Model management") +// ) +// )] +// pub struct V1ApiDoc; diff --git a/src/providers/ollama.rs b/src/providers/ollama.rs deleted file mode 100644 index cd79d19..0000000 --- a/src/providers/ollama.rs +++ /dev/null @@ -1,400 +0,0 @@ -use crate::errors::OllamaError; -use axum::response::sse::Event; -use futures::StreamExt; -use reqwest::Client; -use serde_json::{Value, json}; -use tokio_stream::wrappers::ReceiverStream; - -#[derive(Clone)] -pub struct OllamaProvider { - pub client: Client, - pub base_url: String, -} - -impl OllamaProvider { - pub fn new(base_url: impl Into) -> Self { - Self { - client: Client::new(), - base_url: base_url.into(), - } - } - - // ── private helpers ────────────────────────────────────────────────────── - - fn build_options(body: &Value) -> Value { - json!({ - "temperature": body.get("temperature"), - "top_p": body.get("top_p"), - "num_predict": body.get("max_tokens"), - }) - } - - 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); - - if !exists { - return Err(OllamaError::ModelNotFound(model.to_string())); - } - Ok(()) - } - - fn parse_keep_alive(s: &str) -> Result<(), OllamaError> { - let s = s.trim(); - - // Ollama also accepts plain integers (seconds) or "-1" (load forever) - if s == "-1" || s.parse::().is_ok() { - return Ok(()); - } - - // Otherwise expect: e.g. "10m", "2h", "30s" - let (num, unit) = s - .find(|c: char| c.is_alphabetic()) - .map(|i| s.split_at(i)) - .ok_or_else(|| OllamaError::InvalidKeepAlive(s.to_string()))?; - - num.parse::() - .map_err(|_| OllamaError::InvalidKeepAlive(s.to_string()))?; - - match unit { - "s" | "m" | "h" => Ok(()), - _ => Err(OllamaError::InvalidKeepAlive(s.to_string())), - } - } - - 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)) - } - - fn format_completion_response(&self, res: &Value) -> Value { - json!({ - "id": "cmpl-ollama", - "object": "text_completion", - "model": res.get("model"), - "choices": [{ - "text": res.get("response"), - "index": 0, - "finish_reason": if res.get("done").and_then(|v| v.as_bool()).unwrap_or(false) { - "stop" - } else { - "length" - }, - }], - "usage": { - "prompt_tokens": res.get("prompt_eval_count"), - "completion_tokens": res.get("eval_count"), - "total_tokens": null, - } - }) - } - - fn extract_chat_params<'a>( - &self, - body: &'a Value, - ) -> Result<(&'a str, &'a Vec), OllamaError> { - let model = body - .get("model") - .and_then(|v| v.as_str()) - .unwrap_or("llama3"); - - let messages = body - .get("messages") - .and_then(|v| v.as_array()) - .filter(|arr| !arr.is_empty()) - .ok_or(OllamaError::MissingMessages)?; - - let has_user_msg = messages - .iter() - .any(|m| m.get("role").and_then(|r| r.as_str()) == Some("user")); - - if !has_user_msg { - return Err(OllamaError::MissingMessages); - } - - Ok((model, messages)) - } - - fn format_chat_response(&self, res: &Value) -> Value { - json!({ - "id": "chatcmpl-ollama", - "object": "chat.completion", - "model": res.get("model"), - "choices": [{ - "index": 0, - "message": { - "role": res.get("message").and_then(|m| m.get("role")), - "content": res.get("message").and_then(|m| m.get("content")), - }, - "finish_reason": res - .get("done_reason") - .and_then(|v| v.as_str()) - .unwrap_or("stop"), - }], - "usage": { - "prompt_tokens": res.get("prompt_eval_count"), - "completion_tokens": res.get("eval_count"), - "total_tokens": res.get("prompt_eval_count") - .and_then(|p| p.as_u64()) - .zip(res.get("eval_count").and_then(|e| e.as_u64())) - .map(|(p, e)| p + e), - } - }) - } - - // ── public endpoints ───────────────────────────────────────────────────── - - pub async fn list_models(&self) -> Result { - let url = format!("{}/api/tags", self.base_url); - let res = self.client.get(url).send().await?.json::().await?; - Ok(res) - } - - pub async fn load_model( - &self, - model: &str, - keep_alive: Option<&str>, - ) -> Result { - self.validate_model(model).await?; - - let keep_alive = keep_alive.unwrap_or("5m"); - Self::parse_keep_alive(keep_alive)?; // ← validated before any network call - - let payload = json!({ - "model": model, - "prompt": "", - "keep_alive": keep_alive, - "stream": false, - }); - - let res = self - .client - .post(format!("{}/api/generate", self.base_url)) - .json(&payload) - .send() - .await? - .json::() - .await?; - - Ok(json!({ - "model": res.get("model"), - "status": "loaded", - "keep_alive": keep_alive, - })) - } - - pub async fn unload_model(&self, model: &str) -> Result { - self.validate_model(model).await?; - - let payload = json!({ - "model": model, - "prompt": "", - "keep_alive": "0", - "stream": false, - }); - - let res = self - .client - .post(format!("{}/api/generate", self.base_url)) - .json(&payload) - .send() - .await? - .json::() - .await?; - - Ok(json!({ - "model": res.get("model"), - "status": "unloaded", - })) - } - - pub async fn completions(&self, body: Value) -> Result { - let (prompt, model) = self.extract_completion_params(&body)?; - self.validate_model(model).await?; - - let payload = json!({ - "model": model, - "prompt": prompt, - "stream": false, - "options": Self::build_options(&body), - }); - - let res = self - .client - .post(format!("{}/api/generate", self.base_url)) - .json(&payload) - .send() - .await? - .json::() - .await?; - - Ok(self.format_completion_response(&res)) - } - - pub async fn completions_stream( - &self, - body: Value, - ) -> Result>, OllamaError> { - let (prompt, model) = self.extract_completion_params(&body)?; - self.validate_model(model).await?; - - let payload = json!({ - "model": model, - "prompt": prompt, - "stream": true, - "options": Self::build_options(&body), - }); - - let mut byte_stream = self - .client - .post(format!("{}/api/generate", self.base_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; - } - }; - - if let Ok(json) = serde_json::from_slice::(&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); - - // 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 _ = tx.send(Ok(Event::default().data(event_data))).await; - - if done { - // Final [DONE] sentinel — matches OpenAI streaming protocol - let _ = tx.send(Ok(Event::default().data("[DONE]"))).await; - break; - } - } - } - }); - - Ok(ReceiverStream::new(rx)) - } - - pub async fn chat_completions(&self, body: Value) -> Result { - let (model, messages) = self.extract_chat_params(&body)?; - self.validate_model(model).await?; - - let payload = json!({ - "model": model, - "messages": messages, - "stream": false, - "options": Self::build_options(&body), - }); - - let res = self - .client - .post(format!("{}/api/chat", self.base_url)) - .json(&payload) - .send() - .await? - .json::() - .await?; - - Ok(self.format_chat_response(&res)) - } - - pub async fn chat_completions_stream( - &self, - body: Value, - ) -> Result>, OllamaError> { - let (model, messages) = self.extract_chat_params(&body)?; - self.validate_model(model).await?; - - let payload = json!({ - "model": model, - "messages": messages, - "stream": true, - "options": Self::build_options(&body), - }); - - let mut byte_stream = self - .client - .post(format!("{}/api/chat", self.base_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; - } - }; - - if let Ok(json) = serde_json::from_slice::(&chunk) { - let done = json.get("done").and_then(|v| v.as_bool()).unwrap_or(false); - - let event_data = serde_json::to_string(&json!({ - "id": "chatcmpl-ollama", - "object": "chat.completion.chunk", - "choices": [{ - "index": 0, - "delta": { - "role": json.get("message").and_then(|m| m.get("role")), - "content": json.get("message").and_then(|m| m.get("content")), - }, - "finish_reason": if done { json!("stop") } else { json!(null) }, - }], - })) - .unwrap_or_default(); - - let _ = tx.send(Ok(Event::default().data(event_data))).await; - - if done { - let _ = tx.send(Ok(Event::default().data("[DONE]"))).await; - break; - } - } - } - }); - - Ok(ReceiverStream::new(rx)) - } -} diff --git a/src/providers/ollama/client.rs b/src/providers/ollama/client.rs new file mode 100644 index 0000000..825611e --- /dev/null +++ b/src/providers/ollama/client.rs @@ -0,0 +1,409 @@ +use crate::dto::{api, ollama}; +use crate::errors::OllamaError; +use axum::response::sse::Event; +use futures::StreamExt; +use reqwest::Client; +use serde_json::{Value, json}; +use tokio_stream::wrappers::ReceiverStream; + +#[derive(Clone)] +pub struct OllamaProvider { + pub client: Client, + pub base_url: String, +} + +impl OllamaProvider { + pub fn new(base_url: impl Into) -> Self { + Self { + client: Client::new(), + base_url: base_url.into(), + } + } + + // ── private helpers ────────────────────────────────────────────────────── + + async fn model_exists(&self, model: &str) -> Result { + let url = format!("{}/api/tags", self.base_url); + + let res = self + .client + .get(url) + .send() + .await? + .json::() + .await?; + + 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"), + // }) + // } + + // 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); + + // 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)) + // } + + // fn format_completion_response(&self, res: &Value) -> Value { + // json!({ + // "id": "cmpl-ollama", + // "object": "text_completion", + // "model": res.get("model"), + // "choices": [{ + // "text": res.get("response"), + // "index": 0, + // "finish_reason": if res.get("done").and_then(|v| v.as_bool()).unwrap_or(false) { + // "stop" + // } else { + // "length" + // }, + // }], + // "usage": { + // "prompt_tokens": res.get("prompt_eval_count"), + // "completion_tokens": res.get("eval_count"), + // "total_tokens": null, + // } + // }) + // } + + // fn extract_chat_params<'a>( + // &self, + // body: &'a Value, + // ) -> Result<(&'a str, &'a Vec), OllamaError> { + // let model = body + // .get("model") + // .and_then(|v| v.as_str()) + // .unwrap_or("llama3"); + + // let messages = body + // .get("messages") + // .and_then(|v| v.as_array()) + // .filter(|arr| !arr.is_empty()) + // .ok_or(OllamaError::MissingMessages)?; + + // let has_user_msg = messages + // .iter() + // .any(|m| m.get("role").and_then(|r| r.as_str()) == Some("user")); + + // if !has_user_msg { + // return Err(OllamaError::MissingMessages); + // } + + // Ok((model, messages)) + // } + + // fn format_chat_response(&self, res: &Value) -> Value { + // json!({ + // "id": "chatcmpl-ollama", + // "object": "chat.completion", + // "model": res.get("model"), + // "choices": [{ + // "index": 0, + // "message": { + // "role": res.get("message").and_then(|m| m.get("role")), + // "content": res.get("message").and_then(|m| m.get("content")), + // }, + // "finish_reason": res + // .get("done_reason") + // .and_then(|v| v.as_str()) + // .unwrap_or("stop"), + // }], + // "usage": { + // "prompt_tokens": res.get("prompt_eval_count"), + // "completion_tokens": res.get("eval_count"), + // "total_tokens": res.get("prompt_eval_count") + // .and_then(|p| p.as_u64()) + // .zip(res.get("eval_count").and_then(|e| e.as_u64())) + // .map(|(p, e)| p + e), + // } + // }) + // } + + // // ── public endpoints ───────────────────────────────────────────────────── + + pub async fn list_models(&self) -> Result { + let url = format!("{}/api/tags", self.base_url); + + let res = self + .client + .get(url) + .send() + .await? + .json::() + .await?; + + let models = res.models.into_iter().map(api::ModelInfo::from).collect(); + + Ok(api::ModelsResponse { models }) + } + + pub async fn load_model( + &self, + model: &str, + keep_alive: &str, + ) -> Result { + let url = format!("{}/api/generate", self.base_url); + + let exists = self.model_exists(model).await?; + if !exists { + return Err(OllamaError::ModelNotFound(model.to_string())); + } + + let payload = json!({ + "model": model, + "prompt": "", + "keep_alive": keep_alive, + "stream": false, + }); + + let _res = self + .client + .post(url) + .json(&payload) + .send() + .await? + .json::() + .await?; + + Ok(api::LoadModelResponse { + model: model.to_string(), + status: "loaded".to_string(), + keep_alive: keep_alive.to_string(), + }) + } + + pub async fn unload_model(&self, model: &str) -> Result { + let url = format!("{}/api/generate", self.base_url); + + let exists = self.model_exists(model).await?; + if !exists { + return Err(OllamaError::ModelNotFound(model.to_string())); + } + + let payload = json!({ + "model": model, + "prompt": "", + "keep_alive": "0", + "stream": false, + }); + + let _res = self + .client + .post(url) + .json(&payload) + .send() + .await? + .json::() + .await?; + + Ok(api::UnloadModelResponse { + model: model.to_string(), + status: "unloaded".to_string(), + }) + } + + // pub async fn completions(&self, body: Value) -> Result { + // let (prompt, model) = self.extract_completion_params(&body)?; + // self.validate_model(model).await?; + + // let payload = json!({ + // "model": model, + // "prompt": prompt, + // "stream": false, + // "options": Self::build_options(&body), + // }); + + // let res = self + // .client + // .post(format!("{}/api/generate", self.base_url)) + // .json(&payload) + // .send() + // .await? + // .json::() + // .await?; + + // Ok(self.format_completion_response(&res)) + // } + + // pub async fn completions_stream( + // &self, + // body: Value, + // ) -> Result>, OllamaError> { + // let (prompt, model) = self.extract_completion_params(&body)?; + // self.validate_model(model).await?; + + // let payload = json!({ + // "model": model, + // "prompt": prompt, + // "stream": true, + // "options": Self::build_options(&body), + // }); + + // let mut byte_stream = self + // .client + // .post(format!("{}/api/generate", self.base_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; + // } + // }; + + // if let Ok(json) = serde_json::from_slice::(&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); + + // // 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 _ = tx.send(Ok(Event::default().data(event_data))).await; + + // if done { + // // Final [DONE] sentinel — matches OpenAI streaming protocol + // let _ = tx.send(Ok(Event::default().data("[DONE]"))).await; + // break; + // } + // } + // } + // }); + + // Ok(ReceiverStream::new(rx)) + // } + + // pub async fn chat_completions(&self, body: Value) -> Result { + // let (model, messages) = self.extract_chat_params(&body)?; + // self.validate_model(model).await?; + + // let payload = json!({ + // "model": model, + // "messages": messages, + // "stream": false, + // "options": Self::build_options(&body), + // }); + + // let res = self + // .client + // .post(format!("{}/api/chat", self.base_url)) + // .json(&payload) + // .send() + // .await? + // .json::() + // .await?; + + // Ok(self.format_chat_response(&res)) + // } + + // pub async fn chat_completions_stream( + // &self, + // body: Value, + // ) -> Result>, OllamaError> { + // let (model, messages) = self.extract_chat_params(&body)?; + // self.validate_model(model).await?; + + // let payload = json!({ + // "model": model, + // "messages": messages, + // "stream": true, + // "options": Self::build_options(&body), + // }); + + // let mut byte_stream = self + // .client + // .post(format!("{}/api/chat", self.base_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; + // } + // }; + + // if let Ok(json) = serde_json::from_slice::(&chunk) { + // let done = json.get("done").and_then(|v| v.as_bool()).unwrap_or(false); + + // let event_data = serde_json::to_string(&json!({ + // "id": "chatcmpl-ollama", + // "object": "chat.completion.chunk", + // "choices": [{ + // "index": 0, + // "delta": { + // "role": json.get("message").and_then(|m| m.get("role")), + // "content": json.get("message").and_then(|m| m.get("content")), + // }, + // "finish_reason": if done { json!("stop") } else { json!(null) }, + // }], + // })) + // .unwrap_or_default(); + + // let _ = tx.send(Ok(Event::default().data(event_data))).await; + + // if done { + // let _ = tx.send(Ok(Event::default().data("[DONE]"))).await; + // break; + // } + // } + // } + // }); + + // Ok(ReceiverStream::new(rx)) + // } +} diff --git a/src/providers/ollama/mapper.rs b/src/providers/ollama/mapper.rs new file mode 100644 index 0000000..b7e1052 --- /dev/null +++ b/src/providers/ollama/mapper.rs @@ -0,0 +1,16 @@ +use crate::dto::{api, ollama}; + +impl From for api::ModelInfo { + fn from(m: ollama::OllamaModel) -> Self { + Self { + name: m.name, + + family: m.details.as_ref().and_then(|d| d.family.clone()), + parameter_size: m.details.as_ref().and_then(|d| d.parameter_size.clone()), + quantization: m + .details + .as_ref() + .and_then(|d| d.quantization_level.clone()), + } + } +} diff --git a/src/providers/ollama/mod.rs b/src/providers/ollama/mod.rs new file mode 100644 index 0000000..66ad831 --- /dev/null +++ b/src/providers/ollama/mod.rs @@ -0,0 +1,2 @@ +pub mod client; +pub mod mapper; diff --git a/src/routes/v1/chat.rs b/src/routes/v1/chat.rs index 213e168..2fa030f 100644 --- a/src/routes/v1/chat.rs +++ b/src/routes/v1/chat.rs @@ -18,61 +18,61 @@ use crate::state::app_state::AppState; (status = 200, description = "Chat completion", body = Value), ) )] -pub async fn completions( - State(state): State, - Json(body): Json, -) -> Result { - let wants_stream = body - .get("stream") - .and_then(|v| v.as_bool()) - .unwrap_or(false); +// pub async fn completions( +// State(state): State, +// Json(body): Json, +// ) -> Result { +// let wants_stream = body +// .get("stream") +// .and_then(|v| v.as_bool()) +// .unwrap_or(false); - 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, - Json(body): Json, -) -> Result { - let wants_stream = body - .get("stream") - .and_then(|v| v.as_bool()) - .unwrap_or(false); +// pub async fn chat_completions( +// State(state): State, +// Json(body): Json, +// ) -> Result { +// let wants_stream = body +// .get("stream") +// .and_then(|v| v.as_bool()) +// .unwrap_or(false); - if wants_stream { - let stream = state - .ollama - .chat_completions_stream(body) - .await - .map_err(ollama_err)?; +// if wants_stream { +// let stream = state +// .ollama +// .chat_completions_stream(body) +// .await +// .map_err(ollama_err)?; - Ok(Sse::new(stream) - .keep_alive(KeepAlive::default()) - .into_response()) - } else { - let response = state - .ollama - .chat_completions(body) - .await - .map_err(ollama_err)?; +// Ok(Sse::new(stream) +// .keep_alive(KeepAlive::default()) +// .into_response()) +// } else { +// let response = state +// .ollama +// .chat_completions(body) +// .await +// .map_err(ollama_err)?; - Ok(Json(response).into_response()) - } -} +// Ok(Json(response).into_response()) +// } +// } fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) { match e { diff --git a/src/routes/v1/mod.rs b/src/routes/v1/mod.rs index 78112f9..8c7246b 100644 --- a/src/routes/v1/mod.rs +++ b/src/routes/v1/mod.rs @@ -6,21 +6,20 @@ use crate::auth::middleware::auth_middleware; use crate::state::app_state::AppState; use axum::{Router, middleware, routing::get, routing::post}; -fn public_router() -> Router { - Router::new().route("/openapi.json", get(openapi::openapi_json)) -} +// fn public_router() -> Router { +// Router::new().route("/openapi.json", get(openapi::openapi_json)) +// } pub fn protected_router() -> Router { - 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 { Router::new() - .merge(public_router()) - .merge(protected_router().layer(middleware::from_fn(auth_middleware))) + // .merge(public_router()) + .merge(protected_router()) //.layer(middleware::from_fn(auth_middleware))) } diff --git a/src/routes/v1/models.rs b/src/routes/v1/models.rs index 19aa281..d8a2cfb 100644 --- a/src/routes/v1/models.rs +++ b/src/routes/v1/models.rs @@ -2,19 +2,13 @@ use axum::{ Json, extract::{Path, State}, }; -use serde::Deserialize; -use serde_json::Value; +use crate::dto::api; use crate::{errors::OllamaError, state::app_state::AppState}; -#[derive(Deserialize)] -pub struct LoadModelBody { - pub keep_alive: Option, -} - pub async fn list_models( State(state): State, -) -> Result, (axum::http::StatusCode, String)> { +) -> Result, (axum::http::StatusCode, String)> { match state.ollama.list_models().await { Ok(models) => Ok(Json(models)), Err(e) => Err(ollama_err(e)), @@ -24,27 +18,33 @@ pub async fn list_models( pub async fn load_model( State(state): State, Path(model): Path, - Json(body): Json, -) -> Result, (axum::http::StatusCode, String)> { - match state + Json(body): Json, +) -> Result, (axum::http::StatusCode, String)> { + let keep_alive = body.keep_alive.as_deref().unwrap_or("5m"); + + api::LoadModelBody::parse_keep_alive(keep_alive).map_err(ollama_err)?; + + let response = state .ollama - .load_model(&model, body.keep_alive.as_deref()) + .load_model(&model, keep_alive) .await - { - Ok(response) => Ok(Json(response)), - Err(e) => Err(ollama_err(e)), - } + .map_err(ollama_err)?; + + Ok(Json(response)) } pub async fn unload_model( State(state): State, Path(model): Path, -) -> Result, (axum::http::StatusCode, String)> { - match state.ollama.unload_model(&model).await { - // ← correct method - Ok(response) => Ok(Json(response)), - Err(e) => Err(ollama_err(e)), - } +) -> Result, (axum::http::StatusCode, String)> { + + let response = state + .ollama + .unload_model(&model) + .await + .map_err(ollama_err)?; + + Ok(Json(response)) } fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) { diff --git a/src/routes/v1/openapi.rs b/src/routes/v1/openapi.rs index bb53a58..b5e5508 100644 --- a/src/routes/v1/openapi.rs +++ b/src/routes/v1/openapi.rs @@ -1,8 +1,8 @@ -use axum::Json; -use utoipa::OpenApi; +// use axum::Json; +// use utoipa::OpenApi; -use crate::openapi::V1ApiDoc; +// use crate::openapi::V1ApiDoc; -pub async fn openapi_json() -> Json { - Json(V1ApiDoc::openapi()) -} +// pub async fn openapi_json() -> Json { +// Json(V1ApiDoc::openapi()) +// } diff --git a/src/state/app_state.rs b/src/state/app_state.rs index 1fb247b..36e893a 100644 --- a/src/state/app_state.rs +++ b/src/state/app_state.rs @@ -1,4 +1,4 @@ -use crate::providers::ollama::OllamaProvider; +use crate::providers::ollama::client::OllamaProvider; use std::sync::Arc; #[derive(Clone)]