feat: update test
This commit is contained in:
+4
-30
@@ -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<ModelInfo>,
|
||||
@@ -28,37 +26,13 @@ pub struct LoadModelBody {
|
||||
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)]
|
||||
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,
|
||||
|
||||
@@ -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 <number><unit> (e.g. 30s, 10m, 2h), a plain integer (seconds), or -1"
|
||||
)]
|
||||
|
||||
@@ -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::<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 ─────────────────────────────────────────────────────
|
||||
|
||||
pub async fn list_models(&self) -> Result<api::ModelsResponse, OllamaError> {
|
||||
@@ -77,10 +103,14 @@ impl OllamaProvider {
|
||||
pub async fn load_model(
|
||||
&self,
|
||||
model: &str,
|
||||
keep_alive: &str,
|
||||
keep_alive: Option<&str>,
|
||||
) -> Result<api::LoadModelResponse, OllamaError> {
|
||||
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::<ollama::OllamaGenerateResponse>()
|
||||
.text()
|
||||
.await?;
|
||||
|
||||
Ok(api::LoadModelResponse {
|
||||
@@ -130,7 +160,7 @@ impl OllamaProvider {
|
||||
.json(&payload)
|
||||
.send()
|
||||
.await?
|
||||
.json::<ollama::OllamaGenerateResponse>()
|
||||
.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);
|
||||
}
|
||||
|
||||
@@ -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}'"),
|
||||
|
||||
@@ -20,13 +20,9 @@ pub async fn load_model(
|
||||
Path(model): Path<String>,
|
||||
Json(body): Json<api::LoadModelBody>,
|
||||
) -> 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
|
||||
.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"),
|
||||
|
||||
+132
-105
@@ -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,6 +292,7 @@ async fn test_load_model_ok() {
|
||||
.and(path("/api/generate"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
|
||||
"model": "llama3",
|
||||
"response": "ok",
|
||||
"done": true,
|
||||
})))
|
||||
.mount(&server)
|
||||
@@ -252,34 +300,9 @@ async fn test_load_model_ok() {
|
||||
|
||||
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(_)
|
||||
@@ -335,15 +360,18 @@ async fn test_load_model_keep_alive_plain_integer() {
|
||||
.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]
|
||||
|
||||
Reference in New Issue
Block a user