feat: add proper typing for api
This commit is contained in:
+48
-48
@@ -18,61 +18,61 @@ use crate::state::app_state::AppState;
|
||||
(status = 200, description = "Chat completion", body = Value),
|
||||
)
|
||||
)]
|
||||
pub async fn completions(
|
||||
State(state): State<AppState>,
|
||||
Json(body): Json<Value>,
|
||||
) -> Result<Response, (axum::http::StatusCode, String)> {
|
||||
let wants_stream = body
|
||||
.get("stream")
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(false);
|
||||
// pub async fn completions(
|
||||
// State(state): State<AppState>,
|
||||
// Json(body): Json<Value>,
|
||||
// ) -> Result<Response, (axum::http::StatusCode, String)> {
|
||||
// let wants_stream = body
|
||||
// .get("stream")
|
||||
// .and_then(|v| v.as_bool())
|
||||
// .unwrap_or(false);
|
||||
|
||||
if wants_stream {
|
||||
let stream = state
|
||||
.ollama
|
||||
.completions_stream(body)
|
||||
.await
|
||||
.map_err(ollama_err)?;
|
||||
// if wants_stream {
|
||||
// let stream = state
|
||||
// .ollama
|
||||
// .completions_stream(body)
|
||||
// .await
|
||||
// .map_err(ollama_err)?;
|
||||
|
||||
Ok(Sse::new(stream)
|
||||
.keep_alive(KeepAlive::default())
|
||||
.into_response())
|
||||
} else {
|
||||
let response = state.ollama.completions(body).await.map_err(ollama_err)?;
|
||||
// Ok(Sse::new(stream)
|
||||
// .keep_alive(KeepAlive::default())
|
||||
// .into_response())
|
||||
// } else {
|
||||
// let response = state.ollama.completions(body).await.map_err(ollama_err)?;
|
||||
|
||||
Ok(Json(response).into_response())
|
||||
}
|
||||
}
|
||||
// Ok(Json(response).into_response())
|
||||
// }
|
||||
// }
|
||||
|
||||
pub async fn chat_completions(
|
||||
State(state): State<AppState>,
|
||||
Json(body): Json<Value>,
|
||||
) -> Result<Response, (axum::http::StatusCode, String)> {
|
||||
let wants_stream = body
|
||||
.get("stream")
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(false);
|
||||
// pub async fn chat_completions(
|
||||
// State(state): State<AppState>,
|
||||
// Json(body): Json<Value>,
|
||||
// ) -> Result<Response, (axum::http::StatusCode, String)> {
|
||||
// let wants_stream = body
|
||||
// .get("stream")
|
||||
// .and_then(|v| v.as_bool())
|
||||
// .unwrap_or(false);
|
||||
|
||||
if wants_stream {
|
||||
let stream = state
|
||||
.ollama
|
||||
.chat_completions_stream(body)
|
||||
.await
|
||||
.map_err(ollama_err)?;
|
||||
// if wants_stream {
|
||||
// let stream = state
|
||||
// .ollama
|
||||
// .chat_completions_stream(body)
|
||||
// .await
|
||||
// .map_err(ollama_err)?;
|
||||
|
||||
Ok(Sse::new(stream)
|
||||
.keep_alive(KeepAlive::default())
|
||||
.into_response())
|
||||
} else {
|
||||
let response = state
|
||||
.ollama
|
||||
.chat_completions(body)
|
||||
.await
|
||||
.map_err(ollama_err)?;
|
||||
// Ok(Sse::new(stream)
|
||||
// .keep_alive(KeepAlive::default())
|
||||
// .into_response())
|
||||
// } else {
|
||||
// let response = state
|
||||
// .ollama
|
||||
// .chat_completions(body)
|
||||
// .await
|
||||
// .map_err(ollama_err)?;
|
||||
|
||||
Ok(Json(response).into_response())
|
||||
}
|
||||
}
|
||||
// Ok(Json(response).into_response())
|
||||
// }
|
||||
// }
|
||||
|
||||
fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) {
|
||||
match e {
|
||||
|
||||
+10
-11
@@ -6,21 +6,20 @@ use crate::auth::middleware::auth_middleware;
|
||||
use crate::state::app_state::AppState;
|
||||
use axum::{Router, middleware, routing::get, routing::post};
|
||||
|
||||
fn public_router() -> Router<AppState> {
|
||||
Router::new().route("/openapi.json", get(openapi::openapi_json))
|
||||
}
|
||||
// fn public_router() -> Router<AppState> {
|
||||
// Router::new().route("/openapi.json", get(openapi::openapi_json))
|
||||
// }
|
||||
|
||||
pub fn protected_router() -> Router<AppState> {
|
||||
Router::new()
|
||||
.route("/models", get(models::list_models))
|
||||
.route("/completions", post(chat::completions))
|
||||
.route("/chat/completions", post(chat::chat_completions))
|
||||
.route("/models/{model}/load", post(models::load_model))
|
||||
.route("/models/{model}/unload", post(models::unload_model))
|
||||
Router::new().route("/models", get(models::list_models))
|
||||
// .route("/completions", post(chat::completions))
|
||||
// .route("/chat/completions", post(chat::chat_completions))
|
||||
.route("/models/{model}/load", post(models::load_model))
|
||||
.route("/models/{model}/unload", post(models::unload_model))
|
||||
}
|
||||
|
||||
pub fn router() -> Router<AppState> {
|
||||
Router::new()
|
||||
.merge(public_router())
|
||||
.merge(protected_router().layer(middleware::from_fn(auth_middleware)))
|
||||
// .merge(public_router())
|
||||
.merge(protected_router()) //.layer(middleware::from_fn(auth_middleware)))
|
||||
}
|
||||
|
||||
+22
-22
@@ -2,19 +2,13 @@ use axum::{
|
||||
Json,
|
||||
extract::{Path, State},
|
||||
};
|
||||
use serde::Deserialize;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::dto::api;
|
||||
use crate::{errors::OllamaError, state::app_state::AppState};
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub struct LoadModelBody {
|
||||
pub keep_alive: Option<String>,
|
||||
}
|
||||
|
||||
pub async fn list_models(
|
||||
State(state): State<AppState>,
|
||||
) -> Result<Json<Value>, (axum::http::StatusCode, String)> {
|
||||
) -> Result<Json<api::ModelsResponse>, (axum::http::StatusCode, String)> {
|
||||
match state.ollama.list_models().await {
|
||||
Ok(models) => Ok(Json(models)),
|
||||
Err(e) => Err(ollama_err(e)),
|
||||
@@ -24,27 +18,33 @@ pub async fn list_models(
|
||||
pub async fn load_model(
|
||||
State(state): State<AppState>,
|
||||
Path(model): Path<String>,
|
||||
Json(body): Json<LoadModelBody>,
|
||||
) -> Result<Json<Value>, (axum::http::StatusCode, String)> {
|
||||
match state
|
||||
Json(body): Json<api::LoadModelBody>,
|
||||
) -> Result<Json<api::LoadModelResponse>, (axum::http::StatusCode, String)> {
|
||||
let keep_alive = body.keep_alive.as_deref().unwrap_or("5m");
|
||||
|
||||
api::LoadModelBody::parse_keep_alive(keep_alive).map_err(ollama_err)?;
|
||||
|
||||
let response = state
|
||||
.ollama
|
||||
.load_model(&model, body.keep_alive.as_deref())
|
||||
.load_model(&model, keep_alive)
|
||||
.await
|
||||
{
|
||||
Ok(response) => Ok(Json(response)),
|
||||
Err(e) => Err(ollama_err(e)),
|
||||
}
|
||||
.map_err(ollama_err)?;
|
||||
|
||||
Ok(Json(response))
|
||||
}
|
||||
|
||||
pub async fn unload_model(
|
||||
State(state): State<AppState>,
|
||||
Path(model): Path<String>,
|
||||
) -> Result<Json<Value>, (axum::http::StatusCode, String)> {
|
||||
match state.ollama.unload_model(&model).await {
|
||||
// ← correct method
|
||||
Ok(response) => Ok(Json(response)),
|
||||
Err(e) => Err(ollama_err(e)),
|
||||
}
|
||||
) -> Result<Json<api::UnloadModelResponse>, (axum::http::StatusCode, String)> {
|
||||
|
||||
let response = state
|
||||
.ollama
|
||||
.unload_model(&model)
|
||||
.await
|
||||
.map_err(ollama_err)?;
|
||||
|
||||
Ok(Json(response))
|
||||
}
|
||||
|
||||
fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) {
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
use axum::Json;
|
||||
use utoipa::OpenApi;
|
||||
// use axum::Json;
|
||||
// use utoipa::OpenApi;
|
||||
|
||||
use crate::openapi::V1ApiDoc;
|
||||
// use crate::openapi::V1ApiDoc;
|
||||
|
||||
pub async fn openapi_json() -> Json<utoipa::openapi::OpenApi> {
|
||||
Json(V1ApiDoc::openapi())
|
||||
}
|
||||
// pub async fn openapi_json() -> Json<utoipa::openapi::OpenApi> {
|
||||
// Json(V1ApiDoc::openapi())
|
||||
// }
|
||||
|
||||
Reference in New Issue
Block a user