feat: add chat completions route + test
This commit is contained in:
@@ -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)]
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
pub mod errors;
|
||||
pub mod providers;
|
||||
+100
-31
@@ -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),
|
||||
}
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -8,5 +8,6 @@ pub fn router() -> Router<AppState> {
|
||||
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))
|
||||
}
|
||||
|
||||
@@ -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<AppState>,
|
||||
Json(body): Json<Value>,
|
||||
) -> Result<Json<Value>, (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())),
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user