feat: add proper typing for api

This commit is contained in:
2026-04-10 17:41:25 +02:00
parent a22560c337
commit 562d154480
15 changed files with 656 additions and 506 deletions
+93
View File
@@ -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<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 {
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<ModelInfo>,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ModelInfo {
pub name: String,
pub family: Option<String>,
pub parameter_size: Option<String>,
pub quantization: Option<String>,
}
#[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<String>,
}
impl LoadModelBody {
pub fn parse_keep_alive(s: &str) -> Result<(), OllamaError> {
let s = s.trim();
if s == "-1" || s.parse::<u64>().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::<u64>()
.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,
}
+2
View File
@@ -0,0 +1,2 @@
pub mod api;
pub mod ollama;
+27
View File
@@ -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<OllamaModel>,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct OllamaModel {
pub name: String,
pub details: Option<OllamaModelDetails>,
pub size: Option<u64>,
pub digest: Option<String>,
pub modified_at: Option<String>,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct OllamaModelDetails {
pub family: Option<String>,
pub parameter_size: Option<String>,
pub quantization_level: Option<String>,
}
+1
View File
@@ -1,2 +1,3 @@
pub mod dto;
pub mod errors; pub mod errors;
pub mod providers; pub mod providers;
+2 -1
View File
@@ -1,11 +1,12 @@
mod auth; mod auth;
mod dto;
mod errors; mod errors;
mod openapi; mod openapi;
mod providers; mod providers;
mod routes; mod routes;
mod state; mod state;
use crate::providers::ollama::OllamaProvider; use crate::providers::ollama::client::OllamaProvider;
use crate::state::app_state::AppState; use crate::state::app_state::AppState;
use axum::Router; use axum::Router;
+17 -17
View File
@@ -1,18 +1,18 @@
use utoipa::OpenApi; // use utoipa::OpenApi;
#[derive(OpenApi)] // #[derive(OpenApi)]
#[openapi( // #[openapi(
paths( // paths(
crate::routes::v1::chat::completions // crate::routes::v1::chat::completions
), // ),
components( // components(
schemas( // schemas(
// add your request/response structs here later // // add your request/response structs here later
) // )
), // ),
tags( // tags(
(name = "chat", description = "Chat endpoints"), // (name = "chat", description = "Chat endpoints"),
(name = "models", description = "Model management") // (name = "models", description = "Model management")
) // )
)] // )]
pub struct V1ApiDoc; // pub struct V1ApiDoc;
-400
View File
@@ -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<String>) -> 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::<u64>().is_ok() {
return Ok(());
}
// Otherwise expect: <number><unit> 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::<u64>()
.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<Value>), 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<Value, OllamaError> {
let url = format!("{}/api/tags", self.base_url);
let res = self.client.get(url).send().await?.json::<Value>().await?;
Ok(res)
}
pub async fn load_model(
&self,
model: &str,
keep_alive: Option<&str>,
) -> Result<Value, OllamaError> {
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::<Value>()
.await?;
Ok(json!({
"model": res.get("model"),
"status": "loaded",
"keep_alive": keep_alive,
}))
}
pub async fn unload_model(&self, model: &str) -> Result<Value, OllamaError> {
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::<Value>()
.await?;
Ok(json!({
"model": res.get("model"),
"status": "unloaded",
}))
}
pub async fn completions(&self, body: Value) -> Result<Value, OllamaError> {
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::<Value>()
.await?;
Ok(self.format_completion_response(&res))
}
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 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::<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);
// 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<Value, OllamaError> {
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::<Value>()
.await?;
Ok(self.format_chat_response(&res))
}
pub async fn chat_completions_stream(
&self,
body: Value,
) -> Result<ReceiverStream<Result<Event, OllamaError>>, 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::<Value>(&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))
}
}
+409
View File
@@ -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<String>) -> Self {
Self {
client: Client::new(),
base_url: base_url.into(),
}
}
// ── private helpers ──────────────────────────────────────────────────────
async fn model_exists(&self, model: &str) -> Result<bool, OllamaError> {
let url = format!("{}/api/tags", self.base_url);
let res = self
.client
.get(url)
.send()
.await?
.json::<ollama::OllamaModels>()
.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<Value>), 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<api::ModelsResponse, OllamaError> {
let url = format!("{}/api/tags", self.base_url);
let res = self
.client
.get(url)
.send()
.await?
.json::<ollama::OllamaModels>()
.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<api::LoadModelResponse, 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()));
}
let payload = json!({
"model": model,
"prompt": "",
"keep_alive": keep_alive,
"stream": false,
});
let _res = self
.client
.post(url)
.json(&payload)
.send()
.await?
.json::<Value>()
.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<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()));
}
let payload = json!({
"model": model,
"prompt": "",
"keep_alive": "0",
"stream": false,
});
let _res = self
.client
.post(url)
.json(&payload)
.send()
.await?
.json::<Value>()
.await?;
Ok(api::UnloadModelResponse {
model: model.to_string(),
status: "unloaded".to_string(),
})
}
// pub async fn completions(&self, body: Value) -> Result<Value, OllamaError> {
// 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::<Value>()
// .await?;
// Ok(self.format_completion_response(&res))
// }
// 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 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::<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);
// // 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<Value, OllamaError> {
// 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::<Value>()
// .await?;
// Ok(self.format_chat_response(&res))
// }
// pub async fn chat_completions_stream(
// &self,
// body: Value,
// ) -> Result<ReceiverStream<Result<Event, OllamaError>>, 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::<Value>(&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))
// }
}
+16
View File
@@ -0,0 +1,16 @@
use crate::dto::{api, ollama};
impl From<ollama::OllamaModel> 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()),
}
}
}
+2
View File
@@ -0,0 +1,2 @@
pub mod client;
pub mod mapper;
+48 -48
View File
@@ -18,61 +18,61 @@ use crate::state::app_state::AppState;
(status = 200, description = "Chat completion", body = Value), (status = 200, description = "Chat completion", body = Value),
) )
)] )]
pub async fn completions( // pub async fn completions(
State(state): State<AppState>, // State(state): State<AppState>,
Json(body): Json<Value>, // Json(body): Json<Value>,
) -> Result<Response, (axum::http::StatusCode, String)> { // ) -> Result<Response, (axum::http::StatusCode, String)> {
let wants_stream = body // let wants_stream = body
.get("stream") // .get("stream")
.and_then(|v| v.as_bool()) // .and_then(|v| v.as_bool())
.unwrap_or(false); // .unwrap_or(false);
if wants_stream { // if wants_stream {
let stream = state // let stream = state
.ollama // .ollama
.completions_stream(body) // .completions_stream(body)
.await // .await
.map_err(ollama_err)?; // .map_err(ollama_err)?;
Ok(Sse::new(stream) // Ok(Sse::new(stream)
.keep_alive(KeepAlive::default()) // .keep_alive(KeepAlive::default())
.into_response()) // .into_response())
} else { // } else {
let response = state.ollama.completions(body).await.map_err(ollama_err)?; // 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( // pub async fn chat_completions(
State(state): State<AppState>, // State(state): State<AppState>,
Json(body): Json<Value>, // Json(body): Json<Value>,
) -> Result<Response, (axum::http::StatusCode, String)> { // ) -> Result<Response, (axum::http::StatusCode, String)> {
let wants_stream = body // let wants_stream = body
.get("stream") // .get("stream")
.and_then(|v| v.as_bool()) // .and_then(|v| v.as_bool())
.unwrap_or(false); // .unwrap_or(false);
if wants_stream { // if wants_stream {
let stream = state // let stream = state
.ollama // .ollama
.chat_completions_stream(body) // .chat_completions_stream(body)
.await // .await
.map_err(ollama_err)?; // .map_err(ollama_err)?;
Ok(Sse::new(stream) // Ok(Sse::new(stream)
.keep_alive(KeepAlive::default()) // .keep_alive(KeepAlive::default())
.into_response()) // .into_response())
} else { // } else {
let response = state // let response = state
.ollama // .ollama
.chat_completions(body) // .chat_completions(body)
.await // .await
.map_err(ollama_err)?; // .map_err(ollama_err)?;
Ok(Json(response).into_response()) // Ok(Json(response).into_response())
} // }
} // }
fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) { fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) {
match e { match e {
+8 -9
View File
@@ -6,21 +6,20 @@ use crate::auth::middleware::auth_middleware;
use crate::state::app_state::AppState; use crate::state::app_state::AppState;
use axum::{Router, middleware, routing::get, routing::post}; use axum::{Router, middleware, routing::get, routing::post};
fn public_router() -> Router<AppState> { // fn public_router() -> Router<AppState> {
Router::new().route("/openapi.json", get(openapi::openapi_json)) // Router::new().route("/openapi.json", get(openapi::openapi_json))
} // }
pub fn protected_router() -> Router<AppState> { pub fn protected_router() -> Router<AppState> {
Router::new() Router::new().route("/models", get(models::list_models))
.route("/models", get(models::list_models)) // .route("/completions", post(chat::completions))
.route("/completions", post(chat::completions)) // .route("/chat/completions", post(chat::chat_completions))
.route("/chat/completions", post(chat::chat_completions))
.route("/models/{model}/load", post(models::load_model)) .route("/models/{model}/load", post(models::load_model))
.route("/models/{model}/unload", post(models::unload_model)) .route("/models/{model}/unload", post(models::unload_model))
} }
pub fn router() -> Router<AppState> { pub fn router() -> Router<AppState> {
Router::new() Router::new()
.merge(public_router()) // .merge(public_router())
.merge(protected_router().layer(middleware::from_fn(auth_middleware))) .merge(protected_router()) //.layer(middleware::from_fn(auth_middleware)))
} }
+22 -22
View File
@@ -2,19 +2,13 @@ use axum::{
Json, Json,
extract::{Path, State}, extract::{Path, State},
}; };
use serde::Deserialize;
use serde_json::Value;
use crate::dto::api;
use crate::{errors::OllamaError, state::app_state::AppState}; use crate::{errors::OllamaError, state::app_state::AppState};
#[derive(Deserialize)]
pub struct LoadModelBody {
pub keep_alive: Option<String>,
}
pub async fn list_models( pub async fn list_models(
State(state): State<AppState>, State(state): State<AppState>,
) -> Result<Json<Value>, (axum::http::StatusCode, String)> { ) -> Result<Json<api::ModelsResponse>, (axum::http::StatusCode, String)> {
match state.ollama.list_models().await { match state.ollama.list_models().await {
Ok(models) => Ok(Json(models)), Ok(models) => Ok(Json(models)),
Err(e) => Err(ollama_err(e)), Err(e) => Err(ollama_err(e)),
@@ -24,27 +18,33 @@ pub async fn list_models(
pub async fn load_model( pub async fn load_model(
State(state): State<AppState>, State(state): State<AppState>,
Path(model): Path<String>, Path(model): Path<String>,
Json(body): Json<LoadModelBody>, Json(body): Json<api::LoadModelBody>,
) -> Result<Json<Value>, (axum::http::StatusCode, String)> { ) -> Result<Json<api::LoadModelResponse>, (axum::http::StatusCode, String)> {
match state 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 .ollama
.load_model(&model, body.keep_alive.as_deref()) .load_model(&model, keep_alive)
.await .await
{ .map_err(ollama_err)?;
Ok(response) => Ok(Json(response)),
Err(e) => Err(ollama_err(e)), Ok(Json(response))
}
} }
pub async fn unload_model( pub async fn unload_model(
State(state): State<AppState>, State(state): State<AppState>,
Path(model): Path<String>, Path(model): Path<String>,
) -> Result<Json<Value>, (axum::http::StatusCode, String)> { ) -> Result<Json<api::UnloadModelResponse>, (axum::http::StatusCode, String)> {
match state.ollama.unload_model(&model).await {
// ← correct method let response = state
Ok(response) => Ok(Json(response)), .ollama
Err(e) => Err(ollama_err(e)), .unload_model(&model)
} .await
.map_err(ollama_err)?;
Ok(Json(response))
} }
fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) { fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) {
+6 -6
View File
@@ -1,8 +1,8 @@
use axum::Json; // use axum::Json;
use utoipa::OpenApi; // use utoipa::OpenApi;
use crate::openapi::V1ApiDoc; // use crate::openapi::V1ApiDoc;
pub async fn openapi_json() -> Json<utoipa::openapi::OpenApi> { // pub async fn openapi_json() -> Json<utoipa::openapi::OpenApi> {
Json(V1ApiDoc::openapi()) // Json(V1ApiDoc::openapi())
} // }
+1 -1
View File
@@ -1,4 +1,4 @@
use crate::providers::ollama::OllamaProvider; use crate::providers::ollama::client::OllamaProvider;
use std::sync::Arc; use std::sync::Arc;
#[derive(Clone)] #[derive(Clone)]