From fc391b5d0eea1edb0bef22d29573b54d1c4145c3 Mon Sep 17 00:00:00 2001 From: LucasX Ubuntu Date: Fri, 10 Apr 2026 22:26:45 +0200 Subject: [PATCH] feat: update test --- src/dto/api.rs | 34 +---- src/errors.rs | 3 + src/providers/ollama/client.rs | 40 +++++- src/routes/v1/chat.rs | 4 + src/routes/v1/models.rs | 10 +- tests/ollama_provider.rs | 241 ++++++++++++++++++--------------- 6 files changed, 187 insertions(+), 145 deletions(-) diff --git a/src/dto/api.rs b/src/dto/api.rs index 06b0bde..5a8a1b6 100644 --- a/src/dto/api.rs +++ b/src/dto/api.rs @@ -1,8 +1,6 @@ use serde::{Deserialize, Serialize}; use utoipa::ToSchema; -use crate::errors::OllamaError; - #[derive(Debug, Serialize, Deserialize, ToSchema)] pub struct ModelsResponse { pub models: Vec, @@ -28,37 +26,13 @@ pub struct LoadModelBody { pub keep_alive: Option, } -impl LoadModelBody { - pub fn parse_keep_alive(s: &str) -> Result<(), OllamaError> { - let s = s.trim(); - - if s == "-1" || s.parse::().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::() - .map_err(|_| OllamaError::InvalidKeepAlive(s.to_string()))?; - - match unit { - "s" | "m" | "h" => Ok(()), - _ => Err(OllamaError::InvalidKeepAlive(s.to_string())), - } - } -} - #[derive(Debug, Serialize, Deserialize, ToSchema)] pub struct UnloadModelResponse { pub model: String, pub status: String, } -#[derive(Debug, Deserialize, Serialize, ToSchema)] +#[derive(Debug, Deserialize, Serialize, ToSchema, Default)] pub struct BaseLLMRequest { pub model: String, @@ -89,13 +63,13 @@ pub struct CompletionRequest { pub prompt: String, } -#[derive(Debug, Serialize, Deserialize, ToSchema)] +#[derive(Debug, Serialize, Deserialize, ToSchema, PartialEq, Eq)] #[serde(rename_all = "snake_case")] pub enum CompletionObject { TextCompletion, } -#[derive(Debug, Serialize, Deserialize, ToSchema)] +#[derive(Debug, Serialize, Deserialize, ToSchema, PartialEq, Eq)] #[serde(rename_all = "snake_case")] pub enum FinishReason { Stop, @@ -149,7 +123,7 @@ pub struct Message { pub content: String, } -#[derive(Debug, Deserialize, Serialize, ToSchema)] +#[derive(Debug, Deserialize, Serialize, ToSchema, PartialEq, Eq)] #[serde(rename_all = "lowercase")] pub enum Role { System, diff --git a/src/errors.rs b/src/errors.rs index 1de226c..e380c84 100644 --- a/src/errors.rs +++ b/src/errors.rs @@ -15,6 +15,9 @@ pub enum OllamaError { #[error("model '{0}' is not available — run `ollama pull {0}` first")] ModelNotFound(String), + #[error("keep_alive is required and cannot be empty")] + MissingKeepAlive, + #[error( "invalid keep_alive format '{0}' — expected (e.g. 30s, 10m, 2h), a plain integer (seconds), or -1" )] diff --git a/src/providers/ollama/client.rs b/src/providers/ollama/client.rs index a6db618..8719f69 100644 --- a/src/providers/ollama/client.rs +++ b/src/providers/ollama/client.rs @@ -36,6 +36,10 @@ impl OllamaProvider { 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>( &self, body: &'a api::CompletionRequest, @@ -56,6 +60,28 @@ impl OllamaProvider { Ok((&body.messages, model)) } + pub fn parse_keep_alive(&self, s: &str) -> Result<(), OllamaError> { + let s = s.trim(); + + if s == "-1" || s.parse::().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::() + .map_err(|_| OllamaError::InvalidKeepAlive(s.to_string()))?; + + match unit { + "s" | "m" | "h" => Ok(()), + _ => Err(OllamaError::InvalidKeepAlive(s.to_string())), + } + } + // // ── public endpoints ───────────────────────────────────────────────────── pub async fn list_models(&self) -> Result { @@ -77,10 +103,14 @@ impl OllamaProvider { pub async fn load_model( &self, model: &str, - keep_alive: &str, + keep_alive: Option<&str>, ) -> Result { 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?; if !exists { return Err(OllamaError::ModelNotFound(model.to_string())); @@ -99,7 +129,7 @@ impl OllamaProvider { .json(&payload) .send() .await? - .json::() + .text() .await?; Ok(api::LoadModelResponse { @@ -130,7 +160,7 @@ impl OllamaProvider { .json(&payload) .send() .await? - .json::() + .text() .await?; Ok(api::UnloadModelResponse { @@ -277,6 +307,10 @@ impl OllamaProvider { return Err(OllamaError::MissingMessages); } + if !self.has_user_message(&body.messages) { + return Err(OllamaError::MissingMessages); + } + if model.is_empty() { return Err(OllamaError::MissingModel); } diff --git a/src/routes/v1/chat.rs b/src/routes/v1/chat.rs index 9c7ba44..cef14f7 100644 --- a/src/routes/v1/chat.rs +++ b/src/routes/v1/chat.rs @@ -78,6 +78,10 @@ fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) { axum::http::StatusCode::UNPROCESSABLE_ENTITY, 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) => ( axum::http::StatusCode::BAD_REQUEST, format!("invalid keep_alive '{v}'"), diff --git a/src/routes/v1/models.rs b/src/routes/v1/models.rs index 32db1a9..3513487 100644 --- a/src/routes/v1/models.rs +++ b/src/routes/v1/models.rs @@ -20,13 +20,9 @@ pub async fn load_model( Path(model): Path, Json(body): Json, ) -> Result, (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 .ollama - .load_model(&model, keep_alive) + .load_model(&model, body.keep_alive.as_deref()) .await .map_err(ollama_err)?; @@ -64,6 +60,10 @@ fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) { axum::http::StatusCode::BAD_REQUEST, "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) => ( axum::http::StatusCode::BAD_REQUEST, format!("invalid keep_alive '{v}' — use 30s / 10m / 2h, a plain integer, or -1"), diff --git a/tests/ollama_provider.rs b/tests/ollama_provider.rs index 500607d..2c8bd5d 100644 --- a/tests/ollama_provider.rs +++ b/tests/ollama_provider.rs @@ -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]