feat: chat typing

This commit is contained in:
2026-04-10 21:33:52 +02:00
parent d5856557b4
commit e31cdf131f
6 changed files with 277 additions and 221 deletions
+23 -31
View File
@@ -6,7 +6,6 @@ use axum::{
sse::{KeepAlive, Sse},
},
};
use serde_json::Value;
use crate::dto::api;
use crate::errors::OllamaError;
@@ -23,9 +22,7 @@ pub async fn completions(
State(state): State<AppState>,
Json(body): Json<api::CompletionRequest>,
) -> Result<Response, (axum::http::StatusCode, String)> {
let wants_stream = body.base.stream;
if wants_stream {
if body.base.stream {
let stream = state
.ollama
.completions_stream(&body)
@@ -42,35 +39,30 @@ pub async fn completions(
}
}
// 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<api::ChatRequest>,
) -> Result<Response, (axum::http::StatusCode, String)> {
if body.base.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 {
+2 -2
View File
@@ -14,7 +14,7 @@ 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("/chat/completions", post(chat::chat_completions))
.route("/models/{model}/load", post(models::load_model))
.route("/models/{model}/unload", post(models::unload_model))
}
@@ -22,5 +22,5 @@ pub fn protected_router() -> Router<AppState> {
pub fn router() -> Router<AppState> {
Router::new()
// .merge(public_router())
.merge(protected_router()) //.layer(middleware::from_fn(auth_middleware)))
.merge(protected_router().layer(middleware::from_fn(auth_middleware)))
}