@@ -107,6 +107,31 @@ impl OllamaProvider {
|
||||
}))
|
||||
}
|
||||
|
||||
pub async fn unload_model(&self, model: &str) -> Result<Value, OllamaError> {
|
||||
self.validate_model(model).await?;
|
||||
|
||||
let payload = json!({
|
||||
"model": model,
|
||||
"prompt": "",
|
||||
"keep_alive": "0",
|
||||
"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": "unloaded",
|
||||
}))
|
||||
}
|
||||
|
||||
pub async fn completions(&self, body: Value) -> Result<Value, OllamaError> {
|
||||
let prompt = body
|
||||
.get("prompt")
|
||||
|
||||
@@ -11,5 +11,6 @@ pub fn router() -> Router<AppState> {
|
||||
.route("/completions", post(chat::completions))
|
||||
.route("/chat/completions", post(chat::chat_completions))
|
||||
.route("/models/{model}/load", post(models::load_model))
|
||||
.route("/models/{model}/unload", post(models::unload_model))
|
||||
.layer(middleware::from_fn(auth_middleware))
|
||||
}
|
||||
|
||||
@@ -36,6 +36,17 @@ pub async fn load_model(
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn unload_model(
|
||||
State(state): State<AppState>,
|
||||
Path(model): Path<String>,
|
||||
) -> Result<Json<Value>, (axum::http::StatusCode, String)> {
|
||||
match state.ollama.unload_model(&model).await {
|
||||
// ← correct method
|
||||
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) => (
|
||||
|
||||
@@ -228,3 +228,160 @@ async fn test_chat_completions_model_not_found() {
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, chat::errors::OllamaError::ModelNotFound(_)));
|
||||
}
|
||||
|
||||
// ── load_model ────────────────────────────────────────────────────────────────
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_load_model_ok() {
|
||||
let (server, provider) = setup().await;
|
||||
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/api/tags"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/api/generate"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
|
||||
"model": "llama3",
|
||||
"done": true,
|
||||
})))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let res = provider.load_model("llama3", Some("10m")).await.unwrap();
|
||||
|
||||
assert_eq!(res["model"], "llama3");
|
||||
assert_eq!(res["status"], "loaded");
|
||||
assert_eq!(res["keep_alive"], "10m");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_load_model_default_keep_alive() {
|
||||
let (server, provider) = setup().await;
|
||||
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/api/tags"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/api/generate"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
|
||||
"model": "llama3",
|
||||
"done": true,
|
||||
})))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let res = provider.load_model("llama3", None).await.unwrap();
|
||||
|
||||
assert_eq!(res["status"], "loaded");
|
||||
assert_eq!(res["keep_alive"], "5m"); // default
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_load_model_not_found() {
|
||||
let (server, provider) = setup().await;
|
||||
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/api/tags"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let err = provider.load_model("gpt-4", Some("10m")).await.unwrap_err();
|
||||
assert!(matches!(err, chat::errors::OllamaError::ModelNotFound(_)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_load_model_invalid_keep_alive() {
|
||||
let (server, provider) = setup().await;
|
||||
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/api/tags"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let err = provider
|
||||
.load_model("llama3", Some("10x"))
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(
|
||||
err,
|
||||
chat::errors::OllamaError::InvalidKeepAlive(_)
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_load_model_keep_alive_plain_integer() {
|
||||
let (server, provider) = setup().await;
|
||||
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/api/tags"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/api/generate"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
|
||||
"model": "llama3",
|
||||
"done": true,
|
||||
})))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
// plain integers (seconds) and "-1" are valid
|
||||
let res = provider.load_model("llama3", Some("3600")).await.unwrap();
|
||||
assert_eq!(res["status"], "loaded");
|
||||
|
||||
let res = provider.load_model("llama3", Some("-1")).await.unwrap();
|
||||
assert_eq!(res["status"], "loaded");
|
||||
}
|
||||
|
||||
// ── unload_model ──────────────────────────────────────────────────────────────
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_unload_model_ok() {
|
||||
let (server, provider) = setup().await;
|
||||
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/api/tags"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/api/generate"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
|
||||
"model": "llama3",
|
||||
"done": true,
|
||||
})))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let res = provider.unload_model("llama3").await.unwrap();
|
||||
|
||||
assert_eq!(res["model"], "llama3");
|
||||
assert_eq!(res["status"], "unloaded");
|
||||
// no keep_alive field on unload response
|
||||
assert!(res.get("keep_alive").is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_unload_model_not_found() {
|
||||
let (server, provider) = setup().await;
|
||||
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/api/tags"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
|
||||
.mount(&server)
|
||||
.await;
|
||||
|
||||
let err = provider.unload_model("gpt-4").await.unwrap_err();
|
||||
assert!(matches!(err, chat::errors::OllamaError::ModelNotFound(_)));
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user