feat: update test
CI / Rust CI (push) Successful in 4m35s
Publish & Deploy / Build and Push to Registry (push) Successful in 1m30s
Publish & Deploy / Deploy via SSH (push) Successful in 18s

This commit is contained in:
2026-04-10 22:26:45 +02:00
parent e31cdf131f
commit fc391b5d0e
6 changed files with 187 additions and 145 deletions
+134 -107
View File
@@ -2,7 +2,9 @@ use serde_json::json;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
use chat::providers::ollama::OllamaProvider;
use chat::dto::api;
use chat::errors::OllamaError;
use chat::providers::ollama::client::OllamaProvider;
// ── helpers ──────────────────────────────────────────────────────────────────
@@ -31,7 +33,7 @@ async fn test_list_models_ok() {
.await;
let res = provider.list_models().await.unwrap();
assert_eq!(res["models"][0]["name"], "llama3");
assert_eq!(res.models[0].name, "llama3");
}
// ── completions ───────────────────────────────────────────────────────────────
@@ -58,19 +60,25 @@ async fn test_completions_ok() {
.mount(&server)
.await;
let res = provider
.completions(json!({
"model": "llama3",
"prompt": "Who are you?"
}))
.await
.unwrap();
let req = api::CompletionRequest {
base: api::BaseLLMRequest {
model: "llama3".to_string(),
..Default::default()
},
prompt: "Hello".to_string(),
};
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);
let res = provider.completions(&req).await.unwrap();
assert_eq!(res.object, api::CompletionObject::TextCompletion);
assert_eq!(res.choices.len(), 1);
assert_eq!(res.choices[0].text, "I am a helpful assistant.");
assert_eq!(res.choices[0].finish_reason, api::FinishReason::Stop);
assert_eq!(res.usage.prompt_tokens, 10);
assert_eq!(res.usage.completion_tokens, 8);
assert_eq!(res.usage.total_tokens, 18);
}
#[tokio::test]
@@ -83,11 +91,17 @@ async fn test_completions_missing_prompt() {
.mount(&server)
.await;
let err = provider
.completions(json!({ "model": "llama3" }))
.await
.unwrap_err();
assert!(matches!(err, chat::errors::OllamaError::MissingPrompt));
let req = api::CompletionRequest {
base: api::BaseLLMRequest {
model: "llama3".to_string(),
..Default::default()
},
prompt: "".to_string(),
};
let err = provider.completions(&req).await.unwrap_err();
assert!(matches!(err, OllamaError::MissingPrompt));
}
#[tokio::test]
@@ -100,13 +114,17 @@ async fn test_completions_empty_prompt() {
.mount(&server)
.await;
let err = provider
.completions(json!({
"model": "llama3", "prompt": " "
}))
.await
.unwrap_err();
assert!(matches!(err, chat::errors::OllamaError::MissingPrompt));
let req = api::CompletionRequest {
base: api::BaseLLMRequest {
model: "llama3".to_string(),
..Default::default()
},
prompt: " ".to_string(),
};
let err = provider.completions(&req).await.unwrap_err();
assert!(matches!(err, OllamaError::MissingPrompt));
}
#[tokio::test]
@@ -119,12 +137,16 @@ async fn test_completions_model_not_found() {
.mount(&server)
.await;
let err = provider
.completions(json!({
"model": "gpt-4", "prompt": "hello"
}))
.await
.unwrap_err();
let req = api::CompletionRequest {
base: api::BaseLLMRequest {
model: "gpt-4".to_string(),
..Default::default()
},
prompt: "hello".to_string(),
};
let err = provider.completions(&req).await.unwrap_err();
assert!(matches!(err, chat::errors::OllamaError::ModelNotFound(_)));
}
@@ -146,28 +168,37 @@ async fn test_chat_completions_ok() {
"model": "llama3",
"message": { "role": "assistant", "content": "4." },
"done": true,
"done_reason": "stop",
"prompt_eval_count": 5,
"eval_count": 2,
"eval_count": 2
})))
.mount(&server)
.await;
let res = provider
.chat_completions(json!({
"model": "llama3",
"messages": [
{ "role": "user", "content": "What is 2+2?" }
]
}))
.await
.unwrap();
let req = api::ChatRequest {
base: api::BaseLLMRequest {
model: "llama3".to_string(),
..Default::default()
},
messages: vec![api::Message {
role: api::Role::User,
content: "What is 2+2?".to_string(),
}],
};
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);
let res = provider.chat_completions(&req).await.unwrap();
assert_eq!(res.object, "chat.completion");
assert_eq!(res.choices.len(), 1);
assert_eq!(res.choices[0].message.role, api::Role::Assistant);
assert_eq!(res.choices[0].message.content, "4.");
assert_eq!(res.choices[0].finish_reason, api::FinishReason::Stop);
let usage = res.usage.unwrap();
assert_eq!(usage.prompt_tokens, 5);
assert_eq!(usage.completion_tokens, 2);
assert_eq!(usage.total_tokens, 7);
}
#[tokio::test]
@@ -180,12 +211,16 @@ async fn test_chat_completions_missing_messages() {
.mount(&server)
.await;
let err = provider
.chat_completions(json!({
"model": "llama3", "messages": []
}))
.await
.unwrap_err();
let req = api::ChatRequest {
base: api::BaseLLMRequest {
model: "llama3".to_string(),
..Default::default()
},
messages: vec![],
};
let err = provider.chat_completions(&req).await.unwrap_err();
assert!(matches!(err, chat::errors::OllamaError::MissingMessages));
}
@@ -199,13 +234,19 @@ async fn test_chat_completions_no_user_message() {
.mount(&server)
.await;
let err = provider
.chat_completions(json!({
"model": "llama3",
"messages": [{ "role": "system", "content": "be helpful" }]
}))
.await
.unwrap_err();
let req = api::ChatRequest {
base: api::BaseLLMRequest {
model: "llama3".to_string(),
..Default::default()
},
messages: vec![api::Message {
role: api::Role::System,
content: "be helpful".to_string(),
}],
};
let err = provider.chat_completions(&req).await.unwrap_err();
assert!(matches!(err, chat::errors::OllamaError::MissingMessages));
}
@@ -219,17 +260,23 @@ async fn test_chat_completions_model_not_found() {
.mount(&server)
.await;
let err = provider
.chat_completions(json!({
"model": "gpt-4",
"messages": [{ "role": "user", "content": "hi" }]
}))
.await
.unwrap_err();
let req = api::ChatRequest {
base: api::BaseLLMRequest {
model: "gpt-4".to_string(),
..Default::default()
},
messages: vec![api::Message {
role: api::Role::User,
content: "hi".to_string(),
}],
};
let err = provider.chat_completions(&req).await.unwrap_err();
assert!(matches!(err, chat::errors::OllamaError::ModelNotFound(_)));
}
// ── load_model ────────────────────────────────────────────────────────────────
// // ── load_model ────────────────────────────────────────────────────────────────
#[tokio::test]
async fn test_load_model_ok() {
@@ -245,41 +292,17 @@ async fn test_load_model_ok() {
.and(path("/api/generate"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"model": "llama3",
"done": true,
"response": "ok",
"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
assert_eq!(res.model, "llama3");
assert_eq!(res.status, "loaded");
assert_eq!(res.keep_alive, "10m");
}
#[tokio::test]
@@ -293,6 +316,7 @@ async fn test_load_model_not_found() {
.await;
let err = provider.load_model("gpt-4", Some("10m")).await.unwrap_err();
assert!(matches!(err, chat::errors::OllamaError::ModelNotFound(_)));
}
@@ -310,6 +334,7 @@ async fn test_load_model_invalid_keep_alive() {
.load_model("llama3", Some("10x"))
.await
.unwrap_err();
assert!(matches!(
err,
chat::errors::OllamaError::InvalidKeepAlive(_)
@@ -330,20 +355,23 @@ async fn test_load_model_keep_alive_plain_integer() {
.and(path("/api/generate"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"model": "llama3",
"done": true,
"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");
assert_eq!(res.status, "loaded");
assert_eq!(res.keep_alive, "3600");
let res = provider.load_model("llama3", Some("-1")).await.unwrap();
assert_eq!(res["status"], "loaded");
assert_eq!(res.status, "loaded");
assert_eq!(res.keep_alive, "-1");
}
// ── unload_model ──────────────────────────────────────────────────────────────
// // ── unload_model ──────────────────────────────────────────────────────────────
#[tokio::test]
async fn test_unload_model_ok() {
@@ -359,6 +387,7 @@ async fn test_unload_model_ok() {
.and(path("/api/generate"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"model": "llama3",
"response": "ok",
"done": true,
})))
.mount(&server)
@@ -366,10 +395,8 @@ async fn test_unload_model_ok() {
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());
assert_eq!(res.model, "llama3");
assert_eq!(res.status, "unloaded");
}
#[tokio::test]