diff --git a/Cargo.lock b/Cargo.lock index c4d9d26..8eb57c3 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -154,6 +154,7 @@ dependencies = [ "reqwest", "serde", "serde_json", + "thiserror 2.0.18", "tokio", ] diff --git a/Cargo.toml b/Cargo.toml index 743da72..e22d737 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -11,4 +11,5 @@ serde_json = "1" jsonwebtoken = { version = "10.3.0", features = ["aws_lc_rs"] } reqwest = { version = "0.13.2", features = ["json"] } once_cell = "1" -dotenvy = "0.15" \ No newline at end of file +dotenvy = "0.15" +thiserror = "2.0.18" diff --git a/src/errors.rs b/src/errors.rs new file mode 100644 index 0000000..4e45b57 --- /dev/null +++ b/src/errors.rs @@ -0,0 +1,12 @@ +// errors.rs +use thiserror::Error; + +#[derive(Debug, Error)] +pub enum OllamaError { + #[error("prompt is required and cannot be empty")] + MissingPrompt, + #[error("model '{0}' is not available — run `ollama pull {0}` first")] + ModelNotFound(String), + #[error(transparent)] + Http(#[from] reqwest::Error), +} diff --git a/src/main.rs b/src/main.rs index 11e70de..9f8d122 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,4 +1,5 @@ mod auth; +mod errors; mod providers; mod routes; mod state; diff --git a/src/providers/ollama.rs b/src/providers/ollama.rs index 0cd93b1..9ecfca9 100644 --- a/src/providers/ollama.rs +++ b/src/providers/ollama.rs @@ -1,5 +1,7 @@ use reqwest::Client; -use serde_json::Value; +use serde_json::{Value, json}; + +use crate::errors::OllamaError; #[derive(Clone)] pub struct OllamaProvider { @@ -15,11 +17,82 @@ impl OllamaProvider { } } - pub async fn list_models(&self) -> Result { + 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()) + .filter(|s| !s.trim().is_empty()) + .ok_or(OllamaError::MissingPrompt)?; + + let model = body + .get("model") + .and_then(|v| v.as_str()) + .unwrap_or("llama3"); + + // 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"), + } + }); + + let url = format!("{}/api/generate", self.base_url); + let res = self + .client + .post(&url) + .json(&ollama_payload) + .send() + .await? + .json::() + .await?; + + Ok(json!({ + "id": "cmpl-ollama", + "object": "text_completion", + "model": res.get("model"), + "choices": [{ + "text": res.get("response"), + "index": 0, + "finish_reason": if res.get("done").and_then(|v| v.as_bool()).unwrap_or(false) { + "stop" + } else { + "length" + }, + }], + "usage": { + "prompt_tokens": res.get("prompt_eval_count"), + "completion_tokens": res.get("eval_count"), + "total_tokens": null, + } + })) + } } diff --git a/src/routes/v1/mod.rs b/src/routes/v1/mod.rs index 0e98fa6..ea66a57 100644 --- a/src/routes/v1/mod.rs +++ b/src/routes/v1/mod.rs @@ -1,12 +1,12 @@ pub mod models; use crate::auth::middleware::auth_middleware; -use crate::routes::v1::models::list_models; use crate::state::app_state::AppState; -use axum::{Router, middleware, routing::get}; +use axum::{Router, middleware, routing::get, routing::post}; pub fn router() -> Router { Router::new() - .route("/models", get(list_models)) + .route("/models", get(models::list_models)) + .route("/completions", post(models::completions)) .layer(middleware::from_fn(auth_middleware)) } diff --git a/src/routes/v1/models.rs b/src/routes/v1/models.rs index d2c966b..1eeafc9 100644 --- a/src/routes/v1/models.rs +++ b/src/routes/v1/models.rs @@ -1,6 +1,7 @@ use axum::{Json, extract::State}; use serde_json::Value; +use crate::errors::OllamaError; use crate::state::app_state::AppState; pub async fn list_models( @@ -14,3 +15,23 @@ pub async fn list_models( )), } } + +pub async fn completions( + State(state): State, + Json(body): Json, +) -> Result, (axum::http::StatusCode, String)> { + match state.ollama.completions(body).await { + Ok(response) => Ok(Json(response)), + Err(OllamaError::MissingPrompt) => Err(( + axum::http::StatusCode::BAD_REQUEST, + "prompt is required and cannot be empty".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())) + } + } +}