@@ -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