Compare commits

...
3 Commits
Author SHA1 Message Date
LucasDLTG 7e430bcf00 fix auth debug infos only in debug mode
CI / Rust CI (push) Successful in 1m26s
2026-04-10 12:35:00 +02:00
LucasDLTG 1a06577fa3 feat: add load endpoint 2026-04-10 12:29:34 +02:00
LucasDLTG db14d816cb format: separate routes in file system 2026-04-10 12:08:03 +02:00
6 changed files with 159 additions and 42 deletions
+8 -4
View File
@@ -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);
+8
View File
@@ -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),
} }
+56
View File
@@ -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")
+44
View File
@@ -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())),
}
}
+4 -2
View File
@@ -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
View File
@@ -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()),
} }
} }