feat: add chat completions route + test
CI / Rust CI (push) Successful in 4m23s
Publish & Deploy / Build and Push to Registry (push) Successful in 1m26s
Publish & Deploy / Deploy via SSH (push) Successful in 18s

This commit is contained in:
2026-04-10 11:47:56 +02:00
parent 70fe9bc4da
commit 6dc7231309
10 changed files with 521 additions and 429 deletions
+100 -31
View File
@@ -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<Value, OllamaError> {
let url = format!("{}/api/tags", self.base_url);
let res = self.client.get(url).send().await?.json::<Value>().await?;
Ok(res)
}
pub async fn completions(&self, body: Value) -> Result<Value, OllamaError> {
// 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<Value, OllamaError> {
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::<Value>()
.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),
}
}))
}
}