388 lines
12 KiB
Rust
388 lines
12 KiB
Rust
use serde_json::json;
|
|
use wiremock::matchers::{method, path};
|
|
use wiremock::{Mock, MockServer, ResponseTemplate};
|
|
|
|
use chat::providers::ollama::OllamaProvider;
|
|
|
|
// ── helpers ──────────────────────────────────────────────────────────────────
|
|
|
|
async fn setup() -> (MockServer, OllamaProvider) {
|
|
let server = MockServer::start().await;
|
|
let provider = OllamaProvider::new(server.uri());
|
|
(server, provider)
|
|
}
|
|
|
|
fn models_response(names: &[&str]) -> serde_json::Value {
|
|
json!({
|
|
"models": names.iter().map(|n| json!({ "name": n })).collect::<Vec<_>>()
|
|
})
|
|
}
|
|
|
|
// ── list_models ───────────────────────────────────────────────────────────────
|
|
|
|
#[tokio::test]
|
|
async fn test_list_models_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;
|
|
|
|
let res = provider.list_models().await.unwrap();
|
|
assert_eq!(res["models"][0]["name"], "llama3");
|
|
}
|
|
|
|
// ── completions ───────────────────────────────────────────────────────────────
|
|
|
|
#[tokio::test]
|
|
async fn test_completions_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",
|
|
"response": "I am a helpful assistant.",
|
|
"done": true,
|
|
"prompt_eval_count": 10,
|
|
"eval_count": 8,
|
|
})))
|
|
.mount(&server)
|
|
.await;
|
|
|
|
let res = provider
|
|
.completions(json!({
|
|
"model": "llama3",
|
|
"prompt": "Who are you?"
|
|
}))
|
|
.await
|
|
.unwrap();
|
|
|
|
assert_eq!(res["object"], "text_completion");
|
|
assert_eq!(res["choices"][0]["text"], "I am a helpful assistant.");
|
|
assert_eq!(res["choices"][0]["finish_reason"], "stop");
|
|
assert_eq!(res["usage"]["prompt_tokens"], 10);
|
|
assert_eq!(res["usage"]["completion_tokens"], 8);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_completions_missing_prompt() {
|
|
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
|
|
.completions(json!({ "model": "llama3" }))
|
|
.await
|
|
.unwrap_err();
|
|
assert!(matches!(err, chat::errors::OllamaError::MissingPrompt));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_completions_empty_prompt() {
|
|
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
|
|
.completions(json!({
|
|
"model": "llama3", "prompt": " "
|
|
}))
|
|
.await
|
|
.unwrap_err();
|
|
assert!(matches!(err, chat::errors::OllamaError::MissingPrompt));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_completions_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
|
|
.completions(json!({
|
|
"model": "gpt-4", "prompt": "hello"
|
|
}))
|
|
.await
|
|
.unwrap_err();
|
|
assert!(matches!(err, chat::errors::OllamaError::ModelNotFound(_)));
|
|
}
|
|
|
|
// ── chat_completions ──────────────────────────────────────────────────────────
|
|
|
|
#[tokio::test]
|
|
async fn test_chat_completions_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/chat"))
|
|
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
|
|
"model": "llama3",
|
|
"message": { "role": "assistant", "content": "4." },
|
|
"done": true,
|
|
"done_reason": "stop",
|
|
"prompt_eval_count": 5,
|
|
"eval_count": 2,
|
|
})))
|
|
.mount(&server)
|
|
.await;
|
|
|
|
let res = provider
|
|
.chat_completions(json!({
|
|
"model": "llama3",
|
|
"messages": [
|
|
{ "role": "user", "content": "What is 2+2?" }
|
|
]
|
|
}))
|
|
.await
|
|
.unwrap();
|
|
|
|
assert_eq!(res["object"], "chat.completion");
|
|
assert_eq!(res["choices"][0]["message"]["role"], "assistant");
|
|
assert_eq!(res["choices"][0]["message"]["content"], "4.");
|
|
assert_eq!(res["choices"][0]["finish_reason"], "stop");
|
|
assert_eq!(res["usage"]["total_tokens"], 7);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_chat_completions_missing_messages() {
|
|
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
|
|
.chat_completions(json!({
|
|
"model": "llama3", "messages": []
|
|
}))
|
|
.await
|
|
.unwrap_err();
|
|
assert!(matches!(err, chat::errors::OllamaError::MissingMessages));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_chat_completions_no_user_message() {
|
|
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
|
|
.chat_completions(json!({
|
|
"model": "llama3",
|
|
"messages": [{ "role": "system", "content": "be helpful" }]
|
|
}))
|
|
.await
|
|
.unwrap_err();
|
|
assert!(matches!(err, chat::errors::OllamaError::MissingMessages));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_chat_completions_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
|
|
.chat_completions(json!({
|
|
"model": "gpt-4",
|
|
"messages": [{ "role": "user", "content": "hi" }]
|
|
}))
|
|
.await
|
|
.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(_)));
|
|
}
|