diff --git a/src/providers/ollama.rs b/src/providers/ollama.rs index 979cd2f..5c7cffe 100644 --- a/src/providers/ollama.rs +++ b/src/providers/ollama.rs @@ -107,6 +107,31 @@ impl OllamaProvider { })) } + pub async fn unload_model(&self, model: &str) -> Result { + 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::() + .await?; + + Ok(json!({ + "model": res.get("model"), + "status": "unloaded", + })) + } + 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 9fdecff..30059e7 100644 --- a/src/routes/v1/mod.rs +++ b/src/routes/v1/mod.rs @@ -11,5 +11,6 @@ pub fn router() -> Router { .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)) } diff --git a/src/routes/v1/models.rs b/src/routes/v1/models.rs index fae7eff..2f9d629 100644 --- a/src/routes/v1/models.rs +++ b/src/routes/v1/models.rs @@ -36,6 +36,17 @@ pub async fn load_model( } } +pub async fn unload_model( + State(state): State, + Path(model): Path, +) -> Result, (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) => ( diff --git a/tests/ollama_provider.rs b/tests/ollama_provider.rs index fa6f7b0..500607d 100644 --- a/tests/ollama_provider.rs +++ b/tests/ollama_provider.rs @@ -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(_))); +}