From 1a06577fa3b7fd04bdb9990f6c53ddb6b489bbdc Mon Sep 17 00:00:00 2001 From: LucasX Ubuntu Date: Fri, 10 Apr 2026 12:29:34 +0200 Subject: [PATCH] feat: add load endpoint --- src/errors.rs | 8 ++++++ src/providers/ollama.rs | 56 +++++++++++++++++++++++++++++++++++++++++ src/routes/v1/mod.rs | 1 + src/routes/v1/models.rs | 55 +++++++++++++++++++++++++++++++++++----- 4 files changed, 114 insertions(+), 6 deletions(-) diff --git a/src/errors.rs b/src/errors.rs index 355de13..e36c286 100644 --- a/src/errors.rs +++ b/src/errors.rs @@ -5,10 +5,18 @@ use thiserror::Error; pub enum OllamaError { #[error("prompt is required and cannot be empty")] MissingPrompt, + #[error("messages must be a non-empty array containing at least one user message")] MissingMessages, + #[error("model '{0}' is not available — run `ollama pull {0}` first")] ModelNotFound(String), + + #[error( + "invalid keep_alive format '{0}' — expected (e.g. 30s, 10m, 2h), a plain integer (seconds), or -1" + )] + InvalidKeepAlive(String), + #[error(transparent)] Http(#[from] reqwest::Error), } diff --git a/src/providers/ollama.rs b/src/providers/ollama.rs index b448df3..979cd2f 100644 --- a/src/providers/ollama.rs +++ b/src/providers/ollama.rs @@ -43,6 +43,29 @@ impl OllamaProvider { 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::().is_ok() { + return Ok(()); + } + + // Otherwise expect: 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::() + .map_err(|_| OllamaError::InvalidKeepAlive(s.to_string()))?; + + match unit { + "s" | "m" | "h" => Ok(()), + _ => Err(OllamaError::InvalidKeepAlive(s.to_string())), + } + } + // ── public endpoints ───────────────────────────────────────────────────── pub async fn list_models(&self) -> Result { @@ -51,6 +74,39 @@ impl OllamaProvider { Ok(res) } + pub async fn load_model( + &self, + model: &str, + keep_alive: Option<&str>, + ) -> Result { + 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::() + .await?; + + Ok(json!({ + "model": res.get("model"), + "status": "loaded", + "keep_alive": keep_alive, + })) + } + pub async fn completions(&self, body: Value) -> Result { let prompt = body .get("prompt") diff --git a/src/routes/v1/mod.rs b/src/routes/v1/mod.rs index 768f9a2..9fdecff 100644 --- a/src/routes/v1/mod.rs +++ b/src/routes/v1/mod.rs @@ -10,5 +10,6 @@ pub fn router() -> Router { .route("/models", get(models::list_models)) .route("/completions", post(chat::completions)) .route("/chat/completions", post(chat::chat_completions)) + .route("/models/{model}/load", post(models::load_model)) .layer(middleware::from_fn(auth_middleware)) } diff --git a/src/routes/v1/models.rs b/src/routes/v1/models.rs index d2c966b..fae7eff 100644 --- a/src/routes/v1/models.rs +++ b/src/routes/v1/models.rs @@ -1,16 +1,59 @@ -use axum::{Json, extract::State}; +use axum::{ + Json, + extract::{Path, State}, +}; +use serde::Deserialize; use serde_json::Value; -use crate::state::app_state::AppState; +use crate::{errors::OllamaError, state::app_state::AppState}; + +#[derive(Deserialize)] +pub struct LoadModelBody { + pub keep_alive: Option, +} pub async fn list_models( State(state): State, ) -> Result, (axum::http::StatusCode, String)> { match state.ollama.list_models().await { Ok(models) => Ok(Json(models)), - Err(err) => Err(( - axum::http::StatusCode::INTERNAL_SERVER_ERROR, - err.to_string(), - )), + Err(e) => Err(ollama_err(e)), + } +} + +pub async fn load_model( + State(state): State, + Path(model): Path, + Json(body): Json, +) -> Result, (axum::http::StatusCode, String)> { + match state + .ollama + .load_model(&model, body.keep_alive.as_deref()) + .await + { + Ok(response) => Ok(Json(response)), + Err(e) => Err(ollama_err(e)), + } +} + +fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) { + match e { + OllamaError::ModelNotFound(m) => ( + axum::http::StatusCode::NOT_FOUND, + format!("model '{m}' not found — run `ollama pull {m}`"), + ), + OllamaError::MissingPrompt => ( + axum::http::StatusCode::BAD_REQUEST, + "prompt is required".to_string(), + ), + OllamaError::MissingMessages => ( + axum::http::StatusCode::BAD_REQUEST, + "messages array with at least one user message is required".to_string(), + ), + OllamaError::InvalidKeepAlive(v) => ( + axum::http::StatusCode::BAD_REQUEST, + format!("invalid keep_alive '{v}' — use 30s / 10m / 2h, a plain integer, or -1"), + ), + OllamaError::Http(e) => (axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string()), } }