diff --git a/.gitea/workflows/ci.yml b/.gitea/workflows/ci.yml index b6f19f0..7d80762 100644 --- a/.gitea/workflows/ci.yml +++ b/.gitea/workflows/ci.yml @@ -32,8 +32,8 @@ jobs: restore-keys: | ${{ runner.os }}-cargo- - # - name: Run tests - # run: cargo test --workspace --verbose + - name: Run tests + run: cargo test --workspace --verbose --all - name: Run Clippy (linter) run: cargo clippy --all-targets --all-features -- -D warnings diff --git a/Cargo.lock b/Cargo.lock index 8eb57c3..831ae8b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2,6 +2,25 @@ # It is not intended for manual editing. version = 4 +[[package]] +name = "aho-corasick" +version = "1.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ddd31a130427c27518df266943a5308ed92d4b226cc639f5a8f1002816174301" +dependencies = [ + "memchr", +] + +[[package]] +name = "assert-json-diff" +version = "2.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47e4f2b81832e72834d7518d8487a0396a28cc408186a2e8854c0f98011faf12" +dependencies = [ + "serde", + "serde_json", +] + [[package]] name = "atomic-waker" version = "1.1.2" @@ -156,6 +175,7 @@ dependencies = [ "serde_json", "thiserror 2.0.18", "tokio", + "wiremock", ] [[package]] @@ -203,6 +223,24 @@ version = "0.8.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "773648b94d0e5d620f64f280777445740e61fe701025087ec8b57f45c791888b" +[[package]] +name = "deadpool" +version = "0.12.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0be2b1d1d6ec8d846f05e137292d0b89133caf95ef33695424c09568bdd39b1b" +dependencies = [ + "deadpool-runtime", + "lazy_static", + "num_cpus", + "tokio", +] + +[[package]] +name = "deadpool-runtime" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "092966b41edc516079bdf31ec78a2e0588d1d0c08f78b91d8307215928642b2b" + [[package]] name = "deranged" version = "0.5.8" @@ -287,6 +325,21 @@ version = "1.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c" +[[package]] +name = "futures" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b147ee9d1f6d097cef9ce628cd2ee62288d963e16fb287bd9286455b241382d" +dependencies = [ + "futures-channel", + "futures-core", + "futures-executor", + "futures-io", + "futures-sink", + "futures-task", + "futures-util", +] + [[package]] name = "futures-channel" version = "0.3.32" @@ -294,6 +347,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "07bbe89c50d7a535e539b8c17bc0b49bdb77747034daa8087407d655f3f7cc1d" dependencies = [ "futures-core", + "futures-sink", ] [[package]] @@ -302,6 +356,34 @@ version = "0.3.32" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7e3450815272ef58cec6d564423f6e755e25379b217b0bc688e295ba24df6b1d" +[[package]] +name = "futures-executor" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "baf29c38818342a3b26b5b923639e7b1f4a61fc5e76102d4b1981c6dc7a7579d" +dependencies = [ + "futures-core", + "futures-task", + "futures-util", +] + +[[package]] +name = "futures-io" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cecba35d7ad927e23624b22ad55235f2239cfa44fd10428eecbeba6d6a717718" + +[[package]] +name = "futures-macro" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e835b70203e41293343137df5c0664546da5745f82ec9b84d40be8336958447b" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "futures-sink" version = "0.3.32" @@ -320,8 +402,13 @@ version = "0.3.32" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "389ca41296e6190b48053de0321d02a77f32f8a5d2461dd38762c0593805c6d6" dependencies = [ + "futures-channel", "futures-core", + "futures-io", + "futures-macro", + "futures-sink", "futures-task", + "memchr", "pin-project-lite", "slab", ] @@ -378,6 +465,12 @@ version = "0.16.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "841d1cc9bed7f9236f321df977030373f4a4163ae1a7dbfe1a51a2c1a51d9100" +[[package]] +name = "hermit-abi" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" + [[package]] name = "http" version = "1.4.0" @@ -704,6 +797,12 @@ dependencies = [ "simple_asn1", ] +[[package]] +name = "lazy_static" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" + [[package]] name = "libc" version = "0.2.184" @@ -800,6 +899,16 @@ dependencies = [ "autocfg", ] +[[package]] +name = "num_cpus" +version = "1.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91df4bbde75afed763b708b7eee1e8e7651e02d97f6d5dd763e89367e957b23b" +dependencies = [ + "hermit-abi", + "libc", +] + [[package]] name = "once_cell" version = "1.21.4" @@ -1008,6 +1117,35 @@ dependencies = [ "bitflags", ] +[[package]] +name = "regex" +version = "1.12.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e10754a14b9137dd7b1e3e5b0493cc9171fdd105e0ab477f51b72e7f3ac0e276" +dependencies = [ + "aho-corasick", + "memchr", + "regex-automata", + "regex-syntax", +] + +[[package]] +name = "regex-automata" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6e1dd4122fc1595e8162618945476892eefca7b88c52820e74af6262213cae8f" +dependencies = [ + "aho-corasick", + "memchr", + "regex-syntax", +] + +[[package]] +name = "regex-syntax" +version = "0.8.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc897dd8d9e8bd1ed8cdad82b5966c3e0ecae09fb1907d58efaa013543185d0a" + [[package]] name = "reqwest" version = "0.13.2" @@ -2031,6 +2169,29 @@ version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d6bbff5f0aada427a1e5a6da5f1f98158182f26556f345ac9e04d36d0ebed650" +[[package]] +name = "wiremock" +version = "0.6.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08db1edfb05d9b3c1542e521aea074442088292f00b5f28e435c714a98f85031" +dependencies = [ + "assert-json-diff", + "base64", + "deadpool", + "futures", + "http", + "http-body-util", + "hyper", + "hyper-util", + "log", + "once_cell", + "regex", + "serde", + "serde_json", + "tokio", + "url", +] + [[package]] name = "wit-bindgen" version = "0.51.0" diff --git a/Cargo.toml b/Cargo.toml index e22d737..8cc7e0f 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -3,6 +3,10 @@ name = "chat" version = "0.1.0" edition = "2024" +[dev-dependencies] +wiremock = "0.6" +tokio = { version = "1", features = ["macros", "rt-multi-thread"] } + [dependencies] axum = "0.8.8" tokio = { version = "1", features = ["full"] } diff --git a/readme.md b/readme.md deleted file mode 100644 index 98e7632..0000000 --- a/readme.md +++ /dev/null @@ -1,396 +0,0 @@ -# TODO - -- Race condition on jwks token refresh -- Rate Limiting - - -git tag -d v1.0.0; git push origin :refs/tags/v1.0.0; git tag -a v1.0.0 -m "Release v1.0.0"; git push origin v1.0.0 - - -# ๐Ÿฆ™ Ollama Rust API Wrapper - -A high-performance Rust API wrapper around Ollama, providing an OpenAI-compatible interface, model lifecycle management, and advanced runtime features. - ---- - -# ๐Ÿš€ Features - -* โœ… OpenAI-compatible API (`/v1/...`) -* โšก Streaming (Server-Sent Events) -* ๐Ÿง  Model lifecycle management (load/unload) -* ๐Ÿ” API key authentication (optional) -* ๐Ÿ“Š Usage tracking & observability -* ๐Ÿ”€ Model routing & abstraction -* ๐Ÿงฉ Extensible architecture (multi-provider ready) - ---- - -# ๐Ÿ“ก API Endpoints - -## 1. Core LLM API (OpenAI-compatible) - -### Chat Completions - -``` -POST /v1/chat/completions -``` - -### Text Completions - -``` -POST /v1/completions -``` - -### Embeddings - -``` -POST /v1/embeddings -``` - -### List Models - -``` -GET /v1/models -``` - ---- - -## 2. Model Lifecycle Management - -### Load (Warmup) - -``` -POST /v1/models/{model}/load -``` - -```json -{ - "keep_alive": "10m" -} -``` - ---- - -### Unload (Free Memory) - -``` -POST /v1/models/{model}/unload -``` - -Internally uses: - -```json -{ - "model": "...", - "keep_alive": 0 -} -``` - ---- - -### Reload (Optional) - -``` -POST /v1/models/{model}/reload -``` - ---- - -## 3. Model Management - -### Pull Model - -``` -POST /v1/models/pull -``` - -### Delete Model - -``` -DELETE /v1/models/{model} -``` - ---- - -## 4. Runtime & Observability - -### Model Status - -``` -GET /v1/models/{model}/status -``` - -### List Loaded Models - -``` -GET /v1/runtime/models -``` - ---- - -## 5. Streaming - -Enable streaming with: - -```json -{ - "stream": true -} -``` - -Response format (SSE): - -``` -data: {"choices":[{"delta":{"content":"Hello"}}]} - -data: {"choices":[{"delta":{"content":" world"}}]} - -data: [DONE] -``` - ---- - -## 6. Health Checks - -``` -GET /health -GET /ready -``` - ---- - -# ๐Ÿง  Internal Mapping (Ollama) - -| Wrapper Endpoint | Ollama Endpoint | -| ------------------------- | --------------- | -| /v1/chat/completions | /api/chat | -| /v1/completions | /api/generate | -| /v1/embeddings | /api/embeddings | -| /v1/models | /api/tags | -| /v1/models/pull | /api/pull | -| DELETE /v1/models/{model} | /api/delete | -| load/unload | /api/generate | - ---- - -# โš™๏ธ Configuration - -### Docker (optional default) - -```yaml -environment: - - OLLAMA_KEEP_ALIVE=10m -``` - -> Note: Request-level `keep_alive` overrides this value. - ---- - -# ๐Ÿ”ง Advanced Features - -## ๐Ÿ”€ Model Routing - -Use abstract model names: - -```json -{ - "model": "fast" -} -``` - -Example mapping: - -``` -fast โ†’ llama3:8b -smart โ†’ llama3:70b -code โ†’ deepseek-coder -``` - ---- - -## ๐Ÿ“Š Usage Tracking - -``` -GET /v1/usage -``` - -Tracks: - -* request count -* latency -* per-model usage - ---- - -## ๐Ÿ” Authentication - -``` -Authorization: Bearer sk-xxxx -``` - -Endpoints: - -``` -POST /v1/keys -DELETE /v1/keys/{id} -``` - ---- - -## ๐Ÿšฆ Rate Limiting - -* Requests per minute -* Tokens per minute - -Returns: - -``` -429 Too Many Requests -``` - ---- - -## ๐Ÿง  Sessions (Context Management) - -``` -POST /v1/sessions -POST /v1/sessions/{id}/chat -``` - -Stores conversation history server-side. - ---- - -## โšก Caching - -* Embeddings -* Deterministic prompts (temperature = 0) - ---- - -## ๐Ÿงฉ Tool / Function Calling - -Supports structured tool execution: - -```json -{ - "tools": [ - { - "name": "function_name", - "parameters": {} - } - ] -} -``` - ---- - -## ๐Ÿ“ฆ Batch Requests - -``` -POST /v1/batch -``` - ---- - -## ๐Ÿง  Auto Eviction - -``` -POST /v1/runtime/evict -``` - -Strategies: - -* LRU -* memory threshold - ---- - -## ๐Ÿงพ Logs - -``` -GET /v1/logs -``` - ---- - -## ๐Ÿ“š Embedding Store (Optional) - -``` -POST /v1/documents -POST /v1/search -``` - ---- - -## ๐Ÿ”” Async Jobs / Webhooks - -``` -POST /v1/jobs -``` - ---- - -# ๐Ÿ—๏ธ Architecture - -``` -Client โ†’ Rust API โ†’ Ollama โ†’ Response -``` - -### Layers: - -* HTTP (Axum) -* Service layer (business logic) -* Provider abstraction -* Ollama client - ---- - -# ๐Ÿ”Œ Provider Abstraction (Future-Proof) - -```rust -trait LlmProvider { - async fn chat(...); - async fn embeddings(...); -} -``` - -Supports: - -* Ollama (current) -* OpenAI (future) -* Others - ---- - -# โš ๏ธ Notes - -* Ollama has no native unload endpoint โ†’ simulated via `keep_alive = 0` -* Streaming uses NDJSON โ†’ converted to SSE -* Chunk handling must be robust (partial JSON) - ---- - -# ๐ŸŽฏ Roadmap - -* [ ] Full OpenAI compatibility -* [ ] Multi-node routing -* [ ] GPU-aware scheduling -* [ ] Web UI dashboard -* [ ] Distributed inference - ---- - -# ๐Ÿง  Summary - -This project turns Ollama into: - -๐Ÿ‘‰ A local OpenAI-compatible API -๐Ÿ‘‰ A controllable model runtime -๐Ÿ‘‰ A foundation for a full LLM gateway - ---- - -# ๐Ÿ“œ License - -MIT diff --git a/src/errors.rs b/src/errors.rs index 4e45b57..355de13 100644 --- a/src/errors.rs +++ b/src/errors.rs @@ -5,6 +5,8 @@ use thiserror::Error; pub enum OllamaError { #[error("prompt is required and cannot be empty")] MissingPrompt, + #[error("messages must be a non-empty array containing at least one user message")] + MissingMessages, #[error("model '{0}' is not available โ€” run `ollama pull {0}` first")] ModelNotFound(String), #[error(transparent)] diff --git a/src/lib.rs b/src/lib.rs new file mode 100644 index 0000000..bd9bc1a --- /dev/null +++ b/src/lib.rs @@ -0,0 +1,2 @@ +pub mod errors; +pub mod providers; diff --git a/src/providers/ollama.rs b/src/providers/ollama.rs index 9ecfca9..b448df3 100644 --- a/src/providers/ollama.rs +++ b/src/providers/ollama.rs @@ -1,8 +1,7 @@ +use crate::errors::OllamaError; use reqwest::Client; use serde_json::{Value, json}; -use crate::errors::OllamaError; - #[derive(Clone)] pub struct OllamaProvider { pub client: Client, @@ -17,16 +16,42 @@ impl OllamaProvider { } } + // โ”€โ”€ private helpers โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€ + + fn build_options(body: &Value) -> Value { + json!({ + "temperature": body.get("temperature"), + "top_p": body.get("top_p"), + "num_predict": body.get("max_tokens"), + }) + } + + async fn validate_model(&self, model: &str) -> Result<(), OllamaError> { + let available = self.list_models().await?; + let exists = available + .get("models") + .and_then(|m| m.as_array()) + .map(|arr| { + arr.iter() + .any(|m| m.get("name").and_then(|n| n.as_str()) == Some(model)) + }) + .unwrap_or(false); + + if !exists { + return Err(OllamaError::ModelNotFound(model.to_string())); + } + Ok(()) + } + + // โ”€โ”€ public endpoints โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€ + pub async fn list_models(&self) -> Result { let url = format!("{}/api/tags", self.base_url); - let res = self.client.get(url).send().await?.json::().await?; - Ok(res) } pub async fn completions(&self, body: Value) -> Result { - // Validate prompt let prompt = body .get("prompt") .and_then(|v| v.as_str()) @@ -37,38 +62,18 @@ impl OllamaProvider { .get("model") .and_then(|v| v.as_str()) .unwrap_or("llama3"); + self.validate_model(model).await?; - // Validate model availability - let available = self.list_models().await?; // Http error auto-converts via #[from] - let model_exists = available - .get("models") - .and_then(|m| m.as_array()) - .map(|arr| { - arr.iter() - .any(|m| m.get("name").and_then(|n| n.as_str()) == Some(model)) - }) - .unwrap_or(false); - - if !model_exists { - return Err(OllamaError::ModelNotFound(model.to_string())); - } - - // Build payload (now using validated locals) let ollama_payload = json!({ - "model": model, - "prompt": prompt, - "stream": false, - "options": { - "temperature": body.get("temperature"), - "top_p": body.get("top_p"), - "num_predict": body.get("max_tokens"), - } + "model": model, + "prompt": prompt, + "stream": false, + "options": Self::build_options(&body), }); - let url = format!("{}/api/generate", self.base_url); let res = self .client - .post(&url) + .post(format!("{}/api/generate", self.base_url)) .json(&ollama_payload) .send() .await? @@ -95,4 +100,68 @@ impl OllamaProvider { } })) } + + pub async fn chat_completions(&self, body: Value) -> Result { + let model = body + .get("model") + .and_then(|v| v.as_str()) + .unwrap_or("llama3"); + + let messages = body + .get("messages") + .and_then(|v| v.as_array()) + .filter(|arr| !arr.is_empty()) + .ok_or(OllamaError::MissingMessages)?; + + let has_user_msg = messages + .iter() + .any(|m| m.get("role").and_then(|r| r.as_str()) == Some("user")); + + if !has_user_msg { + return Err(OllamaError::MissingMessages); + } + + self.validate_model(model).await?; + + let ollama_payload = json!({ + "model": model, + "messages": messages, + "stream": false, + "options": Self::build_options(&body), + }); + + let res = self + .client + .post(format!("{}/api/chat", self.base_url)) + .json(&ollama_payload) + .send() + .await? + .json::() + .await?; + + Ok(json!({ + "id": "chatcmpl-ollama", + "object": "chat.completion", + "model": res.get("model"), + "choices": [{ + "index": 0, + "message": { + "role": res.get("message").and_then(|m| m.get("role")), + "content": res.get("message").and_then(|m| m.get("content")), + }, + "finish_reason": res + .get("done_reason") + .and_then(|v| v.as_str()) + .unwrap_or("stop"), + }], + "usage": { + "prompt_tokens": res.get("prompt_eval_count"), + "completion_tokens": res.get("eval_count"), + "total_tokens": res.get("prompt_eval_count") + .and_then(|p| p.as_u64()) + .zip(res.get("eval_count").and_then(|e| e.as_u64())) + .map(|(p, e)| p + e), + } + })) + } } diff --git a/src/routes/v1/mod.rs b/src/routes/v1/mod.rs index ea66a57..3e1f749 100644 --- a/src/routes/v1/mod.rs +++ b/src/routes/v1/mod.rs @@ -8,5 +8,6 @@ pub fn router() -> Router { Router::new() .route("/models", get(models::list_models)) .route("/completions", post(models::completions)) + .route("/chat/completions", post(models::chat_completions)) .layer(middleware::from_fn(auth_middleware)) } diff --git a/src/routes/v1/models.rs b/src/routes/v1/models.rs index 1eeafc9..169fa38 100644 --- a/src/routes/v1/models.rs +++ b/src/routes/v1/models.rs @@ -30,8 +30,27 @@ pub async fn completions( axum::http::StatusCode::UNPROCESSABLE_ENTITY, format!("model '{m}' is not available โ€” run `ollama pull {m}` first"), )), + Err(e) => Err((axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string())), + } +} + +pub async fn chat_completions( + State(state): State, + Json(body): Json, +) -> Result, (axum::http::StatusCode, String)> { + match state.ollama.chat_completions(body).await { + Ok(response) => Ok(Json(response)), + Err(OllamaError::MissingMessages) => Err(( + axum::http::StatusCode::BAD_REQUEST, + "messages must be a non-empty array with at least one user message".to_string(), + )), + Err(OllamaError::ModelNotFound(m)) => Err(( + axum::http::StatusCode::UNPROCESSABLE_ENTITY, + format!("model '{m}' is not available โ€” run `ollama pull {m}` first"), + )), Err(OllamaError::Http(e)) => { Err((axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string())) } + Err(e) => Err((axum::http::StatusCode::BAD_REQUEST, e.to_string())), } } diff --git a/tests/ollama_provider.rs b/tests/ollama_provider.rs new file mode 100644 index 0000000..fa6f7b0 --- /dev/null +++ b/tests/ollama_provider.rs @@ -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::>() + }) +} + +// โ”€โ”€ 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(_))); +}