feat: add chat completions route + test
CI / Rust CI (push) Successful in 4m23s
Publish & Deploy / Build and Push to Registry (push) Successful in 1m26s
Publish & Deploy / Deploy via SSH (push) Successful in 18s

This commit is contained in:
2026-04-10 11:47:56 +02:00
parent 70fe9bc4da
commit 6dc7231309
10 changed files with 521 additions and 429 deletions
+230
View File
@@ -0,0 +1,230 @@
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(_)));
}