This commit is contained in:
Generated
+1
@@ -154,6 +154,7 @@ dependencies = [
|
|||||||
"reqwest",
|
"reqwest",
|
||||||
"serde",
|
"serde",
|
||||||
"serde_json",
|
"serde_json",
|
||||||
|
"thiserror 2.0.18",
|
||||||
"tokio",
|
"tokio",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|||||||
@@ -12,3 +12,4 @@ jsonwebtoken = { version = "10.3.0", features = ["aws_lc_rs"] }
|
|||||||
reqwest = { version = "0.13.2", features = ["json"] }
|
reqwest = { version = "0.13.2", features = ["json"] }
|
||||||
once_cell = "1"
|
once_cell = "1"
|
||||||
dotenvy = "0.15"
|
dotenvy = "0.15"
|
||||||
|
thiserror = "2.0.18"
|
||||||
|
|||||||
@@ -0,0 +1,12 @@
|
|||||||
|
// errors.rs
|
||||||
|
use thiserror::Error;
|
||||||
|
|
||||||
|
#[derive(Debug, Error)]
|
||||||
|
pub enum OllamaError {
|
||||||
|
#[error("prompt is required and cannot be empty")]
|
||||||
|
MissingPrompt,
|
||||||
|
#[error("model '{0}' is not available — run `ollama pull {0}` first")]
|
||||||
|
ModelNotFound(String),
|
||||||
|
#[error(transparent)]
|
||||||
|
Http(#[from] reqwest::Error),
|
||||||
|
}
|
||||||
@@ -1,4 +1,5 @@
|
|||||||
mod auth;
|
mod auth;
|
||||||
|
mod errors;
|
||||||
mod providers;
|
mod providers;
|
||||||
mod routes;
|
mod routes;
|
||||||
mod state;
|
mod state;
|
||||||
|
|||||||
+75
-2
@@ -1,5 +1,7 @@
|
|||||||
use reqwest::Client;
|
use reqwest::Client;
|
||||||
use serde_json::Value;
|
use serde_json::{Value, json};
|
||||||
|
|
||||||
|
use crate::errors::OllamaError;
|
||||||
|
|
||||||
#[derive(Clone)]
|
#[derive(Clone)]
|
||||||
pub struct OllamaProvider {
|
pub struct OllamaProvider {
|
||||||
@@ -15,11 +17,82 @@ impl OllamaProvider {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_models(&self) -> Result<Value, reqwest::Error> {
|
pub async fn list_models(&self) -> Result<Value, OllamaError> {
|
||||||
let url = format!("{}/api/tags", self.base_url);
|
let url = format!("{}/api/tags", self.base_url);
|
||||||
|
|
||||||
let res = self.client.get(url).send().await?.json::<Value>().await?;
|
let res = self.client.get(url).send().await?.json::<Value>().await?;
|
||||||
|
|
||||||
Ok(res)
|
Ok(res)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub async fn completions(&self, body: Value) -> Result<Value, OllamaError> {
|
||||||
|
// Validate prompt
|
||||||
|
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())
|
||||||
|
.unwrap_or("llama3");
|
||||||
|
|
||||||
|
// Validate model availability
|
||||||
|
let available = self.list_models().await?; // Http error auto-converts via #[from]
|
||||||
|
let model_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 !model_exists {
|
||||||
|
return Err(OllamaError::ModelNotFound(model.to_string()));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build payload (now using validated locals)
|
||||||
|
let ollama_payload = json!({
|
||||||
|
"model": model,
|
||||||
|
"prompt": prompt,
|
||||||
|
"stream": false,
|
||||||
|
"options": {
|
||||||
|
"temperature": body.get("temperature"),
|
||||||
|
"top_p": body.get("top_p"),
|
||||||
|
"num_predict": body.get("max_tokens"),
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
let url = format!("{}/api/generate", self.base_url);
|
||||||
|
let res = self
|
||||||
|
.client
|
||||||
|
.post(&url)
|
||||||
|
.json(&ollama_payload)
|
||||||
|
.send()
|
||||||
|
.await?
|
||||||
|
.json::<Value>()
|
||||||
|
.await?;
|
||||||
|
|
||||||
|
Ok(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,
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,12 +1,12 @@
|
|||||||
pub mod models;
|
pub mod models;
|
||||||
|
|
||||||
use crate::auth::middleware::auth_middleware;
|
use crate::auth::middleware::auth_middleware;
|
||||||
use crate::routes::v1::models::list_models;
|
|
||||||
use crate::state::app_state::AppState;
|
use crate::state::app_state::AppState;
|
||||||
use axum::{Router, middleware, routing::get};
|
use axum::{Router, middleware, routing::get, routing::post};
|
||||||
|
|
||||||
pub fn router() -> Router<AppState> {
|
pub fn router() -> Router<AppState> {
|
||||||
Router::new()
|
Router::new()
|
||||||
.route("/models", get(list_models))
|
.route("/models", get(models::list_models))
|
||||||
|
.route("/completions", post(models::completions))
|
||||||
.layer(middleware::from_fn(auth_middleware))
|
.layer(middleware::from_fn(auth_middleware))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
use axum::{Json, extract::State};
|
use axum::{Json, extract::State};
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
|
|
||||||
|
use crate::errors::OllamaError;
|
||||||
use crate::state::app_state::AppState;
|
use crate::state::app_state::AppState;
|
||||||
|
|
||||||
pub async fn list_models(
|
pub async fn list_models(
|
||||||
@@ -14,3 +15,23 @@ pub async fn list_models(
|
|||||||
)),
|
)),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub async fn completions(
|
||||||
|
State(state): State<AppState>,
|
||||||
|
Json(body): Json<Value>,
|
||||||
|
) -> Result<Json<Value>, (axum::http::StatusCode, String)> {
|
||||||
|
match state.ollama.completions(body).await {
|
||||||
|
Ok(response) => Ok(Json(response)),
|
||||||
|
Err(OllamaError::MissingPrompt) => Err((
|
||||||
|
axum::http::StatusCode::BAD_REQUEST,
|
||||||
|
"prompt is required and cannot be empty".to_string(),
|
||||||
|
)),
|
||||||
|
Err(OllamaError::ModelNotFound(m)) => Err((
|
||||||
|
axum::http::StatusCode::UNPROCESSABLE_ENTITY,
|
||||||
|
format!("model '{m}' is not available — run `ollama pull {m}` first"),
|
||||||
|
)),
|
||||||
|
Err(OllamaError::Http(e)) => {
|
||||||
|
Err((axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user