diff --git a/src/docs.rs b/src/docs.rs new file mode 100644 index 0000000..48a62a1 --- /dev/null +++ b/src/docs.rs @@ -0,0 +1,46 @@ +use utoipa::OpenApi; + +use crate::dto::api; +use crate::routes; + +#[derive(OpenApi)] +#[openapi( + paths( + routes::v1::chat::completions, + routes::v1::chat::chat_completions, + // routes::v1::models::list_models, + // routes::v1::models::load_model, + // routes::v1::models::unload_model, + ), + components( + schemas( + api::ErrorResponse, + api::ModelsResponse, + api::ModelInfo, + api::LoadModelResponse, + api::LoadModelBody, + api::UnloadModelResponse, + api::BaseLLMRequest, + api::CompletionRequest, + api::CompletionObject, + api::FinishReason, + api::CompletionResponse, + api::Choice, + api::Usage, + api::CompletionChunk, + api::ChatRequest, + api::Message, + api::Role, + api::ChatCompletionResponse, + api::ChatChoice, + api::ChatCompletionChunk, + api::ChatChunkChoice, + api::ChatDelta, + ) + ), + tags( + (name = "chat", description = "Chat & completions"), + (name = "models", description = "Model management") + ) +)] +pub struct ApiDoc; diff --git a/src/dto/api.rs b/src/dto/api.rs index 5a8a1b6..cf0e199 100644 --- a/src/dto/api.rs +++ b/src/dto/api.rs @@ -1,6 +1,17 @@ use serde::{Deserialize, Serialize}; use utoipa::ToSchema; +#[derive(Debug, Serialize, ToSchema)] +pub struct ErrorResponse { + pub error: String, +} + +impl ErrorResponse { + pub fn new(msg: impl Into) -> Self { + Self { error: msg.into() } + } +} + #[derive(Debug, Serialize, Deserialize, ToSchema)] pub struct ModelsResponse { pub models: Vec, @@ -148,21 +159,21 @@ pub struct ChatChoice { pub finish_reason: FinishReason, } -#[derive(Debug, Serialize)] +#[derive(Debug, Serialize, ToSchema)] pub struct ChatCompletionChunk { pub id: String, pub object: String, pub choices: Vec, } -#[derive(Debug, Serialize)] +#[derive(Debug, Serialize, ToSchema)] pub struct ChatChunkChoice { pub index: u32, pub delta: ChatDelta, pub finish_reason: Option, } -#[derive(Debug, Serialize)] +#[derive(Debug, Serialize, ToSchema)] pub struct ChatDelta { pub role: Option, pub content: Option, diff --git a/src/main.rs b/src/main.rs index 125bbed..4ca9145 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,7 +1,7 @@ mod auth; +mod docs; mod dto; mod errors; -mod openapi; mod providers; mod routes; mod state; diff --git a/src/openapi.rs b/src/openapi.rs deleted file mode 100644 index f510f44..0000000 --- a/src/openapi.rs +++ /dev/null @@ -1,18 +0,0 @@ -// use utoipa::OpenApi; - -// #[derive(OpenApi)] -// #[openapi( -// paths( -// crate::routes::v1::chat::completions -// ), -// components( -// schemas( -// // add your request/response structs here later -// ) -// ), -// tags( -// (name = "chat", description = "Chat endpoints"), -// (name = "models", description = "Model management") -// ) -// )] -// pub struct V1ApiDoc; diff --git a/src/routes/v1/chat.rs b/src/routes/v1/chat.rs index cef14f7..19f6499 100644 --- a/src/routes/v1/chat.rs +++ b/src/routes/v1/chat.rs @@ -13,32 +13,105 @@ use crate::state::app_state::AppState; #[utoipa::path( post, - path = "/chat/completions", + path = "/completions", + tag = "chat", + request_body( + content = api::CompletionRequest, + description = "Text completion request", + content_type = "application/json" + ), responses( - (status = 200, description = "Chat completion", body = Value), + ( + status = 200, + description = "Text completion response. If stream=true, response is SSE stream of chunks ending in [DONE].", + body = api::CompletionResponse, + content_type = "application/json" + ), + ( + status = 400, + description = "Invalid request: missing prompt, model, or invalid format", + body = api::ErrorResponse, + example = json!({ "error": "prompt is required and cannot be empty" }) + ), + ( + status = 422, + description = "Model not found or not available locally", + body = api::ErrorResponse, + example = json!({ "error": "model 'llama3' is not available — run `ollama pull llama3` first" }) + ), + ( + status = 500, + description = "Internal server error (Ollama or network failure)", + body = api::ErrorResponse, + example = json!({ "error": "connection refused" }) + ) ) )] pub async fn completions( State(state): State, Json(body): Json, -) -> Result { +) -> Result)> { if body.base.stream { - let stream = state - .ollama - .completions_stream(&body) - .await - .map_err(ollama_err)?; + let stream = state.ollama.completions_stream(&body).await.map_err(|e| { + let (code, msg) = ollama_err(e); + (code, Json(api::ErrorResponse::new(msg))) + })?; Ok(Sse::new(stream) .keep_alive(KeepAlive::default()) .into_response()) } else { - let response = state.ollama.completions(&body).await.map_err(ollama_err)?; + let response = state.ollama.completions(&body).await.map_err(|e| { + let (code, msg) = ollama_err(e); + (code, Json(api::ErrorResponse::new(msg))) + })?; Ok(Json(response).into_response()) } } +#[utoipa::path( + post, + path = "/chat/completions", + tag = "chat", + request_body( + content = api::ChatRequest, + description = "Chat completion request with message history", + content_type = "application/json" + ), + responses( + ( + status = 200, + description = "Chat completion response. If stream=false returns JSON. If stream=true returns SSE stream of chunks ending with [DONE].", + body = api::ChatCompletionResponse, + content_type = "application/json" + ), + ( + status = 400, + description = "Invalid request", + body = api::ErrorResponse, + example = json!({ "error": "messages array with at least one user message is required" }) + ), + ( + status = 401, + description = "Unauthorized", + body = api::ErrorResponse, + example = json!({ "error": "missing or invalid token" }) + ), + ( + status = 422, + description = "Model not found or unavailable", + body = api::ErrorResponse, + example = json!({ "error": "model 'llama3' is not available — run `ollama pull llama3` first" }) + ), + ( + status = 500, + description = "Internal server error", + body = api::ErrorResponse, + example = json!({ "error": "connection refused" }) + ) + ) +)] pub async fn chat_completions( State(state): State, Json(body): Json, diff --git a/src/routes/v1/mod.rs b/src/routes/v1/mod.rs index 0d7f6a0..a4b3028 100644 --- a/src/routes/v1/mod.rs +++ b/src/routes/v1/mod.rs @@ -1,14 +1,19 @@ pub mod chat; pub mod models; -mod openapi; use crate::auth::middleware::auth_middleware; +use crate::docs::ApiDoc; use crate::state::app_state::AppState; -use axum::{Router, middleware, routing::get, routing::post}; +use axum::{Json, Router, middleware, routing::get, routing::post}; +use utoipa::OpenApi; -// fn public_router() -> Router { -// Router::new().route("/openapi.json", get(openapi::openapi_json)) -// } +async fn openapi_json() -> Json { + Json(ApiDoc::openapi()) +} + +fn public_router() -> Router { + Router::new().route("/openapi.json", get(openapi_json)) +} pub fn protected_router() -> Router { Router::new() @@ -21,6 +26,6 @@ pub fn protected_router() -> Router { pub fn router() -> Router { Router::new() - // .merge(public_router()) + .merge(public_router()) .merge(protected_router().layer(middleware::from_fn(auth_middleware))) } diff --git a/src/routes/v1/models.rs b/src/routes/v1/models.rs index 3513487..58a1d7a 100644 --- a/src/routes/v1/models.rs +++ b/src/routes/v1/models.rs @@ -6,6 +6,25 @@ use axum::{ use crate::dto::api; use crate::{errors::OllamaError, state::app_state::AppState}; +#[utoipa::path( + get, + path = "/models", + tag = "models", + responses( + ( + status = 200, + description = "List of locally available Ollama models", + body = api::ModelsResponse, + content_type = "application/json", + ), + ( + status = 500, + description = "Internal server error (Ollama or network failure)", + body = api::ErrorResponse, + example = json!({ "error": "connection refused" }) + ) + ) +)] pub async fn list_models( State(state): State, ) -> Result, (axum::http::StatusCode, String)> { @@ -15,6 +34,49 @@ pub async fn list_models( } } +#[utoipa::path( + post, + path = "/models/{model}/load", + tag = "models", + params( + ("model" = String, Path, description = "Name of the model to load into memory (e.g. 'llama3')") + ), + request_body( + content = api::LoadModelBody, + description = "Load model request", + content_type = "application/json", + example = json!({ "keep_alive": "10m" }) + ), + responses( + ( + status = 200, + description = "Model successfully loaded into memory", + body = api::LoadModelResponse, + content_type = "application/json", + ), + ( + status = 400, + description = "Invalid or missing keep_alive format", + body = api::ErrorResponse, + examples( + ("Missing" = (value = json!({ "error": "keep alive is required and cannot be empty" }))), + ("Invalid" = (value = json!({ "error": "invalid keep_alive '10x' — use 30s / 10m / 2h, a plain integer, or -1" }))) + ) + ), + ( + status = 404, + description = "Model not found locally", + body = api::ErrorResponse, + example = json!({ "error": "model 'llama3' not found — run `ollama pull llama3`" }) + ), + ( + status = 500, + description = "Internal server error (Ollama or network failure)", + body = api::ErrorResponse, + example = json!({ "error": "connection refused" }) + ) + ) +)] pub async fn load_model( State(state): State, Path(model): Path, @@ -29,6 +91,34 @@ pub async fn load_model( Ok(Json(response)) } +#[utoipa::path( + delete, + path = "/models/{model}/load", + tag = "models", + params( + ("model" = String, Path, description = "Name of the model to unload from memory (e.g. 'llama3')") + ), + responses( + ( + status = 200, + description = "Model successfully unloaded from memory", + body = api::UnloadModelResponse, + content_type = "application/json", + ), + ( + status = 404, + description = "Model not found locally", + body = api::ErrorResponse, + example = json!({ "error": "model 'llama3' not found — run `ollama pull llama3`" }) + ), + ( + status = 500, + description = "Internal server error (Ollama or network failure)", + body = api::ErrorResponse, + example = json!({ "error": "connection refused" }) + ) + ) +)] pub async fn unload_model( State(state): State, Path(model): Path,