feat: add load endpoint

This commit is contained in:
2026-04-10 12:29:34 +02:00
parent db14d816cb
commit 1a06577fa3
4 changed files with 114 additions and 6 deletions
+1
View File
@@ -10,5 +10,6 @@ pub fn router() -> Router<AppState> {
.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))
}
+49 -6
View File
@@ -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<String>,
}
pub async fn list_models(
State(state): State<AppState>,
) -> Result<Json<Value>, (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<AppState>,
Path(model): Path<String>,
Json(body): Json<LoadModelBody>,
) -> Result<Json<Value>, (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()),
}
}