feat: add proper typing for api
This commit is contained in:
@@ -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,
|
||||||
|
}
|
||||||
@@ -0,0 +1,2 @@
|
|||||||
|
pub mod api;
|
||||||
|
pub mod ollama;
|
||||||
@@ -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,2 +1,3 @@
|
|||||||
|
pub mod dto;
|
||||||
pub mod errors;
|
pub mod errors;
|
||||||
pub mod providers;
|
pub mod providers;
|
||||||
|
|||||||
+2
-1
@@ -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
@@ -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;
|
||||||
|
|||||||
@@ -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))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -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))
|
||||||
|
// }
|
||||||
|
}
|
||||||
@@ -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()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,2 @@
|
|||||||
|
pub mod client;
|
||||||
|
pub mod mapper;
|
||||||
+48
-48
@@ -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 {
|
||||||
|
|||||||
@@ -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
@@ -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) {
|
||||||
|
|||||||
@@ -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,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)]
|
||||||
|
|||||||
Reference in New Issue
Block a user