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
+4 -30
View File
@@ -1,8 +1,6 @@
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use utoipa::ToSchema; use utoipa::ToSchema;
use crate::errors::OllamaError;
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ModelsResponse { pub struct ModelsResponse {
pub models: Vec<ModelInfo>, pub models: Vec<ModelInfo>,
@@ -28,37 +26,13 @@ pub struct LoadModelBody {
pub keep_alive: Option<String>, pub keep_alive: Option<String>,
} }
impl LoadModelBody {
pub fn parse_keep_alive(s: &str) -> Result<(), OllamaError> {
let s = s.trim();
if s == "-1" || s.parse::<u64>().is_ok() {
return Ok(());
}
let split = s
.find(|c: char| c.is_alphabetic())
.ok_or_else(|| OllamaError::InvalidKeepAlive(s.to_string()))?;
let (num, unit) = s.split_at(split);
num.parse::<u64>()
.map_err(|_| OllamaError::InvalidKeepAlive(s.to_string()))?;
match unit {
"s" | "m" | "h" => Ok(()),
_ => Err(OllamaError::InvalidKeepAlive(s.to_string())),
}
}
}
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct UnloadModelResponse { pub struct UnloadModelResponse {
pub model: String, pub model: String,
pub status: String, pub status: String,
} }
#[derive(Debug, Deserialize, Serialize, ToSchema)] #[derive(Debug, Deserialize, Serialize, ToSchema, Default)]
pub struct BaseLLMRequest { pub struct BaseLLMRequest {
pub model: String, pub model: String,
@@ -89,13 +63,13 @@ pub struct CompletionRequest {
pub prompt: String, pub prompt: String,
} }
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema, PartialEq, Eq)]
#[serde(rename_all = "snake_case")] #[serde(rename_all = "snake_case")]
pub enum CompletionObject { pub enum CompletionObject {
TextCompletion, TextCompletion,
} }
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema, PartialEq, Eq)]
#[serde(rename_all = "snake_case")] #[serde(rename_all = "snake_case")]
pub enum FinishReason { pub enum FinishReason {
Stop, Stop,
@@ -149,7 +123,7 @@ pub struct Message {
pub content: String, pub content: String,
} }
#[derive(Debug, Deserialize, Serialize, ToSchema)] #[derive(Debug, Deserialize, Serialize, ToSchema, PartialEq, Eq)]
#[serde(rename_all = "lowercase")] #[serde(rename_all = "lowercase")]
pub enum Role { pub enum Role {
System, System,
+3
View File
@@ -15,6 +15,9 @@ pub enum OllamaError {
#[error("model '{0}' is not available — run `ollama pull {0}` first")] #[error("model '{0}' is not available — run `ollama pull {0}` first")]
ModelNotFound(String), ModelNotFound(String),
#[error("keep_alive is required and cannot be empty")]
MissingKeepAlive,
#[error( #[error(
"invalid keep_alive format '{0}' — expected <number><unit> (e.g. 30s, 10m, 2h), a plain integer (seconds), or -1" "invalid keep_alive format '{0}' — expected <number><unit> (e.g. 30s, 10m, 2h), a plain integer (seconds), or -1"
)] )]
+37 -3
View File
@@ -36,6 +36,10 @@ impl OllamaProvider {
Ok(res.models.iter().any(|m| m.name == model)) Ok(res.models.iter().any(|m| m.name == model))
} }
fn has_user_message(&self, messages: &[api::Message]) -> bool {
messages.iter().any(|m| matches!(m.role, api::Role::User))
}
fn extract_completion_params<'a>( fn extract_completion_params<'a>(
&self, &self,
body: &'a api::CompletionRequest, body: &'a api::CompletionRequest,
@@ -56,6 +60,28 @@ impl OllamaProvider {
Ok((&body.messages, model)) Ok((&body.messages, model))
} }
pub fn parse_keep_alive(&self, s: &str) -> Result<(), OllamaError> {
let s = s.trim();
if s == "-1" || s.parse::<u64>().is_ok() {
return Ok(());
}
let split = s
.find(|c: char| c.is_alphabetic())
.ok_or_else(|| OllamaError::InvalidKeepAlive(s.to_string()))?;
let (num, unit) = s.split_at(split);
num.parse::<u64>()
.map_err(|_| OllamaError::InvalidKeepAlive(s.to_string()))?;
match unit {
"s" | "m" | "h" => Ok(()),
_ => Err(OllamaError::InvalidKeepAlive(s.to_string())),
}
}
// // ── public endpoints ───────────────────────────────────────────────────── // // ── public endpoints ─────────────────────────────────────────────────────
pub async fn list_models(&self) -> Result<api::ModelsResponse, OllamaError> { pub async fn list_models(&self) -> Result<api::ModelsResponse, OllamaError> {
@@ -77,10 +103,14 @@ impl OllamaProvider {
pub async fn load_model( pub async fn load_model(
&self, &self,
model: &str, model: &str,
keep_alive: &str, keep_alive: Option<&str>,
) -> Result<api::LoadModelResponse, OllamaError> { ) -> Result<api::LoadModelResponse, OllamaError> {
let url = format!("{}/api/generate", self.base_url); let url = format!("{}/api/generate", self.base_url);
let keep_alive = keep_alive.ok_or(OllamaError::MissingKeepAlive)?;
self.parse_keep_alive(keep_alive)?;
let exists = self.model_exists(model).await?; let exists = self.model_exists(model).await?;
if !exists { if !exists {
return Err(OllamaError::ModelNotFound(model.to_string())); return Err(OllamaError::ModelNotFound(model.to_string()));
@@ -99,7 +129,7 @@ impl OllamaProvider {
.json(&payload) .json(&payload)
.send() .send()
.await? .await?
.json::<ollama::OllamaGenerateResponse>() .text()
.await?; .await?;
Ok(api::LoadModelResponse { Ok(api::LoadModelResponse {
@@ -130,7 +160,7 @@ impl OllamaProvider {
.json(&payload) .json(&payload)
.send() .send()
.await? .await?
.json::<ollama::OllamaGenerateResponse>() .text()
.await?; .await?;
Ok(api::UnloadModelResponse { Ok(api::UnloadModelResponse {
@@ -277,6 +307,10 @@ impl OllamaProvider {
return Err(OllamaError::MissingMessages); return Err(OllamaError::MissingMessages);
} }
if !self.has_user_message(&body.messages) {
return Err(OllamaError::MissingMessages);
}
if model.is_empty() { if model.is_empty() {
return Err(OllamaError::MissingModel); return Err(OllamaError::MissingModel);
} }
+4
View File
@@ -78,6 +78,10 @@ fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) {
axum::http::StatusCode::UNPROCESSABLE_ENTITY, axum::http::StatusCode::UNPROCESSABLE_ENTITY,
format!("model '{m}' is not available — run `ollama pull {m}` first"), format!("model '{m}' is not available — run `ollama pull {m}` first"),
), ),
OllamaError::MissingKeepAlive => (
axum::http::StatusCode::BAD_REQUEST,
"keep alive is required and cannot be empty".to_string(),
),
OllamaError::InvalidKeepAlive(v) => ( OllamaError::InvalidKeepAlive(v) => (
axum::http::StatusCode::BAD_REQUEST, axum::http::StatusCode::BAD_REQUEST,
format!("invalid keep_alive '{v}'"), format!("invalid keep_alive '{v}'"),
+5 -5
View File
@@ -20,13 +20,9 @@ pub async fn load_model(
Path(model): Path<String>, Path(model): Path<String>,
Json(body): Json<api::LoadModelBody>, Json(body): Json<api::LoadModelBody>,
) -> Result<Json<api::LoadModelResponse>, (axum::http::StatusCode, String)> { ) -> Result<Json<api::LoadModelResponse>, (axum::http::StatusCode, String)> {
let keep_alive = body.keep_alive.as_deref().unwrap_or("5m");
api::LoadModelBody::parse_keep_alive(keep_alive).map_err(ollama_err)?;
let response = state let response = state
.ollama .ollama
.load_model(&model, keep_alive) .load_model(&model, body.keep_alive.as_deref())
.await .await
.map_err(ollama_err)?; .map_err(ollama_err)?;
@@ -64,6 +60,10 @@ fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) {
axum::http::StatusCode::BAD_REQUEST, axum::http::StatusCode::BAD_REQUEST,
"messages array with at least one user message is required".to_string(), "messages array with at least one user message is required".to_string(),
), ),
OllamaError::MissingKeepAlive => (
axum::http::StatusCode::BAD_REQUEST,
"keep alive is required and cannot be empty".to_string(),
),
OllamaError::InvalidKeepAlive(v) => ( OllamaError::InvalidKeepAlive(v) => (
axum::http::StatusCode::BAD_REQUEST, axum::http::StatusCode::BAD_REQUEST,
format!("invalid keep_alive '{v}' — use 30s / 10m / 2h, a plain integer, or -1"), format!("invalid keep_alive '{v}' — use 30s / 10m / 2h, a plain integer, or -1"),
+134 -107
View File
@@ -2,7 +2,9 @@ use serde_json::json;
use wiremock::matchers::{method, path}; use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate}; 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 ────────────────────────────────────────────────────────────────── // ── helpers ──────────────────────────────────────────────────────────────────
@@ -31,7 +33,7 @@ async fn test_list_models_ok() {
.await; .await;
let res = provider.list_models().await.unwrap(); let res = provider.list_models().await.unwrap();
assert_eq!(res["models"][0]["name"], "llama3"); assert_eq!(res.models[0].name, "llama3");
} }
// ── completions ─────────────────────────────────────────────────────────────── // ── completions ───────────────────────────────────────────────────────────────
@@ -58,19 +60,25 @@ async fn test_completions_ok() {
.mount(&server) .mount(&server)
.await; .await;
let res = provider let req = api::CompletionRequest {
.completions(json!({ base: api::BaseLLMRequest {
"model": "llama3", model: "llama3".to_string(),
"prompt": "Who are you?" ..Default::default()
})) },
.await prompt: "Hello".to_string(),
.unwrap(); };
assert_eq!(res["object"], "text_completion"); let res = provider.completions(&req).await.unwrap();
assert_eq!(res["choices"][0]["text"], "I am a helpful assistant.");
assert_eq!(res["choices"][0]["finish_reason"], "stop"); assert_eq!(res.object, api::CompletionObject::TextCompletion);
assert_eq!(res["usage"]["prompt_tokens"], 10);
assert_eq!(res["usage"]["completion_tokens"], 8); 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] #[tokio::test]
@@ -83,11 +91,17 @@ async fn test_completions_missing_prompt() {
.mount(&server) .mount(&server)
.await; .await;
let err = provider let req = api::CompletionRequest {
.completions(json!({ "model": "llama3" })) base: api::BaseLLMRequest {
.await model: "llama3".to_string(),
.unwrap_err(); ..Default::default()
assert!(matches!(err, chat::errors::OllamaError::MissingPrompt)); },
prompt: "".to_string(),
};
let err = provider.completions(&req).await.unwrap_err();
assert!(matches!(err, OllamaError::MissingPrompt));
} }
#[tokio::test] #[tokio::test]
@@ -100,13 +114,17 @@ async fn test_completions_empty_prompt() {
.mount(&server) .mount(&server)
.await; .await;
let err = provider let req = api::CompletionRequest {
.completions(json!({ base: api::BaseLLMRequest {
"model": "llama3", "prompt": " " model: "llama3".to_string(),
})) ..Default::default()
.await },
.unwrap_err(); prompt: " ".to_string(),
assert!(matches!(err, chat::errors::OllamaError::MissingPrompt)); };
let err = provider.completions(&req).await.unwrap_err();
assert!(matches!(err, OllamaError::MissingPrompt));
} }
#[tokio::test] #[tokio::test]
@@ -119,12 +137,16 @@ async fn test_completions_model_not_found() {
.mount(&server) .mount(&server)
.await; .await;
let err = provider let req = api::CompletionRequest {
.completions(json!({ base: api::BaseLLMRequest {
"model": "gpt-4", "prompt": "hello" model: "gpt-4".to_string(),
})) ..Default::default()
.await },
.unwrap_err(); prompt: "hello".to_string(),
};
let err = provider.completions(&req).await.unwrap_err();
assert!(matches!(err, chat::errors::OllamaError::ModelNotFound(_))); assert!(matches!(err, chat::errors::OllamaError::ModelNotFound(_)));
} }
@@ -146,28 +168,37 @@ async fn test_chat_completions_ok() {
"model": "llama3", "model": "llama3",
"message": { "role": "assistant", "content": "4." }, "message": { "role": "assistant", "content": "4." },
"done": true, "done": true,
"done_reason": "stop",
"prompt_eval_count": 5, "prompt_eval_count": 5,
"eval_count": 2, "eval_count": 2
}))) })))
.mount(&server) .mount(&server)
.await; .await;
let res = provider let req = api::ChatRequest {
.chat_completions(json!({ base: api::BaseLLMRequest {
"model": "llama3", model: "llama3".to_string(),
"messages": [ ..Default::default()
{ "role": "user", "content": "What is 2+2?" } },
] messages: vec![api::Message {
})) role: api::Role::User,
.await content: "What is 2+2?".to_string(),
.unwrap(); }],
};
assert_eq!(res["object"], "chat.completion"); let res = provider.chat_completions(&req).await.unwrap();
assert_eq!(res["choices"][0]["message"]["role"], "assistant");
assert_eq!(res["choices"][0]["message"]["content"], "4."); assert_eq!(res.object, "chat.completion");
assert_eq!(res["choices"][0]["finish_reason"], "stop"); assert_eq!(res.choices.len(), 1);
assert_eq!(res["usage"]["total_tokens"], 7);
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] #[tokio::test]
@@ -180,12 +211,16 @@ async fn test_chat_completions_missing_messages() {
.mount(&server) .mount(&server)
.await; .await;
let err = provider let req = api::ChatRequest {
.chat_completions(json!({ base: api::BaseLLMRequest {
"model": "llama3", "messages": [] model: "llama3".to_string(),
})) ..Default::default()
.await },
.unwrap_err(); messages: vec![],
};
let err = provider.chat_completions(&req).await.unwrap_err();
assert!(matches!(err, chat::errors::OllamaError::MissingMessages)); assert!(matches!(err, chat::errors::OllamaError::MissingMessages));
} }
@@ -199,13 +234,19 @@ async fn test_chat_completions_no_user_message() {
.mount(&server) .mount(&server)
.await; .await;
let err = provider let req = api::ChatRequest {
.chat_completions(json!({ base: api::BaseLLMRequest {
"model": "llama3", model: "llama3".to_string(),
"messages": [{ "role": "system", "content": "be helpful" }] ..Default::default()
})) },
.await messages: vec![api::Message {
.unwrap_err(); 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)); assert!(matches!(err, chat::errors::OllamaError::MissingMessages));
} }
@@ -219,17 +260,23 @@ async fn test_chat_completions_model_not_found() {
.mount(&server) .mount(&server)
.await; .await;
let err = provider let req = api::ChatRequest {
.chat_completions(json!({ base: api::BaseLLMRequest {
"model": "gpt-4", model: "gpt-4".to_string(),
"messages": [{ "role": "user", "content": "hi" }] ..Default::default()
})) },
.await messages: vec![api::Message {
.unwrap_err(); role: api::Role::User,
content: "hi".to_string(),
}],
};
let err = provider.chat_completions(&req).await.unwrap_err();
assert!(matches!(err, chat::errors::OllamaError::ModelNotFound(_))); assert!(matches!(err, chat::errors::OllamaError::ModelNotFound(_)));
} }
// ── load_model ──────────────────────────────────────────────────────────────── // // ── load_model ────────────────────────────────────────────────────────────────
#[tokio::test] #[tokio::test]
async fn test_load_model_ok() { async fn test_load_model_ok() {
@@ -245,41 +292,17 @@ async fn test_load_model_ok() {
.and(path("/api/generate")) .and(path("/api/generate"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({ .respond_with(ResponseTemplate::new(200).set_body_json(json!({
"model": "llama3", "model": "llama3",
"done": true, "response": "ok",
"done": true,
}))) })))
.mount(&server) .mount(&server)
.await; .await;
let res = provider.load_model("llama3", Some("10m")).await.unwrap(); let res = provider.load_model("llama3", Some("10m")).await.unwrap();
assert_eq!(res["model"], "llama3"); assert_eq!(res.model, "llama3");
assert_eq!(res["status"], "loaded"); assert_eq!(res.status, "loaded");
assert_eq!(res["keep_alive"], "10m"); 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] #[tokio::test]
@@ -293,6 +316,7 @@ async fn test_load_model_not_found() {
.await; .await;
let err = provider.load_model("gpt-4", Some("10m")).await.unwrap_err(); let err = provider.load_model("gpt-4", Some("10m")).await.unwrap_err();
assert!(matches!(err, chat::errors::OllamaError::ModelNotFound(_))); assert!(matches!(err, chat::errors::OllamaError::ModelNotFound(_)));
} }
@@ -310,6 +334,7 @@ async fn test_load_model_invalid_keep_alive() {
.load_model("llama3", Some("10x")) .load_model("llama3", Some("10x"))
.await .await
.unwrap_err(); .unwrap_err();
assert!(matches!( assert!(matches!(
err, err,
chat::errors::OllamaError::InvalidKeepAlive(_) chat::errors::OllamaError::InvalidKeepAlive(_)
@@ -330,20 +355,23 @@ async fn test_load_model_keep_alive_plain_integer() {
.and(path("/api/generate")) .and(path("/api/generate"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({ .respond_with(ResponseTemplate::new(200).set_body_json(json!({
"model": "llama3", "model": "llama3",
"done": true, "done": true,
}))) })))
.mount(&server) .mount(&server)
.await; .await;
// plain integers (seconds) and "-1" are valid
let res = provider.load_model("llama3", Some("3600")).await.unwrap(); 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(); 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] #[tokio::test]
async fn test_unload_model_ok() { async fn test_unload_model_ok() {
@@ -359,6 +387,7 @@ async fn test_unload_model_ok() {
.and(path("/api/generate")) .and(path("/api/generate"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({ .respond_with(ResponseTemplate::new(200).set_body_json(json!({
"model": "llama3", "model": "llama3",
"response": "ok",
"done": true, "done": true,
}))) })))
.mount(&server) .mount(&server)
@@ -366,10 +395,8 @@ async fn test_unload_model_ok() {
let res = provider.unload_model("llama3").await.unwrap(); let res = provider.unload_model("llama3").await.unwrap();
assert_eq!(res["model"], "llama3"); assert_eq!(res.model, "llama3");
assert_eq!(res["status"], "unloaded"); assert_eq!(res.status, "unloaded");
// no keep_alive field on unload response
assert!(res.get("keep_alive").is_none());
} }
#[tokio::test] #[tokio::test]