use serde_json::json; use wiremock::matchers::{method, path}; use wiremock::{Mock, MockServer, ResponseTemplate}; use chat::dto::api; use chat::errors::OllamaError; use chat::providers::ollama::client::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::>() }) } // ── 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 req = api::CompletionRequest { base: api::BaseLLMRequest { model: "llama3".to_string(), ..Default::default() }, prompt: "Hello".to_string(), }; 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] 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 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] 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 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] 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 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(_))); } // ── 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, "prompt_eval_count": 5, "eval_count": 2 }))) .mount(&server) .await; 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(), }], }; 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] 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 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)); } #[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 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)); } #[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 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 ──────────────────────────────────────────────────────────────── #[tokio::test] async fn test_load_model_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": "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_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.load_model("gpt-4", Some("10m")).await.unwrap_err(); assert!(matches!(err, chat::errors::OllamaError::ModelNotFound(_))); } #[tokio::test] async fn test_load_model_invalid_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; let err = provider .load_model("llama3", Some("10x")) .await .unwrap_err(); assert!(matches!( err, chat::errors::OllamaError::InvalidKeepAlive(_) )); } #[tokio::test] async fn test_load_model_keep_alive_plain_integer() { 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", Some("3600")).await.unwrap(); 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.keep_alive, "-1"); } // // ── unload_model ────────────────────────────────────────────────────────────── #[tokio::test] async fn test_unload_model_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": "ok", "done": true, }))) .mount(&server) .await; let res = provider.unload_model("llama3").await.unwrap(); assert_eq!(res.model, "llama3"); assert_eq!(res.status, "unloaded"); } #[tokio::test] async fn test_unload_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.unload_model("gpt-4").await.unwrap_err(); assert!(matches!(err, chat::errors::OllamaError::ModelNotFound(_))); }