Compare commits
3
Commits
6dc7231309
...
7e430bcf00
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7e430bcf00 | ||
|
|
1a06577fa3 | ||
|
|
db14d816cb |
@@ -6,15 +6,18 @@ use crate::auth::{
|
|||||||
};
|
};
|
||||||
|
|
||||||
pub async fn auth_middleware(mut request: Request, next: Next) -> Result<Response, StatusCode> {
|
pub async fn auth_middleware(mut request: Request, next: Next) -> Result<Response, StatusCode> {
|
||||||
dbg!("Middleware hit");
|
#[cfg(debug_assertions)]
|
||||||
|
println!("Middleware hit");
|
||||||
|
|
||||||
let headers = request.headers();
|
let headers = request.headers();
|
||||||
|
|
||||||
dbg!("Headers extracted");
|
#[cfg(debug_assertions)]
|
||||||
|
println!("Headers extracted");
|
||||||
|
|
||||||
let auth_header = headers.get("authorization").and_then(|v| v.to_str().ok());
|
let auth_header = headers.get("authorization").and_then(|v| v.to_str().ok());
|
||||||
|
|
||||||
dbg!("Auth header: {:?}", auth_header);
|
#[cfg(debug_assertions)]
|
||||||
|
println!("Auth header: {:?}", auth_header);
|
||||||
|
|
||||||
let auth_header = auth_header.ok_or(StatusCode::UNAUTHORIZED)?;
|
let auth_header = auth_header.ok_or(StatusCode::UNAUTHORIZED)?;
|
||||||
|
|
||||||
@@ -26,7 +29,8 @@ pub async fn auth_middleware(mut request: Request, next: Next) -> Result<Respons
|
|||||||
|
|
||||||
match validate_token(token, &jwks) {
|
match validate_token(token, &jwks) {
|
||||||
Ok(claims) => {
|
Ok(claims) => {
|
||||||
dbg!("Token valid");
|
#[cfg(debug_assertions)]
|
||||||
|
println!("Token valid");
|
||||||
|
|
||||||
request.extensions_mut().insert(claims);
|
request.extensions_mut().insert(claims);
|
||||||
|
|
||||||
|
|||||||
@@ -5,10 +5,18 @@ use thiserror::Error;
|
|||||||
pub enum OllamaError {
|
pub enum OllamaError {
|
||||||
#[error("prompt is required and cannot be empty")]
|
#[error("prompt is required and cannot be empty")]
|
||||||
MissingPrompt,
|
MissingPrompt,
|
||||||
|
|
||||||
#[error("messages must be a non-empty array containing at least one user message")]
|
#[error("messages must be a non-empty array containing at least one user message")]
|
||||||
MissingMessages,
|
MissingMessages,
|
||||||
|
|
||||||
#[error("model '{0}' is not available — run `ollama pull {0}` first")]
|
#[error("model '{0}' is not available — run `ollama pull {0}` first")]
|
||||||
ModelNotFound(String),
|
ModelNotFound(String),
|
||||||
|
|
||||||
|
#[error(
|
||||||
|
"invalid keep_alive format '{0}' — expected <number><unit> (e.g. 30s, 10m, 2h), a plain integer (seconds), or -1"
|
||||||
|
)]
|
||||||
|
InvalidKeepAlive(String),
|
||||||
|
|
||||||
#[error(transparent)]
|
#[error(transparent)]
|
||||||
Http(#[from] reqwest::Error),
|
Http(#[from] reqwest::Error),
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -43,6 +43,29 @@ impl OllamaProvider {
|
|||||||
Ok(())
|
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())),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// ── public endpoints ─────────────────────────────────────────────────────
|
// ── public endpoints ─────────────────────────────────────────────────────
|
||||||
|
|
||||||
pub async fn list_models(&self) -> Result<Value, OllamaError> {
|
pub async fn list_models(&self) -> Result<Value, OllamaError> {
|
||||||
@@ -51,6 +74,39 @@ impl OllamaProvider {
|
|||||||
Ok(res)
|
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 completions(&self, body: Value) -> Result<Value, OllamaError> {
|
pub async fn completions(&self, body: Value) -> Result<Value, OllamaError> {
|
||||||
let prompt = body
|
let prompt = body
|
||||||
.get("prompt")
|
.get("prompt")
|
||||||
|
|||||||
@@ -0,0 +1,44 @@
|
|||||||
|
use axum::{Json, extract::State};
|
||||||
|
use serde_json::Value;
|
||||||
|
|
||||||
|
use crate::errors::OllamaError;
|
||||||
|
use crate::state::app_state::AppState;
|
||||||
|
|
||||||
|
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(e) => Err((axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string())),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn chat_completions(
|
||||||
|
State(state): State<AppState>,
|
||||||
|
Json(body): Json<Value>,
|
||||||
|
) -> Result<Json<Value>, (axum::http::StatusCode, String)> {
|
||||||
|
match state.ollama.chat_completions(body).await {
|
||||||
|
Ok(response) => Ok(Json(response)),
|
||||||
|
Err(OllamaError::MissingMessages) => Err((
|
||||||
|
axum::http::StatusCode::BAD_REQUEST,
|
||||||
|
"messages must be a non-empty array with at least one user message".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()))
|
||||||
|
}
|
||||||
|
Err(e) => Err((axum::http::StatusCode::BAD_REQUEST, e.to_string())),
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,3 +1,4 @@
|
|||||||
|
pub mod chat;
|
||||||
pub mod models;
|
pub mod models;
|
||||||
|
|
||||||
use crate::auth::middleware::auth_middleware;
|
use crate::auth::middleware::auth_middleware;
|
||||||
@@ -7,7 +8,8 @@ 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(models::list_models))
|
.route("/models", get(models::list_models))
|
||||||
.route("/completions", post(models::completions))
|
.route("/completions", post(chat::completions))
|
||||||
.route("/chat/completions", post(models::chat_completions))
|
.route("/chat/completions", post(chat::chat_completions))
|
||||||
|
.route("/models/{model}/load", post(models::load_model))
|
||||||
.layer(middleware::from_fn(auth_middleware))
|
.layer(middleware::from_fn(auth_middleware))
|
||||||
}
|
}
|
||||||
|
|||||||
+39
-36
@@ -1,56 +1,59 @@
|
|||||||
use axum::{Json, extract::State};
|
use axum::{
|
||||||
|
Json,
|
||||||
|
extract::{Path, State},
|
||||||
|
};
|
||||||
|
use serde::Deserialize;
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
|
|
||||||
use crate::errors::OllamaError;
|
use crate::{errors::OllamaError, state::app_state::AppState};
|
||||||
use crate::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<Value>, (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(err) => Err((
|
Err(e) => Err(ollama_err(e)),
|
||||||
axum::http::StatusCode::INTERNAL_SERVER_ERROR,
|
|
||||||
err.to_string(),
|
|
||||||
)),
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn completions(
|
pub async fn load_model(
|
||||||
State(state): State<AppState>,
|
State(state): State<AppState>,
|
||||||
Json(body): Json<Value>,
|
Path(model): Path<String>,
|
||||||
|
Json(body): Json<LoadModelBody>,
|
||||||
) -> Result<Json<Value>, (axum::http::StatusCode, String)> {
|
) -> Result<Json<Value>, (axum::http::StatusCode, String)> {
|
||||||
match state.ollama.completions(body).await {
|
match state
|
||||||
|
.ollama
|
||||||
|
.load_model(&model, body.keep_alive.as_deref())
|
||||||
|
.await
|
||||||
|
{
|
||||||
Ok(response) => Ok(Json(response)),
|
Ok(response) => Ok(Json(response)),
|
||||||
Err(OllamaError::MissingPrompt) => Err((
|
Err(e) => Err(ollama_err(e)),
|
||||||
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(e) => Err((axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string())),
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn chat_completions(
|
fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) {
|
||||||
State(state): State<AppState>,
|
match e {
|
||||||
Json(body): Json<Value>,
|
OllamaError::ModelNotFound(m) => (
|
||||||
) -> Result<Json<Value>, (axum::http::StatusCode, String)> {
|
axum::http::StatusCode::NOT_FOUND,
|
||||||
match state.ollama.chat_completions(body).await {
|
format!("model '{m}' not found — run `ollama pull {m}`"),
|
||||||
Ok(response) => Ok(Json(response)),
|
),
|
||||||
Err(OllamaError::MissingMessages) => Err((
|
OllamaError::MissingPrompt => (
|
||||||
axum::http::StatusCode::BAD_REQUEST,
|
axum::http::StatusCode::BAD_REQUEST,
|
||||||
"messages must be a non-empty array with at least one user message".to_string(),
|
"prompt is required".to_string(),
|
||||||
)),
|
),
|
||||||
Err(OllamaError::ModelNotFound(m)) => Err((
|
OllamaError::MissingMessages => (
|
||||||
axum::http::StatusCode::UNPROCESSABLE_ENTITY,
|
axum::http::StatusCode::BAD_REQUEST,
|
||||||
format!("model '{m}' is not available — run `ollama pull {m}` first"),
|
"messages array with at least one user message is required".to_string(),
|
||||||
)),
|
),
|
||||||
Err(OllamaError::Http(e)) => {
|
OllamaError::InvalidKeepAlive(v) => (
|
||||||
Err((axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))
|
axum::http::StatusCode::BAD_REQUEST,
|
||||||
}
|
format!("invalid keep_alive '{v}' — use 30s / 10m / 2h, a plain integer, or -1"),
|
||||||
Err(e) => Err((axum::http::StatusCode::BAD_REQUEST, e.to_string())),
|
),
|
||||||
|
OllamaError::Http(e) => (axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string()),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user