Compare commits
5
Commits
77bd729d95
..
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bdbd9413fb | ||
|
|
f1b8c310c4 | ||
|
|
387e0a0cfb | ||
|
|
67312cb7a4 | ||
|
|
63e57153dd |
Generated
+56
-19
@@ -42,6 +42,28 @@ dependencies = [
|
||||
"serde_json",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "async-stream"
|
||||
version = "0.3.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "0b5a71a6f37880a80d1d7f19efd781e4b5de42c88f0722cc13bcb6cc2cfe8476"
|
||||
dependencies = [
|
||||
"async-stream-impl",
|
||||
"futures-core",
|
||||
"pin-project-lite",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "async-stream-impl"
|
||||
version = "0.3.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c7c24de15d275a1ecfd47a380fb4d5ec9bfe0933f309ed5e705b775596a3574d"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "atoi"
|
||||
version = "2.0.0"
|
||||
@@ -241,6 +263,7 @@ checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724"
|
||||
name = "chat"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"async-stream",
|
||||
"axum",
|
||||
"base64",
|
||||
"chrono",
|
||||
@@ -257,7 +280,7 @@ dependencies = [
|
||||
"thiserror 2.0.18",
|
||||
"tokio",
|
||||
"tokio-stream",
|
||||
"tower-http",
|
||||
"tower-http 0.7.0",
|
||||
"tracing",
|
||||
"tracing-subscriber",
|
||||
"utoipa",
|
||||
@@ -267,9 +290,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "chrono"
|
||||
version = "0.4.44"
|
||||
version = "0.4.45"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c673075a2e0e5f4a1dde27ce9dee1ea4558c7ffe648f576438a20ca1d2acc4b0"
|
||||
checksum = "1aa79e62e7697b8e29b513a68abacf485adcd1fe8284a4316c5ae868e6633327"
|
||||
dependencies = [
|
||||
"iana-time-zone",
|
||||
"js-sys",
|
||||
@@ -1222,14 +1245,14 @@ checksum = "b6d2cec3eae94f9f509c767b45932f1ada8350c4bdb85af2fcab4a3c14807981"
|
||||
|
||||
[[package]]
|
||||
name = "libredox"
|
||||
version = "0.1.16"
|
||||
version = "0.1.18"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e02f3bb43d335493c96bf3fd3a321600bf6bd07ed34bc64118e9293bdffea46c"
|
||||
checksum = "c943259e342f1e06ff2da7a83eabdfe7f92ce10262688dbf1895ff0b3e6e4652"
|
||||
dependencies = [
|
||||
"bitflags",
|
||||
"libc",
|
||||
"plain",
|
||||
"redox_syscall 0.7.5",
|
||||
"redox_syscall 0.9.0",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -1369,11 +1392,10 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "num-iter"
|
||||
version = "0.1.45"
|
||||
version = "0.1.46"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1429034a0490724d0075ebb2bc9e875d6503c3cf69e235a8941aa757d83ef5bf"
|
||||
checksum = "c92800bd69a1eac91786bcfe9da64a897eb72911b8dc3095decbd07429e8048b"
|
||||
dependencies = [
|
||||
"autocfg",
|
||||
"num-integer",
|
||||
"num-traits",
|
||||
]
|
||||
@@ -1693,9 +1715,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "redox_syscall"
|
||||
version = "0.7.5"
|
||||
version = "0.9.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "4666a1a60d8412eab19d94f6d13dcc9cea0a5ef4fdf6a5db306537413c661b1b"
|
||||
checksum = "c5102a6aaa05aa011a238e178e6bca86d2cb56fc9f586d37cb80f5bca6e07759"
|
||||
dependencies = [
|
||||
"bitflags",
|
||||
]
|
||||
@@ -1763,7 +1785,7 @@ dependencies = [
|
||||
"tokio-rustls",
|
||||
"tokio-util",
|
||||
"tower",
|
||||
"tower-http",
|
||||
"tower-http 0.6.11",
|
||||
"tower-service",
|
||||
"url",
|
||||
"wasm-bindgen",
|
||||
@@ -2021,9 +2043,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "sha1"
|
||||
version = "0.10.6"
|
||||
version = "0.10.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba"
|
||||
checksum = "a978451301f4db1d02937a4ab3ccce137717b81826e79b7d49ffe3244a13c3b8"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cpufeatures 0.2.17",
|
||||
@@ -2617,6 +2639,21 @@ dependencies = [
|
||||
"url",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tower-http"
|
||||
version = "0.7.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b11f75e912b0c2be01b63d8cf8057b8c3f97cf34abb3d431a3a4c8675498e233"
|
||||
dependencies = [
|
||||
"bitflags",
|
||||
"bytes",
|
||||
"http",
|
||||
"percent-encoding",
|
||||
"pin-project-lite",
|
||||
"tower-layer",
|
||||
"tower-service",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tower-layer"
|
||||
version = "0.3.3"
|
||||
@@ -2793,9 +2830,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "uuid"
|
||||
version = "1.23.2"
|
||||
version = "1.23.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d258b83ceec21034727ecee8c382cfa6c3e133699b0742c64571814fb420c9f7"
|
||||
checksum = "ea5fab0d6c3c01ae70085a09cb03d4c7a1d6314e2b3e075392783396d724ca0a"
|
||||
dependencies = [
|
||||
"getrandom 0.4.2",
|
||||
"js-sys",
|
||||
@@ -3007,14 +3044,14 @@ version = "0.26.11"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "521bc38abb08001b01866da9f51eb7c5d647a19260e00054a8c7fd5f9e57f7a9"
|
||||
dependencies = [
|
||||
"webpki-roots 1.0.7",
|
||||
"webpki-roots 1.0.8",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "webpki-roots"
|
||||
version = "1.0.7"
|
||||
version = "1.0.8"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "52f5ee44c96cf55f1b349600768e3ece3a8f26010c05265ab73f945bb1a2eb9d"
|
||||
checksum = "bf85cb06032201fa7c6f829d7db5a7e5aa45bcc0655327713065f6f0576731bf"
|
||||
dependencies = [
|
||||
"rustls-pki-types",
|
||||
]
|
||||
|
||||
+5
-4
@@ -20,12 +20,13 @@ dotenvy = "0.15"
|
||||
thiserror = "2.0.18"
|
||||
tokio-stream = "0.1"
|
||||
futures = "0.3"
|
||||
chrono = { version = "0.4.44", features = ["serde"] }
|
||||
uuid = { version = "1.23.2", features = ["v4", "serde"] }
|
||||
tower-http = { version = "0.6.11", features = ["cors"] }
|
||||
chrono = { version = "0.4.45", features = ["serde"] }
|
||||
uuid = { version = "1.23.5", features = ["v4", "serde"] }
|
||||
tower-http = { version = "0.7.0", features = ["cors"] }
|
||||
tracing = "0.1.44"
|
||||
tracing-subscriber = { version = "0.3", features = ["env-filter"]}
|
||||
sqlx = { version = "0.8.6", features = ["runtime-tokio-rustls", "postgres", "uuid", "chrono", "macros"] }
|
||||
rand = "0.8"
|
||||
base64 = "0.22.1"
|
||||
sha2 = "0.11.0"
|
||||
sha2 = "0.11.0"
|
||||
async-stream = "0.3"
|
||||
+2
-2
@@ -19,14 +19,14 @@ pub async fn build_app() -> Router {
|
||||
.await
|
||||
.expect("Fatal error");
|
||||
let conversation_service = ConversationService::new(pool.clone());
|
||||
let api_key_service = AuthService::new(pool.clone());
|
||||
let auth_service = AuthService::new(pool.clone());
|
||||
|
||||
let ollama = OllamaProvider::new(OLLAMA_URL.as_str());
|
||||
let chat_service = ChatService::new(ollama, conversation_service.clone());
|
||||
|
||||
let state = Arc::new(api::state::AppState {
|
||||
conversation_service,
|
||||
auth_service: api_key_service,
|
||||
auth_service,
|
||||
chat_service,
|
||||
});
|
||||
|
||||
|
||||
+14
-8
@@ -15,33 +15,39 @@ use crate::api::routes;
|
||||
),
|
||||
),
|
||||
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,
|
||||
routes::v1::llm::list_models,
|
||||
routes::v1::llm::completions,
|
||||
routes::v1::llm::chat_completions,
|
||||
routes::v1::llm::load_model,
|
||||
routes::v1::llm::unload_model,
|
||||
),
|
||||
components(
|
||||
schemas(
|
||||
api::errors::ErrorResponse,
|
||||
|
||||
api::types::ApiModelsResponse,
|
||||
api::types::ApiModelInfo,
|
||||
api::types::ApiModelMetadata,
|
||||
|
||||
api::types::ApiLoadModelResponse,
|
||||
api::types::ApiLoadModelRequest,
|
||||
api::types::ApiUnloadModelResponse,
|
||||
|
||||
api::types::ApiLlmOptions,
|
||||
|
||||
api::types::ApiCompletionRequest,
|
||||
api::types::ApiCompletionObject,
|
||||
api::types::ApiFinishReason,
|
||||
api::types::ApiCompletionResponse,
|
||||
api::types::ApiCompletionObject,
|
||||
api::types::Choice,
|
||||
api::types::Usage,
|
||||
api::types::CompletionChunk,
|
||||
api::types::ApiFinishReason,
|
||||
|
||||
api::types::ApiChatRequest,
|
||||
api::types::ApiMessage,
|
||||
api::types::ApiRole,
|
||||
api::types::ApiChatResponse,
|
||||
api::types::ApiChatChoice,
|
||||
|
||||
api::types::ChatCompletionChunk,
|
||||
api::types::ChatChunkChoice,
|
||||
api::types::Delta,
|
||||
|
||||
+24
-30
@@ -1,5 +1,6 @@
|
||||
use crate::api::middlewares::errors::AuthMiddlewareError;
|
||||
use crate::databases::errors::DbError;
|
||||
use crate::providers::keycloak::errors::AuthError;
|
||||
use crate::providers::keycloak::errors::JwtValidationError;
|
||||
use crate::providers::ollama::errors::LlmError;
|
||||
use crate::services::errors::ServiceError;
|
||||
|
||||
@@ -7,7 +8,6 @@ use axum::Json;
|
||||
use axum::http::StatusCode;
|
||||
use axum::response::{IntoResponse, Response};
|
||||
use serde::Serialize;
|
||||
use thiserror::Error;
|
||||
use utoipa::ToSchema;
|
||||
|
||||
#[derive(Serialize, ToSchema)]
|
||||
@@ -25,18 +25,6 @@ pub struct ApiError {
|
||||
pub message: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum AuthMiddlewareError {
|
||||
#[error("invalid authorization format")]
|
||||
InvalidAuthorizationFormat,
|
||||
|
||||
#[error("authentication required")]
|
||||
AuthenticationRequired,
|
||||
|
||||
#[error("insufficient permissions")]
|
||||
Forbidden,
|
||||
}
|
||||
|
||||
impl From<AuthMiddlewareError> for ApiError {
|
||||
fn from(err: AuthMiddlewareError) -> Self {
|
||||
match err {
|
||||
@@ -104,35 +92,41 @@ impl From<ServiceError> for ApiError {
|
||||
},
|
||||
|
||||
ServiceError::Auth(e) => match e {
|
||||
AuthError::InvalidToken | AuthError::TokenValidationFailed => Self {
|
||||
status: StatusCode::UNAUTHORIZED,
|
||||
code: "AUTH_INVALID_TOKEN",
|
||||
message: "invalid or expired token".into(),
|
||||
},
|
||||
AuthError::InvalidHeader
|
||||
// | AuthError::InvalidHeaderDecode
|
||||
| AuthError::MissingKid => Self {
|
||||
JwtValidationError::InvalidToken | JwtValidationError::TokenValidationFailed => {
|
||||
Self {
|
||||
status: StatusCode::UNAUTHORIZED,
|
||||
code: "AUTH_INVALID_TOKEN",
|
||||
message: "invalid or expired token".into(),
|
||||
}
|
||||
}
|
||||
JwtValidationError::InvalidHeader | JwtValidationError::MissingKid => Self {
|
||||
status: StatusCode::UNAUTHORIZED,
|
||||
code: "AUTH_INVALID_HEADER",
|
||||
message: "invalid authorization header".into(),
|
||||
},
|
||||
AuthError::JwkNotFound
|
||||
| AuthError::InvalidJwks
|
||||
| AuthError::MissingModulus
|
||||
| AuthError::MissingExponent
|
||||
| AuthError::InvalidDecodingKey => Self {
|
||||
JwtValidationError::JwkNotFound
|
||||
| JwtValidationError::InvalidJwks
|
||||
| JwtValidationError::MissingModulus
|
||||
| JwtValidationError::MissingExponent
|
||||
| JwtValidationError::InvalidDecodingKey => Self {
|
||||
status: StatusCode::INTERNAL_SERVER_ERROR,
|
||||
code: "AUTH_JWKS_ERROR",
|
||||
message: "key validation error".into(),
|
||||
},
|
||||
AuthError::JwksFetchFailed
|
||||
| AuthError::JwksRefreshFailed
|
||||
| AuthError::Reqwest(_) => Self {
|
||||
JwtValidationError::JwksFetchFailed
|
||||
| JwtValidationError::JwksRefreshFailed
|
||||
| JwtValidationError::Reqwest(_) => Self {
|
||||
status: StatusCode::SERVICE_UNAVAILABLE,
|
||||
code: "AUTH_JWKS_FETCH",
|
||||
message: "failed to fetch authorization keys".into(),
|
||||
},
|
||||
},
|
||||
|
||||
ServiceError::Internal(e) => Self {
|
||||
status: StatusCode::INTERNAL_SERVER_ERROR,
|
||||
code: "INTERNAL_SERVER_ERROR",
|
||||
message: e,
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+54
-44
@@ -1,4 +1,5 @@
|
||||
use crate::api::errors::{ApiError, AuthMiddlewareError};
|
||||
use crate::api::errors::ApiError;
|
||||
use crate::api::middlewares::errors::AuthMiddlewareError;
|
||||
use crate::api::state::SharedState;
|
||||
|
||||
use axum::{
|
||||
@@ -19,9 +20,47 @@ enum ApiKeyError {
|
||||
Service(ApiError),
|
||||
}
|
||||
|
||||
fn resolve_auth_error(jwt_err: JwtError, api_key_err: ApiKeyError) -> ApiError {
|
||||
match (jwt_err, api_key_err) {
|
||||
// Both headers absent
|
||||
(JwtError::MissingHeader, ApiKeyError::MissingHeader) => {
|
||||
tracing::debug!("auth failed: no credentials provided");
|
||||
AuthMiddlewareError::AuthenticationRequired.into()
|
||||
}
|
||||
|
||||
// JWT header present but malformed — api key result irrelevant
|
||||
(JwtError::InvalidFormat, _) => {
|
||||
tracing::debug!("auth failed: malformed Authorization header");
|
||||
AuthMiddlewareError::InvalidAuthorizationFormat.into()
|
||||
}
|
||||
|
||||
// JWT service failure, api key not attempted
|
||||
(JwtError::Service(e), ApiKeyError::MissingHeader) => {
|
||||
tracing::warn!(error = ?e, "auth failed: jwt validation error");
|
||||
e
|
||||
}
|
||||
|
||||
// Both services failed
|
||||
(JwtError::Service(jwt_e), ApiKeyError::Service(api_e)) => {
|
||||
tracing::warn!(
|
||||
jwt_error = ?jwt_e,
|
||||
api_key_error = ?api_e,
|
||||
"auth failed: both jwt and api key validation errored"
|
||||
);
|
||||
jwt_e
|
||||
}
|
||||
|
||||
// JWT missing, api key service failed
|
||||
(JwtError::MissingHeader, ApiKeyError::Service(e)) => {
|
||||
tracing::warn!(error = ?e, "auth failed: api key validation error");
|
||||
e
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn auth_middleware(
|
||||
State(state): State<SharedState>,
|
||||
req: Request,
|
||||
mut req: Request,
|
||||
next: Next,
|
||||
) -> Response {
|
||||
let headers = req.headers();
|
||||
@@ -31,48 +70,27 @@ pub async fn auth_middleware(
|
||||
Err(jwt_err) => match try_api_key(&state, headers).await {
|
||||
Ok(auth) => auth,
|
||||
Err(api_key_err) => {
|
||||
let err = match (jwt_err, api_key_err) {
|
||||
// Both headers absent
|
||||
(JwtError::MissingHeader, ApiKeyError::MissingHeader) => {
|
||||
AuthMiddlewareError::AuthenticationRequired
|
||||
}
|
||||
// JWT header present but malformed — surface it, api key result irrelevant
|
||||
(JwtError::InvalidFormat, _) => AuthMiddlewareError::InvalidAuthorizationFormat,
|
||||
// JWT service failure — api key header was missing, so JWT was the intended method
|
||||
(JwtError::Service(e), ApiKeyError::MissingHeader) => {
|
||||
return e.into_response();
|
||||
}
|
||||
// Both services failed
|
||||
(JwtError::Service(e), ApiKeyError::Service(_)) => {
|
||||
return e.into_response();
|
||||
}
|
||||
// JWT missing, api key service failed
|
||||
(JwtError::MissingHeader, ApiKeyError::Service(e)) => {
|
||||
return e.into_response();
|
||||
}
|
||||
};
|
||||
|
||||
return ApiError::from(err).into_response();
|
||||
return resolve_auth_error(jwt_err, api_key_err).into_response();
|
||||
}
|
||||
},
|
||||
};
|
||||
|
||||
match handle_auth(&state, req, next, auth).await {
|
||||
Ok(response) => {
|
||||
tracing::debug!("User authentified");
|
||||
response
|
||||
}
|
||||
Err(err) => {
|
||||
tracing::debug!("Error during authentification {:?}", err);
|
||||
err.into_response()
|
||||
}
|
||||
if let Err(err) = record_auth(&state, &auth).await {
|
||||
tracing::debug!("Error during authentification {:?}", err);
|
||||
return err.into_response();
|
||||
}
|
||||
|
||||
tracing::debug!("User authentified");
|
||||
req.extensions_mut().insert(auth);
|
||||
next.run(req).await
|
||||
}
|
||||
|
||||
async fn try_jwt(
|
||||
state: &SharedState,
|
||||
headers: &HeaderMap,
|
||||
) -> Result<crate::core::auth::Auth, JwtError> {
|
||||
println!("{:?}", headers);
|
||||
|
||||
let token = headers
|
||||
.get("authorization")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
@@ -107,13 +125,8 @@ async fn try_api_key(
|
||||
Ok(crate::core::auth::Auth::ApiKey(auth))
|
||||
}
|
||||
|
||||
async fn handle_auth(
|
||||
state: &SharedState,
|
||||
mut request: Request,
|
||||
next: Next,
|
||||
auth: crate::core::auth::Auth,
|
||||
) -> Result<Response, ApiError> {
|
||||
match &auth {
|
||||
async fn record_auth(state: &SharedState, auth: &crate::core::auth::Auth) -> Result<(), ApiError> {
|
||||
match auth {
|
||||
crate::core::auth::Auth::Jwt(_) => {
|
||||
state.auth_service.create_user(&auth.user_id()).await?;
|
||||
}
|
||||
@@ -124,10 +137,7 @@ async fn handle_auth(
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
|
||||
request.extensions_mut().insert(auth);
|
||||
|
||||
Ok(next.run(request).await)
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ── Role guard ───────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
use thiserror::Error;
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum AuthMiddlewareError {
|
||||
#[error("invalid authorization format")]
|
||||
InvalidAuthorizationFormat,
|
||||
|
||||
#[error("authentication required")]
|
||||
AuthenticationRequired,
|
||||
|
||||
#[error("insufficient permissions")]
|
||||
Forbidden,
|
||||
}
|
||||
@@ -1 +1,2 @@
|
||||
pub mod auth;
|
||||
pub mod errors;
|
||||
|
||||
@@ -15,9 +15,142 @@ use axum::{
|
||||
use futures::StreamExt;
|
||||
use uuid::Uuid;
|
||||
|
||||
#[utoipa::path(
|
||||
get,
|
||||
path = "/llm/models",
|
||||
tag = "models",
|
||||
responses(
|
||||
(
|
||||
status = 200,
|
||||
description = "List of locally available Ollama models",
|
||||
body = api::types::ApiModelsResponse,
|
||||
content_type = "application/json",
|
||||
),
|
||||
(
|
||||
status = 500,
|
||||
description = "Internal server error (Ollama or network failure)",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "connection refused" })
|
||||
)
|
||||
)
|
||||
)]
|
||||
pub async fn list_models(
|
||||
State(state): State<SharedState>,
|
||||
) -> Result<Json<api::types::ApiModelsResponse>, api::errors::ApiError> {
|
||||
let models = state.chat_service.list_models().await?;
|
||||
|
||||
Ok(Json(models.into()))
|
||||
}
|
||||
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/completions",
|
||||
path = "/llm/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::types::ApiLoadModelRequest,
|
||||
description = "Load model request",
|
||||
content_type = "application/json",
|
||||
example = json!({ "keep_alive": "10m" })
|
||||
),
|
||||
responses(
|
||||
(
|
||||
status = 200,
|
||||
description = "Model successfully loaded into memory",
|
||||
body = api::types::ApiLoadModelResponse,
|
||||
content_type = "application/json",
|
||||
),
|
||||
(
|
||||
status = 400,
|
||||
description = "Invalid or missing keep_alive format",
|
||||
body = api::errors::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::errors::ErrorResponse,
|
||||
example = json!({ "error": "model 'llama3' not found — run `ollama pull llama3`" })
|
||||
),
|
||||
(
|
||||
status = 500,
|
||||
description = "Internal server error (Ollama or network failure)",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "connection refused" })
|
||||
)
|
||||
)
|
||||
)]
|
||||
pub async fn load_model(
|
||||
State(state): State<SharedState>,
|
||||
Path(model): Path<String>,
|
||||
Json(body): Json<api::types::ApiLoadModelRequest>,
|
||||
) -> Result<Json<api::types::ApiLoadModelResponse>, api::errors::ApiError> {
|
||||
let response = state
|
||||
.chat_service
|
||||
.load_model(crate::core::llm::models::LoadModelRequest {
|
||||
model,
|
||||
keep_alive: body.keep_alive.clone(),
|
||||
})
|
||||
.await?;
|
||||
|
||||
Ok(Json(api::types::ApiLoadModelResponse {
|
||||
model: response.model,
|
||||
keep_alive: body.keep_alive,
|
||||
status: "loaded".to_string(),
|
||||
}))
|
||||
}
|
||||
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/llm/models/{model}/unload",
|
||||
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::types::ApiUnloadModelResponse,
|
||||
content_type = "application/json",
|
||||
),
|
||||
(
|
||||
status = 404,
|
||||
description = "Model not found locally",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "model 'llama3' not found — run `ollama pull llama3`" })
|
||||
),
|
||||
(
|
||||
status = 500,
|
||||
description = "Internal server error (Ollama or network failure)",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "connection refused" })
|
||||
)
|
||||
)
|
||||
)]
|
||||
pub async fn unload_model(
|
||||
State(state): State<SharedState>,
|
||||
Path(model): Path<String>,
|
||||
) -> Result<Json<api::types::ApiUnloadModelResponse>, api::errors::ApiError> {
|
||||
let response = state
|
||||
.chat_service
|
||||
.unload_model(crate::core::llm::models::UnloadModelRequest { model })
|
||||
.await?;
|
||||
|
||||
Ok(Json(api::types::ApiUnloadModelResponse {
|
||||
model: response.model,
|
||||
status: "unloaded".to_string(),
|
||||
}))
|
||||
}
|
||||
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/llm/completions",
|
||||
tag = "chat",
|
||||
request_body(
|
||||
content = api::types::ApiCompletionRequest,
|
||||
@@ -81,7 +214,7 @@ pub async fn completions(
|
||||
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/chat/completions",
|
||||
path = "/llm/chat/completions",
|
||||
tag = "chat",
|
||||
request_body(
|
||||
content = api::types::ApiChatRequest,
|
||||
@@ -147,31 +280,36 @@ pub async fn chat_completions(
|
||||
conversation_id,
|
||||
message_id,
|
||||
created_at,
|
||||
model,
|
||||
}) => {
|
||||
println!("RECEIVED STRATTT");
|
||||
let payload = serde_json::to_string(&api::types::StreamEvent::Start(
|
||||
api::types::StartEventData {
|
||||
conversation_id,
|
||||
created: created_at,
|
||||
id: message_id,
|
||||
model,
|
||||
},
|
||||
))
|
||||
.unwrap_or_default();
|
||||
Ok(Event::default().data(payload))
|
||||
}
|
||||
|
||||
Ok(crate::core::llm::chat::ChatCompletionStreamEvent::Token(tok)) => {
|
||||
Ok(crate::core::llm::chat::ChatCompletionStreamEvent::Token {
|
||||
content,
|
||||
id,
|
||||
..
|
||||
}) => {
|
||||
let chunk = api::types::ChatCompletionChunk {
|
||||
id: String::new(),
|
||||
id,
|
||||
object: "chat.completion.chunk".to_string(),
|
||||
choices: vec![api::types::ChatChunkChoice {
|
||||
index: 0,
|
||||
delta: api::types::Delta {
|
||||
content: Some(tok),
|
||||
role: None,
|
||||
content: Some(content),
|
||||
role: Some(api::types::ApiRole::Assistant),
|
||||
},
|
||||
finish_reason: None,
|
||||
}],
|
||||
usage: None,
|
||||
};
|
||||
let payload = serde_json::to_string(&api::types::StreamEvent::Delta(chunk))
|
||||
.unwrap_or_default();
|
||||
@@ -181,13 +319,13 @@ pub async fn chat_completions(
|
||||
Ok(crate::core::llm::chat::ChatCompletionStreamEvent::Final(res)) => {
|
||||
let payload = serde_json::to_string(&api::types::StreamEvent::End(
|
||||
api::types::EndEventData {
|
||||
created: res.created_at.parse().unwrap_or(0),
|
||||
id: res.id,
|
||||
usage: api::types::Usage {
|
||||
prompt_tokens: res.prompt_tokens,
|
||||
completion_tokens: res.completion_tokens,
|
||||
total_tokens: res.prompt_tokens + res.completion_tokens,
|
||||
},
|
||||
finish_reason: res.finish_reason.into(),
|
||||
},
|
||||
))
|
||||
.unwrap_or_default();
|
||||
+29
-24
@@ -1,6 +1,5 @@
|
||||
pub mod apikey;
|
||||
pub mod chat;
|
||||
pub mod models;
|
||||
pub mod llm;
|
||||
|
||||
use crate::api::docs::ApiDoc;
|
||||
use crate::api::middlewares::auth::auth_middleware;
|
||||
@@ -17,37 +16,43 @@ async fn openapi_json() -> Json<utoipa::openapi::OpenApi> {
|
||||
Json(ApiDoc::openapi())
|
||||
}
|
||||
|
||||
fn public_router() -> Router<SharedState> {
|
||||
Router::new().route("/docs.json", get(openapi_json))
|
||||
fn llm_router() -> Router<SharedState> {
|
||||
Router::new()
|
||||
.route("/models", get(llm::list_models))
|
||||
.route("/completions", post(llm::completions))
|
||||
.route("/chat/completions", post(llm::chat_completions))
|
||||
.route("/models/{model}/load", post(llm::load_model))
|
||||
.route("/models/{model}/unload", post(llm::unload_model))
|
||||
}
|
||||
|
||||
pub fn protected_router() -> Router<SharedState> {
|
||||
fn keys_router() -> Router<SharedState> {
|
||||
Router::new().route(
|
||||
"/generate",
|
||||
post(apikey::create_api_key).route_layer(role_guard!(Some("admin"), None)),
|
||||
)
|
||||
}
|
||||
|
||||
fn log_router() -> Router<SharedState> {
|
||||
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))
|
||||
.route(
|
||||
"/keys/generate",
|
||||
post(apikey::create_api_key) // Usage
|
||||
.route_layer(role_guard!(Some("admin"), None)),
|
||||
)
|
||||
.route("/conversations", get(chat::get_conversations))
|
||||
.route("/conversations", get(llm::get_conversations))
|
||||
.route(
|
||||
"/conversations/{conversation_id}/messages",
|
||||
get(chat::get_messages),
|
||||
get(llm::get_messages),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn router(state: SharedState) -> Router<SharedState> {
|
||||
let protected = Router::new()
|
||||
.nest("/llm", llm_router())
|
||||
.nest("/keys", keys_router())
|
||||
.nest("/log", log_router())
|
||||
.route_layer(middleware::from_fn_with_state(
|
||||
state.clone(),
|
||||
auth_middleware,
|
||||
));
|
||||
|
||||
Router::new()
|
||||
.merge(public_router())
|
||||
.merge(
|
||||
protected_router().route_layer(middleware::from_fn_with_state(
|
||||
state.clone(),
|
||||
auth_middleware,
|
||||
)),
|
||||
)
|
||||
.route("/docs.json", get(openapi_json))
|
||||
.merge(protected)
|
||||
.with_state(state)
|
||||
}
|
||||
|
||||
@@ -1,138 +0,0 @@
|
||||
use crate::api::{self, state::SharedState};
|
||||
|
||||
use axum::{
|
||||
Json,
|
||||
extract::{Path, State},
|
||||
};
|
||||
|
||||
#[utoipa::path(
|
||||
get,
|
||||
path = "/models",
|
||||
tag = "models",
|
||||
responses(
|
||||
(
|
||||
status = 200,
|
||||
description = "List of locally available Ollama models",
|
||||
body = api::types::ApiModelsResponse,
|
||||
content_type = "application/json",
|
||||
),
|
||||
(
|
||||
status = 500,
|
||||
description = "Internal server error (Ollama or network failure)",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "connection refused" })
|
||||
)
|
||||
)
|
||||
)]
|
||||
// #[axum::debug_handler]
|
||||
pub async fn list_models(
|
||||
State(state): State<SharedState>,
|
||||
) -> Result<Json<api::types::ApiModelsResponse>, api::errors::ApiError> {
|
||||
let models = state.chat_service.list_models().await?;
|
||||
|
||||
Ok(Json(models.into()))
|
||||
}
|
||||
|
||||
#[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::types::ApiLoadModelRequest,
|
||||
description = "Load model request",
|
||||
content_type = "application/json",
|
||||
example = json!({ "keep_alive": "10m" })
|
||||
),
|
||||
responses(
|
||||
(
|
||||
status = 200,
|
||||
description = "Model successfully loaded into memory",
|
||||
body = api::types::ApiLoadModelResponse,
|
||||
content_type = "application/json",
|
||||
),
|
||||
(
|
||||
status = 400,
|
||||
description = "Invalid or missing keep_alive format",
|
||||
body = api::errors::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::errors::ErrorResponse,
|
||||
example = json!({ "error": "model 'llama3' not found — run `ollama pull llama3`" })
|
||||
),
|
||||
(
|
||||
status = 500,
|
||||
description = "Internal server error (Ollama or network failure)",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "connection refused" })
|
||||
)
|
||||
)
|
||||
)]
|
||||
pub async fn load_model(
|
||||
State(state): State<SharedState>,
|
||||
Path(model): Path<String>,
|
||||
Json(body): Json<api::types::ApiLoadModelRequest>,
|
||||
) -> Result<Json<api::types::ApiLoadModelResponse>, api::errors::ApiError> {
|
||||
let response = state
|
||||
.chat_service
|
||||
.load_model(crate::core::llm::models::LoadModelRequest {
|
||||
model,
|
||||
keep_alive: body.keep_alive.clone(),
|
||||
})
|
||||
.await?;
|
||||
|
||||
Ok(Json(api::types::ApiLoadModelResponse {
|
||||
model: response.model,
|
||||
keep_alive: body.keep_alive,
|
||||
status: "loaded".to_string(),
|
||||
}))
|
||||
}
|
||||
|
||||
// #[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::types::UnloadModelResponse,
|
||||
// content_type = "application/json",
|
||||
// ),
|
||||
// (
|
||||
// status = 404,
|
||||
// description = "Model not found locally",
|
||||
// body = api::errors::ErrorResponse,
|
||||
// example = json!({ "error": "model 'llama3' not found — run `ollama pull llama3`" })
|
||||
// ),
|
||||
// (
|
||||
// status = 500,
|
||||
// description = "Internal server error (Ollama or network failure)",
|
||||
// body = api::errors::ErrorResponse,
|
||||
// example = json!({ "error": "connection refused" })
|
||||
// )
|
||||
// )
|
||||
// )]
|
||||
// pub async fn unload_model(
|
||||
// State(state): State<AppState>,
|
||||
// Path(model): Path<String>,
|
||||
// ) -> Result<Json<api::types::UnloadModelResponse>, (axum::http::StatusCode, String)> {
|
||||
// let response = state
|
||||
// .ollama
|
||||
// .unload_model(&model)
|
||||
// .await
|
||||
// .map_err(into_http_response)?;
|
||||
|
||||
// Ok(Json(response))
|
||||
// }
|
||||
@@ -0,0 +1,574 @@
|
||||
use axum::{
|
||||
body::Body,
|
||||
extract::{Multipart, Path, Query, State},
|
||||
http::header,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use utoipa::ToSchema;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::api;
|
||||
use crate::api::errors::ApiError;
|
||||
use crate::SharedState;
|
||||
|
||||
// ---------------------------------------------------------------------
|
||||
// Types
|
||||
// ---------------------------------------------------------------------
|
||||
|
||||
#[derive(Debug, Clone, Copy, Serialize, Deserialize, ToSchema)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum AudioFormat {
|
||||
Mp3,
|
||||
Wav,
|
||||
Ogg,
|
||||
Pcm,
|
||||
}
|
||||
|
||||
impl AudioFormat {
|
||||
pub fn content_type(&self) -> &'static str {
|
||||
match self {
|
||||
AudioFormat::Mp3 => "audio/mpeg",
|
||||
AudioFormat::Wav => "audio/wav",
|
||||
AudioFormat::Ogg => "audio/ogg",
|
||||
AudioFormat::Pcm => "audio/L16",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize, ToSchema)]
|
||||
pub struct ApiSpeechRequest {
|
||||
/// Text to synthesize.
|
||||
#[schema(example = "Hello there, how can I help you today?")]
|
||||
pub text: String,
|
||||
/// Voice id, as returned by `GET /v1/voices`.
|
||||
#[schema(example = "voice_en_us_amy")]
|
||||
pub voice: String,
|
||||
/// TTS model/engine to use. Defaults to the server's default model.
|
||||
#[schema(example = "xtts-v2")]
|
||||
pub model: Option<String>,
|
||||
/// BCP-47 language code. Defaults to the voice's native language.
|
||||
#[schema(example = "en-US")]
|
||||
pub language: Option<String>,
|
||||
/// Output audio format. Defaults to mp3.
|
||||
pub format: Option<AudioFormat>,
|
||||
/// Playback speed multiplier (0.5-2.0). Defaults to 1.0.
|
||||
#[schema(example = 1.0)]
|
||||
pub speed: Option<f32>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, ToSchema)]
|
||||
pub struct ApiVoice {
|
||||
pub id: String,
|
||||
pub name: String,
|
||||
/// BCP-47 language code, e.g. "en-US".
|
||||
pub language: String,
|
||||
pub sample_rate: u32,
|
||||
/// URL to a short preview clip, if available.
|
||||
pub preview_url: Option<String>,
|
||||
pub is_cloned: bool,
|
||||
pub created_at: chrono::DateTime<chrono::Utc>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, ToSchema)]
|
||||
pub struct ApiVoiceListResponse {
|
||||
pub voices: Vec<ApiVoice>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize, ToSchema)]
|
||||
pub struct ApiRegisterVoiceRequest {
|
||||
/// Display name for the new voice.
|
||||
#[schema(example = "My Cloned Voice")]
|
||||
pub name: String,
|
||||
/// BCP-47 language code for the voice.
|
||||
#[schema(example = "en-US")]
|
||||
pub language: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, ToSchema)]
|
||||
pub struct ApiRegisterVoiceResponse {
|
||||
pub voice: ApiVoice,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, ToSchema)]
|
||||
pub struct ApiDeleteVoiceResponse {
|
||||
pub id: String,
|
||||
pub deleted: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, ToSchema)]
|
||||
pub struct ApiLanguage {
|
||||
/// BCP-47 language code, e.g. "en-US".
|
||||
pub code: String,
|
||||
pub name: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, ToSchema)]
|
||||
pub struct ApiLanguageListResponse {
|
||||
pub languages: Vec<ApiLanguage>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, ToSchema)]
|
||||
pub struct ApiTtsModel {
|
||||
pub id: String,
|
||||
pub name: String,
|
||||
pub description: Option<String>,
|
||||
pub supported_languages: Vec<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, ToSchema)]
|
||||
pub struct ApiTtsModelListResponse {
|
||||
pub models: Vec<ApiTtsModel>,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------
|
||||
// POST /v1/audio/speech
|
||||
// ---------------------------------------------------------------------
|
||||
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/v1/audio/speech",
|
||||
tag = "audio",
|
||||
request_body(
|
||||
content = ApiSpeechRequest,
|
||||
description = "Speech synthesis request",
|
||||
content_type = "application/json",
|
||||
example = json!({ "text": "Hello there!", "voice": "voice_en_us_amy", "format": "mp3" })
|
||||
),
|
||||
responses(
|
||||
(
|
||||
status = 200,
|
||||
description = "Full audio file generated from the input text",
|
||||
content_type = "audio/mpeg",
|
||||
),
|
||||
(
|
||||
status = 400,
|
||||
description = "Invalid request (e.g. empty text, unsupported speed)",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "text must not be empty" })
|
||||
),
|
||||
(
|
||||
status = 404,
|
||||
description = "Voice or model not found",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "voice 'voice_en_us_amy' not found" })
|
||||
),
|
||||
(
|
||||
status = 500,
|
||||
description = "Internal server error during synthesis",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "synthesis engine crashed" })
|
||||
)
|
||||
)
|
||||
)]
|
||||
pub async fn generate_speech(
|
||||
State(state): State<SharedState>,
|
||||
Json(body): Json<ApiSpeechRequest>,
|
||||
) -> Result<Response, ApiError> {
|
||||
let audio = state
|
||||
.tts_service
|
||||
.synthesize(crate::core::tts::SynthesizeRequest {
|
||||
text: body.text,
|
||||
voice: body.voice,
|
||||
model: body.model,
|
||||
language: body.language,
|
||||
format: body.format.unwrap_or(AudioFormat::Mp3),
|
||||
speed: body.speed.unwrap_or(1.0),
|
||||
})
|
||||
.await?;
|
||||
|
||||
let content_type = audio.format.content_type();
|
||||
Ok((
|
||||
[(header::CONTENT_TYPE, content_type)],
|
||||
Body::from(audio.bytes),
|
||||
)
|
||||
.into_response())
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------
|
||||
// POST /v1/audio/speech/stream
|
||||
// ---------------------------------------------------------------------
|
||||
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/v1/audio/speech/stream",
|
||||
tag = "audio",
|
||||
request_body(
|
||||
content = ApiSpeechRequest,
|
||||
description = "Speech synthesis request, streamed back as audio is generated",
|
||||
content_type = "application/json",
|
||||
example = json!({ "text": "Hello there!", "voice": "voice_en_us_amy", "format": "mp3" })
|
||||
),
|
||||
responses(
|
||||
(
|
||||
status = 200,
|
||||
description = "Chunked audio stream (chunk-transfer-encoded); the same audio the sync endpoint returns, sent incrementally",
|
||||
content_type = "audio/mpeg",
|
||||
),
|
||||
(
|
||||
status = 400,
|
||||
description = "Invalid request",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "text must not be empty" })
|
||||
),
|
||||
(
|
||||
status = 404,
|
||||
description = "Voice or model not found",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "voice 'voice_en_us_amy' not found" })
|
||||
),
|
||||
(
|
||||
status = 500,
|
||||
description = "Internal server error during synthesis",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "synthesis engine crashed" })
|
||||
)
|
||||
)
|
||||
)]
|
||||
pub async fn generate_speech_stream(
|
||||
State(state): State<SharedState>,
|
||||
Json(body): Json<ApiSpeechRequest>,
|
||||
) -> Result<Response, ApiError> {
|
||||
let format = body.format.unwrap_or(AudioFormat::Mp3);
|
||||
|
||||
let stream = state
|
||||
.tts_service
|
||||
.synthesize_stream(crate::core::tts::SynthesizeRequest {
|
||||
text: body.text,
|
||||
voice: body.voice,
|
||||
model: body.model,
|
||||
language: body.language,
|
||||
format,
|
||||
speed: body.speed.unwrap_or(1.0),
|
||||
})
|
||||
.await?;
|
||||
|
||||
Ok((
|
||||
[(header::CONTENT_TYPE, format.content_type())],
|
||||
Body::from_stream(stream),
|
||||
)
|
||||
.into_response())
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------
|
||||
// GET /v1/voices
|
||||
// ---------------------------------------------------------------------
|
||||
|
||||
#[derive(Debug, Deserialize, ToSchema)]
|
||||
pub struct ListVoicesQuery {
|
||||
/// Optional BCP-47 language filter, e.g. "en-US".
|
||||
pub language: Option<String>,
|
||||
}
|
||||
|
||||
#[utoipa::path(
|
||||
get,
|
||||
path = "/v1/voices",
|
||||
tag = "voices",
|
||||
params(
|
||||
("language" = Option<String>, Query, description = "Filter voices by BCP-47 language code")
|
||||
),
|
||||
responses(
|
||||
(
|
||||
status = 200,
|
||||
description = "List of available voices",
|
||||
body = ApiVoiceListResponse,
|
||||
content_type = "application/json",
|
||||
),
|
||||
(
|
||||
status = 500,
|
||||
description = "Internal server error",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "failed to load voice registry" })
|
||||
)
|
||||
)
|
||||
)]
|
||||
pub async fn list_voices(
|
||||
State(state): State<SharedState>,
|
||||
Query(query): Query<ListVoicesQuery>,
|
||||
) -> Result<Json<ApiVoiceListResponse>, ApiError> {
|
||||
let voices = state.tts_service.list_voices(query.language).await?;
|
||||
Ok(Json(ApiVoiceListResponse {
|
||||
voices: voices.into_iter().map(Into::into).collect(),
|
||||
}))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------
|
||||
// POST /v1/voices (register / clone a voice)
|
||||
// ---------------------------------------------------------------------
|
||||
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/v1/voices",
|
||||
tag = "voices",
|
||||
request_body(
|
||||
content = ApiRegisterVoiceRequest,
|
||||
description = "Multipart form: JSON fields (name, language) plus an `audio_sample` file field containing the reference audio to clone",
|
||||
content_type = "multipart/form-data",
|
||||
),
|
||||
responses(
|
||||
(
|
||||
status = 200,
|
||||
description = "Voice registered/cloned successfully",
|
||||
body = ApiRegisterVoiceResponse,
|
||||
content_type = "application/json",
|
||||
),
|
||||
(
|
||||
status = 400,
|
||||
description = "Invalid request (missing name, missing/unsupported audio sample)",
|
||||
body = api::errors::ErrorResponse,
|
||||
examples(
|
||||
("Missing name" = (value = json!({ "error": "name is required and cannot be empty" }))),
|
||||
("Bad sample" = (value = json!({ "error": "audio_sample must be a wav or mp3 file under 30s" })))
|
||||
)
|
||||
),
|
||||
(
|
||||
status = 500,
|
||||
description = "Internal server error while cloning the voice",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "voice cloning engine failed" })
|
||||
)
|
||||
)
|
||||
)]
|
||||
pub async fn register_voice(
|
||||
State(state): State<SharedState>,
|
||||
mut multipart: Multipart,
|
||||
) -> Result<Json<ApiRegisterVoiceResponse>, ApiError> {
|
||||
let mut name: Option<String> = None;
|
||||
let mut language: Option<String> = None;
|
||||
let mut audio_sample: Option<Vec<u8>> = None;
|
||||
|
||||
while let Some(field) = multipart
|
||||
.next_field()
|
||||
.await
|
||||
.map_err(|e| ApiError::bad_request(format!("invalid multipart body: {e}")))?
|
||||
{
|
||||
match field.name() {
|
||||
Some("name") => {
|
||||
name = Some(field.text().await.map_err(|e| {
|
||||
ApiError::bad_request(format!("invalid name field: {e}"))
|
||||
})?);
|
||||
}
|
||||
Some("language") => {
|
||||
language = Some(field.text().await.map_err(|e| {
|
||||
ApiError::bad_request(format!("invalid language field: {e}"))
|
||||
})?);
|
||||
}
|
||||
Some("audio_sample") => {
|
||||
audio_sample = Some(
|
||||
field
|
||||
.bytes()
|
||||
.await
|
||||
.map_err(|e| {
|
||||
ApiError::bad_request(format!("invalid audio_sample field: {e}"))
|
||||
})?
|
||||
.to_vec(),
|
||||
);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
let name = name.ok_or_else(|| {
|
||||
ApiError::bad_request("name is required and cannot be empty".to_string())
|
||||
})?;
|
||||
let audio_sample = audio_sample.ok_or_else(|| {
|
||||
ApiError::bad_request("audio_sample is required".to_string())
|
||||
})?;
|
||||
|
||||
let voice = state
|
||||
.tts_service
|
||||
.register_voice(crate::core::tts::RegisterVoiceRequest {
|
||||
name,
|
||||
language,
|
||||
audio_sample,
|
||||
})
|
||||
.await?;
|
||||
|
||||
Ok(Json(ApiRegisterVoiceResponse {
|
||||
voice: voice.into(),
|
||||
}))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------
|
||||
// GET /v1/voices/{id}
|
||||
// ---------------------------------------------------------------------
|
||||
|
||||
#[utoipa::path(
|
||||
get,
|
||||
path = "/v1/voices/{id}",
|
||||
tag = "voices",
|
||||
params(
|
||||
("id" = String, Path, description = "Voice id, e.g. 'voice_en_us_amy'")
|
||||
),
|
||||
responses(
|
||||
(
|
||||
status = 200,
|
||||
description = "Voice metadata (language, sample rate, preview URL)",
|
||||
body = ApiVoice,
|
||||
content_type = "application/json",
|
||||
),
|
||||
(
|
||||
status = 404,
|
||||
description = "Voice not found",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "voice 'voice_en_us_amy' not found" })
|
||||
),
|
||||
(
|
||||
status = 500,
|
||||
description = "Internal server error",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "failed to load voice registry" })
|
||||
)
|
||||
)
|
||||
)]
|
||||
pub async fn get_voice(
|
||||
State(state): State<SharedState>,
|
||||
Path(id): Path<String>,
|
||||
) -> Result<Json<ApiVoice>, ApiError> {
|
||||
let voice = state.tts_service.get_voice(&id).await?;
|
||||
Ok(Json(voice.into()))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------
|
||||
// DELETE /v1/voices/{id}
|
||||
// ---------------------------------------------------------------------
|
||||
|
||||
#[utoipa::path(
|
||||
delete,
|
||||
path = "/v1/voices/{id}",
|
||||
tag = "voices",
|
||||
params(
|
||||
("id" = String, Path, description = "Voice id to remove")
|
||||
),
|
||||
responses(
|
||||
(
|
||||
status = 200,
|
||||
description = "Voice removed successfully",
|
||||
body = ApiDeleteVoiceResponse,
|
||||
content_type = "application/json",
|
||||
),
|
||||
(
|
||||
status = 404,
|
||||
description = "Voice not found",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "voice 'voice_en_us_amy' not found" })
|
||||
),
|
||||
(
|
||||
status = 400,
|
||||
description = "Voice is a built-in voice and cannot be deleted",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "built-in voices cannot be deleted" })
|
||||
),
|
||||
(
|
||||
status = 500,
|
||||
description = "Internal server error",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "failed to update voice registry" })
|
||||
)
|
||||
)
|
||||
)]
|
||||
pub async fn delete_voice(
|
||||
State(state): State<SharedState>,
|
||||
Path(id): Path<String>,
|
||||
) -> Result<Json<ApiDeleteVoiceResponse>, ApiError> {
|
||||
state.tts_service.delete_voice(&id).await?;
|
||||
Ok(Json(ApiDeleteVoiceResponse {
|
||||
id,
|
||||
deleted: true,
|
||||
}))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------
|
||||
// GET /v1/languages
|
||||
// ---------------------------------------------------------------------
|
||||
|
||||
#[utoipa::path(
|
||||
get,
|
||||
path = "/v1/languages",
|
||||
tag = "languages",
|
||||
responses(
|
||||
(
|
||||
status = 200,
|
||||
description = "List of supported languages",
|
||||
body = ApiLanguageListResponse,
|
||||
content_type = "application/json",
|
||||
),
|
||||
(
|
||||
status = 500,
|
||||
description = "Internal server error",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "failed to load language list" })
|
||||
)
|
||||
)
|
||||
)]
|
||||
pub async fn list_languages(
|
||||
State(state): State<SharedState>,
|
||||
) -> Result<Json<ApiLanguageListResponse>, ApiError> {
|
||||
let languages = state.tts_service.list_languages().await?;
|
||||
Ok(Json(ApiLanguageListResponse {
|
||||
languages: languages
|
||||
.into_iter()
|
||||
.map(|l| ApiLanguage {
|
||||
code: l.code,
|
||||
name: l.name,
|
||||
})
|
||||
.collect(),
|
||||
}))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------
|
||||
// GET /v1/models
|
||||
// ---------------------------------------------------------------------
|
||||
|
||||
#[utoipa::path(
|
||||
get,
|
||||
path = "/v1/models",
|
||||
tag = "models",
|
||||
responses(
|
||||
(
|
||||
status = 200,
|
||||
description = "List of available TTS models/engines",
|
||||
body = ApiTtsModelListResponse,
|
||||
content_type = "application/json",
|
||||
),
|
||||
(
|
||||
status = 500,
|
||||
description = "Internal server error",
|
||||
body = api::errors::ErrorResponse,
|
||||
example = json!({ "error": "failed to enumerate models" })
|
||||
)
|
||||
)
|
||||
)]
|
||||
pub async fn list_models(
|
||||
State(state): State<SharedState>,
|
||||
) -> Result<Json<ApiTtsModelListResponse>, ApiError> {
|
||||
let models = state.tts_service.list_models().await?;
|
||||
Ok(Json(ApiTtsModelListResponse {
|
||||
models: models
|
||||
.into_iter()
|
||||
.map(|m| ApiTtsModel {
|
||||
id: m.id,
|
||||
name: m.name,
|
||||
description: m.description,
|
||||
supported_languages: m.supported_languages,
|
||||
})
|
||||
.collect(),
|
||||
}))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------
|
||||
// Router wiring (example)
|
||||
// ---------------------------------------------------------------------
|
||||
|
||||
pub fn router() -> axum::Router<SharedState> {
|
||||
use axum::routing::{delete, get, post};
|
||||
|
||||
axum::Router::new()
|
||||
.route("/v1/audio/speech", post(generate_speech))
|
||||
.route("/v1/audio/speech/stream", post(generate_speech_stream))
|
||||
.route("/v1/voices", get(list_voices).post(register_voice))
|
||||
.route("/v1/voices/{id}", get(get_voice).delete(delete_voice))
|
||||
.route("/v1/languages", get(list_languages))
|
||||
.route("/v1/models", get(list_models))
|
||||
}
|
||||
+4
-14
@@ -105,8 +105,6 @@ pub enum ApiCompletionObject {
|
||||
pub enum ApiFinishReason {
|
||||
Stop,
|
||||
Length,
|
||||
ContentFilter,
|
||||
ToolCalls,
|
||||
Error,
|
||||
}
|
||||
|
||||
@@ -124,13 +122,6 @@ pub struct Usage {
|
||||
pub total_tokens: u32,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||
pub struct CompletionChunk {
|
||||
pub id: String,
|
||||
pub object: String,
|
||||
pub choices: Vec<Choice>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize, Serialize, ToSchema)]
|
||||
pub struct ApiChatRequest {
|
||||
#[serde(flatten)]
|
||||
@@ -179,17 +170,15 @@ pub struct ApiChatChoice {
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||
pub struct ChatCompletionChunk {
|
||||
pub id: String,
|
||||
pub id: Uuid,
|
||||
pub object: String,
|
||||
pub choices: Vec<ChatChunkChoice>,
|
||||
pub usage: Option<Usage>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||
pub struct ChatChunkChoice {
|
||||
pub index: u32,
|
||||
pub delta: Delta,
|
||||
pub finish_reason: Option<ApiFinishReason>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||
@@ -201,15 +190,16 @@ pub struct Delta {
|
||||
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||
pub struct StartEventData {
|
||||
pub conversation_id: Uuid,
|
||||
pub created: u64,
|
||||
pub created: String,
|
||||
pub id: Uuid,
|
||||
pub model: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||
pub struct EndEventData {
|
||||
pub created: u64,
|
||||
pub id: Uuid,
|
||||
pub usage: Usage,
|
||||
pub finish_reason: ApiFinishReason,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||
|
||||
+15
-3
@@ -51,20 +51,32 @@ pub struct ChatCompletionResultNoStream {
|
||||
pub prompt_tokens: u32,
|
||||
pub completion_tokens: u32,
|
||||
|
||||
pub done_reason: String,
|
||||
pub finish_reason: FinishReason,
|
||||
|
||||
pub total_duration: u64,
|
||||
pub load_duration: u64,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Clone)]
|
||||
pub enum FinishReason {
|
||||
Stop,
|
||||
Length,
|
||||
Error,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub enum ChatCompletionStreamEvent {
|
||||
Start {
|
||||
conversation_id: Uuid,
|
||||
message_id: Uuid,
|
||||
created_at: u64,
|
||||
created_at: String,
|
||||
model: String,
|
||||
},
|
||||
Token {
|
||||
content: String,
|
||||
id: uuid::Uuid,
|
||||
created_at: String,
|
||||
},
|
||||
Token(String),
|
||||
Final(ChatCompletionResultNoStream),
|
||||
}
|
||||
|
||||
|
||||
@@ -34,3 +34,13 @@ pub struct LoadModelRequest {
|
||||
pub struct LoadModelResponse {
|
||||
pub model: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct UnloadModelRequest {
|
||||
pub model: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct UnloadModelResponse {
|
||||
pub model: String,
|
||||
}
|
||||
|
||||
@@ -134,3 +134,13 @@ impl From<core::llm::chat::ChatCompletionResultNoStream> for api::types::ApiChat
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<core::llm::chat::FinishReason> for api::types::ApiFinishReason {
|
||||
fn from(f: core::llm::chat::FinishReason) -> Self {
|
||||
match f {
|
||||
core::llm::chat::FinishReason::Stop => Self::Stop,
|
||||
core::llm::chat::FinishReason::Length => Self::Length,
|
||||
core::llm::chat::FinishReason::Error => Self::Error,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -48,3 +48,13 @@ impl From<ollama::types::OllamaMessage> for core::llm::chat::Message {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<ollama::types::OllamaFinishReason> for core::llm::chat::FinishReason {
|
||||
fn from(f: ollama::types::OllamaFinishReason) -> Self {
|
||||
match f {
|
||||
ollama::types::OllamaFinishReason::Stop => Self::Stop,
|
||||
ollama::types::OllamaFinishReason::Length => Self::Length,
|
||||
ollama::types::OllamaFinishReason::Error => Self::Error,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use thiserror::Error;
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum AuthError {
|
||||
pub enum JwtValidationError {
|
||||
#[error("invalid authorization header")]
|
||||
InvalidHeader,
|
||||
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
use super::errors::AuthError;
|
||||
use super::errors::JwtValidationError;
|
||||
|
||||
use once_cell::sync::Lazy;
|
||||
use serde_json::Value;
|
||||
@@ -17,7 +17,7 @@ static JWK_CACHE: Lazy<Arc<RwLock<Option<JwksCache>>>> = Lazy::new(|| Arc::new(R
|
||||
|
||||
static JWKS_URL: Lazy<String> = Lazy::new(|| env::var("JWKS_URL").expect("JWKS_URL not set"));
|
||||
|
||||
async fn fetch_jwks() -> Result<Value, AuthError> {
|
||||
async fn fetch_jwks() -> Result<Value, JwtValidationError> {
|
||||
let jwks = reqwest::get(JWKS_URL.as_str())
|
||||
.await?
|
||||
.json::<Value>()
|
||||
@@ -26,7 +26,7 @@ async fn fetch_jwks() -> Result<Value, AuthError> {
|
||||
Ok(jwks)
|
||||
}
|
||||
|
||||
pub async fn refresh_jwks() -> Result<Value, AuthError> {
|
||||
pub async fn refresh_jwks() -> Result<Value, JwtValidationError> {
|
||||
let jwks = fetch_jwks().await?;
|
||||
|
||||
let mut write = JWK_CACHE.write().await;
|
||||
@@ -39,7 +39,7 @@ pub async fn refresh_jwks() -> Result<Value, AuthError> {
|
||||
Ok(jwks)
|
||||
}
|
||||
|
||||
pub async fn get_jwks() -> Result<Value, AuthError> {
|
||||
pub async fn get_jwks() -> Result<Value, JwtValidationError> {
|
||||
let ttl = Duration::from_secs(3600); // 1 hour
|
||||
|
||||
{
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use super::claims::KeycloakClaims;
|
||||
use super::errors::AuthError;
|
||||
use super::errors::JwtValidationError;
|
||||
|
||||
use jsonwebtoken::{Algorithm, DecodingKey, Validation, decode, decode_header};
|
||||
use once_cell::sync::Lazy;
|
||||
@@ -8,26 +8,32 @@ use std::env;
|
||||
|
||||
static ISSUER: Lazy<String> = Lazy::new(|| env::var("ISSUER").expect("ISSUER not set"));
|
||||
|
||||
fn validate_token(token: &str, jwks: &Value) -> Result<KeycloakClaims, AuthError> {
|
||||
fn validate_token(token: &str, jwks: &Value) -> Result<KeycloakClaims, JwtValidationError> {
|
||||
// 1. Decode header
|
||||
let header = decode_header(token).map_err(|_| AuthError::InvalidHeader)?;
|
||||
let header = decode_header(token).map_err(|_| JwtValidationError::InvalidHeader)?;
|
||||
|
||||
let kid = header.kid.ok_or(AuthError::MissingKid)?;
|
||||
let kid = header.kid.ok_or(JwtValidationError::MissingKid)?;
|
||||
|
||||
// 2. Find matching key
|
||||
let keys = jwks["keys"].as_array().ok_or(AuthError::InvalidJwks)?;
|
||||
let keys = jwks["keys"]
|
||||
.as_array()
|
||||
.ok_or(JwtValidationError::InvalidJwks)?;
|
||||
|
||||
let key = keys
|
||||
.iter()
|
||||
.find(|k| k["kid"] == kid)
|
||||
.ok_or(AuthError::InvalidDecodingKey)?;
|
||||
.ok_or(JwtValidationError::InvalidDecodingKey)?;
|
||||
|
||||
// 3. Extract RSA components
|
||||
let n = key["n"].as_str().ok_or(AuthError::MissingModulus)?;
|
||||
let e = key["e"].as_str().ok_or(AuthError::MissingExponent)?;
|
||||
let n = key["n"]
|
||||
.as_str()
|
||||
.ok_or(JwtValidationError::MissingModulus)?;
|
||||
let e = key["e"]
|
||||
.as_str()
|
||||
.ok_or(JwtValidationError::MissingExponent)?;
|
||||
|
||||
let decoding_key =
|
||||
DecodingKey::from_rsa_components(n, e).map_err(|_| AuthError::JwkNotFound)?;
|
||||
DecodingKey::from_rsa_components(n, e).map_err(|_| JwtValidationError::JwkNotFound)?;
|
||||
|
||||
// 4. Setup validation rules
|
||||
let mut validation = Validation::new(Algorithm::RS256);
|
||||
@@ -39,15 +45,15 @@ fn validate_token(token: &str, jwks: &Value) -> Result<KeycloakClaims, AuthError
|
||||
|
||||
// 5. Decode & verify
|
||||
let token_data = decode::<KeycloakClaims>(token, &decoding_key, &validation)
|
||||
.map_err(|_| AuthError::TokenValidationFailed)?;
|
||||
.map_err(|_| JwtValidationError::TokenValidationFailed)?;
|
||||
|
||||
Ok(token_data.claims)
|
||||
}
|
||||
|
||||
pub async fn authenticate_jwt(token: &str) -> Result<KeycloakClaims, AuthError> {
|
||||
pub async fn authenticate_jwt(token: &str) -> Result<KeycloakClaims, JwtValidationError> {
|
||||
let jwks = super::jwks::get_jwks()
|
||||
.await
|
||||
.map_err(|_| AuthError::JwksFetchFailed)?;
|
||||
.map_err(|_| JwtValidationError::JwksFetchFailed)?;
|
||||
|
||||
match validate_token(token, &jwks) {
|
||||
Ok(claims) => Ok(claims),
|
||||
@@ -56,9 +62,9 @@ pub async fn authenticate_jwt(token: &str) -> Result<KeycloakClaims, AuthError>
|
||||
// one retry with refresh
|
||||
let fresh = super::jwks::refresh_jwks()
|
||||
.await
|
||||
.map_err(|_| AuthError::JwksRefreshFailed)?;
|
||||
.map_err(|_| JwtValidationError::JwksRefreshFailed)?;
|
||||
|
||||
validate_token(token, &fresh).map_err(|_| AuthError::InvalidToken)
|
||||
validate_token(token, &fresh).map_err(|_| JwtValidationError::InvalidToken)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -285,62 +285,70 @@ impl OllamaProvider {
|
||||
.await?
|
||||
.bytes_stream();
|
||||
|
||||
let stream = byte_stream
|
||||
.flat_map(|chunk_result| {
|
||||
let mut out: Vec<Result<super::types::OllamaChatStreamEvent, LlmError>> =
|
||||
Vec::new();
|
||||
let chunk = match chunk_result {
|
||||
Ok(b) => b,
|
||||
Err(e) => {
|
||||
tracing::debug!("Error: {:?}", e);
|
||||
out.push(Err(LlmError::Http(e)));
|
||||
return futures::stream::iter(out);
|
||||
}
|
||||
};
|
||||
let stream = byte_stream.flat_map(|chunk_result| {
|
||||
let mut events: Vec<Result<super::types::OllamaChatStreamEvent, LlmError>> = Vec::new();
|
||||
|
||||
tracing::debug!("Raw line: {:?}", std::str::from_utf8(&chunk));
|
||||
|
||||
for line in chunk.split(|&b| b == b'\n') {
|
||||
if line.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let parsed: ollama::types::OllamaChatResponse =
|
||||
match serde_json::from_slice(line) {
|
||||
Ok(v) => v,
|
||||
Err(_) => continue,
|
||||
};
|
||||
|
||||
tracing::debug!("Parsed: {:?}", parsed);
|
||||
|
||||
if !parsed.message.content.is_empty() {
|
||||
out.push(Ok(super::types::OllamaChatStreamEvent::Token(
|
||||
parsed.message.content.clone(),
|
||||
)));
|
||||
}
|
||||
|
||||
if parsed.done {
|
||||
out.push(Ok(super::types::OllamaChatStreamEvent::Final(parsed)));
|
||||
return futures::stream::iter(out);
|
||||
}
|
||||
let chunk = match chunk_result {
|
||||
Ok(b) => b,
|
||||
Err(e) => {
|
||||
tracing::debug!("Error: {:?}", e);
|
||||
events.push(Err(LlmError::Http(e)));
|
||||
return futures::stream::iter(events);
|
||||
}
|
||||
futures::stream::iter(out)
|
||||
})
|
||||
.scan(String::new(), |acc, event| {
|
||||
let result = match event {
|
||||
Ok(ollama::types::OllamaChatStreamEvent::Token(ref tok)) => {
|
||||
acc.push_str(tok);
|
||||
Some(event)
|
||||
}
|
||||
Ok(ollama::types::OllamaChatStreamEvent::Final(mut resp)) => {
|
||||
tracing::debug!("Acc: {:?}", acc);
|
||||
resp.message.content = std::mem::take(acc);
|
||||
Some(Ok(ollama::types::OllamaChatStreamEvent::Final(resp)))
|
||||
}
|
||||
Err(_) => Some(event),
|
||||
};
|
||||
futures::future::ready(result)
|
||||
});
|
||||
};
|
||||
|
||||
let text = match std::str::from_utf8(&chunk) {
|
||||
Ok(v) => v,
|
||||
Err(e) => {
|
||||
tracing::debug!("Invalid UTF8 from Ollama: {:?}", e);
|
||||
return futures::stream::iter(events);
|
||||
}
|
||||
};
|
||||
|
||||
for line in text.lines() {
|
||||
if line.trim().is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let parsed: ollama::types::OllamaChatStreamResponse =
|
||||
match serde_json::from_str(line) {
|
||||
Ok(v) => v,
|
||||
Err(e) => {
|
||||
tracing::debug!("Failed parsing Ollama line {:?}: {:?}", line, e);
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
tracing::debug!("Parsed: {:?}", parsed);
|
||||
|
||||
// End of generation
|
||||
if parsed.done {
|
||||
println!("final: {:?}", &parsed);
|
||||
events.push(Ok(super::types::OllamaChatStreamEvent::Final(
|
||||
super::types::OllamaChatResponse {
|
||||
model: parsed.model,
|
||||
created_at: parsed.created_at,
|
||||
message: parsed.message,
|
||||
done_reason: parsed
|
||||
.done_reason
|
||||
.unwrap_or(super::types::OllamaFinishReason::Error),
|
||||
total_duration: parsed.total_duration.unwrap_or(0),
|
||||
load_duration: parsed.load_duration.unwrap_or(0),
|
||||
prompt_eval_count: parsed.prompt_eval_count.unwrap_or(0),
|
||||
eval_count: parsed.eval_count.unwrap_or(0),
|
||||
},
|
||||
)));
|
||||
break;
|
||||
}
|
||||
|
||||
// Normal generated token
|
||||
if !parsed.message.content.is_empty() {
|
||||
events.push(Ok(super::types::OllamaChatStreamEvent::Token(parsed)));
|
||||
}
|
||||
}
|
||||
|
||||
futures::stream::iter(events)
|
||||
});
|
||||
|
||||
Ok(Box::pin(stream))
|
||||
}
|
||||
|
||||
@@ -116,8 +116,7 @@ pub struct OllamaChatResponse {
|
||||
|
||||
pub message: OllamaMessage,
|
||||
|
||||
pub done: bool,
|
||||
pub done_reason: String,
|
||||
pub done_reason: OllamaFinishReason,
|
||||
|
||||
pub total_duration: u64,
|
||||
pub load_duration: u64,
|
||||
@@ -126,9 +125,34 @@ pub struct OllamaChatResponse {
|
||||
pub eval_count: u32,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct OllamaChatStreamResponse {
|
||||
pub model: String,
|
||||
pub created_at: String,
|
||||
|
||||
pub message: OllamaMessage,
|
||||
|
||||
pub done: bool,
|
||||
pub done_reason: Option<OllamaFinishReason>,
|
||||
|
||||
pub total_duration: Option<u64>,
|
||||
pub load_duration: Option<u64>,
|
||||
|
||||
pub prompt_eval_count: Option<u32>,
|
||||
pub eval_count: Option<u32>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize, Serialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum OllamaFinishReason {
|
||||
Stop,
|
||||
Length,
|
||||
Error,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum OllamaChatStreamEvent {
|
||||
Token(String),
|
||||
Token(OllamaChatStreamResponse),
|
||||
Final(OllamaChatResponse),
|
||||
}
|
||||
|
||||
|
||||
@@ -9,8 +9,11 @@ use crate::services::errors::ServiceError;
|
||||
|
||||
use super::ConversationService;
|
||||
|
||||
use async_stream::try_stream;
|
||||
use futures::StreamExt;
|
||||
use std::boxed::Box;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::Mutex;
|
||||
use uuid::Uuid;
|
||||
|
||||
#[derive(Clone)]
|
||||
@@ -50,6 +53,23 @@ impl ChatService {
|
||||
Ok(crate::core::llm::models::LoadModelResponse { model: body.model })
|
||||
}
|
||||
|
||||
pub async fn unload_model(
|
||||
&self,
|
||||
body: crate::core::llm::models::UnloadModelRequest,
|
||||
) -> Result<crate::core::llm::models::UnloadModelResponse, ServiceError> {
|
||||
let b = crate::providers::ollama::types::OllamaGenerateRequest {
|
||||
model: body.model.clone(),
|
||||
prompt: "unload".to_string(),
|
||||
stream: false,
|
||||
keep_alive: "0s".to_string(),
|
||||
options: None,
|
||||
};
|
||||
|
||||
self.ollama.completions(&b).await?;
|
||||
|
||||
Ok(crate::core::llm::models::UnloadModelResponse { model: body.model })
|
||||
}
|
||||
|
||||
pub async fn complete(
|
||||
&self,
|
||||
body: core::llm::completions::CompletionRequest,
|
||||
@@ -136,47 +156,77 @@ impl ChatService {
|
||||
request.messages = history.into_iter().map(Into::into).collect();
|
||||
|
||||
if stream {
|
||||
let created_at = 3;
|
||||
|
||||
let start_event = futures::stream::once(async move {
|
||||
Ok(core::llm::chat::ChatCompletionStreamEvent::Start {
|
||||
conversation_id,
|
||||
message_id: user_msg_id,
|
||||
created_at,
|
||||
})
|
||||
});
|
||||
|
||||
let ollama_stream = self.ollama.chat_completions_stream(&request).await?;
|
||||
let mut ollama_stream = Box::pin(self.ollama.chat_completions_stream(&request).await?);
|
||||
|
||||
let conversation_svc = self.conversation.clone();
|
||||
let user_id = auth.user_id();
|
||||
|
||||
let mapped = ollama_stream.then(move |item| {
|
||||
let conversation_svc = conversation_svc.clone();
|
||||
async move {
|
||||
let out = try_stream! {
|
||||
let accumulated = Arc::new(Mutex::new(String::new()));
|
||||
|
||||
// Pull first item to get real model/created_at for Start.
|
||||
let first = ollama_stream.next().await;
|
||||
|
||||
let Some(first) = first else {
|
||||
Err(ServiceError::Internal(
|
||||
"provider stream ended before producing any events".to_string(),
|
||||
))?;
|
||||
return;
|
||||
};
|
||||
|
||||
let first = first?; // propagates provider error via `?` inside try_stream!
|
||||
|
||||
let (model, created_at) = match &first {
|
||||
OllamaChatStreamEvent::Token(tok) => (tok.model.clone(), tok.created_at.clone()),
|
||||
OllamaChatStreamEvent::Final(resp) => (resp.model.clone(), resp.created_at.clone()),
|
||||
};
|
||||
|
||||
yield core::llm::chat::ChatCompletionStreamEvent::Start {
|
||||
model,
|
||||
conversation_id,
|
||||
message_id: user_msg_id,
|
||||
created_at,
|
||||
};
|
||||
|
||||
// Helper closure-like inline handling so we don't duplicate match logic;
|
||||
// process `first`, then continue draining the rest of the stream.
|
||||
let mut pending = Some(first);
|
||||
|
||||
loop {
|
||||
let item = match pending.take() {
|
||||
Some(item) => item,
|
||||
None => match ollama_stream.next().await {
|
||||
Some(res) => res?,
|
||||
None => break,
|
||||
},
|
||||
};
|
||||
|
||||
match item {
|
||||
Err(e) => Err(e.into()),
|
||||
Ok(OllamaChatStreamEvent::Token(tok)) => {
|
||||
Ok(core::llm::chat::ChatCompletionStreamEvent::Token(tok))
|
||||
OllamaChatStreamEvent::Token(tok) => {
|
||||
accumulated.lock().await.push_str(&tok.message.content);
|
||||
|
||||
yield core::llm::chat::ChatCompletionStreamEvent::Token {
|
||||
content: tok.message.content,
|
||||
created_at: tok.created_at,
|
||||
id: Uuid::new_v4(),
|
||||
};
|
||||
}
|
||||
Ok(OllamaChatStreamEvent::Final(resp)) => {
|
||||
tracing::debug!("Inserting {:?}", resp.message.content);
|
||||
|
||||
OllamaChatStreamEvent::Final(mut resp) => {
|
||||
let content = accumulated.lock().await.clone();
|
||||
resp.message.content = content.clone();
|
||||
|
||||
let assistant_message_id = conversation_svc
|
||||
.log_assistant_message(
|
||||
user_id,
|
||||
conversation_id,
|
||||
user_msg_id,
|
||||
&resp.message.content,
|
||||
&content,
|
||||
resp.eval_count,
|
||||
)
|
||||
.await?;
|
||||
|
||||
conversation_svc
|
||||
.update_message_tokens(user_id, user_msg_id, resp.prompt_eval_count)
|
||||
.await?;
|
||||
|
||||
Ok(core::llm::chat::ChatCompletionStreamEvent::Final(
|
||||
yield core::llm::chat::ChatCompletionStreamEvent::Final(
|
||||
core::llm::chat::ChatCompletionResultNoStream {
|
||||
id: assistant_message_id,
|
||||
conversation_id,
|
||||
@@ -185,19 +235,17 @@ impl ChatService {
|
||||
message: resp.message.into(),
|
||||
prompt_tokens: resp.prompt_eval_count,
|
||||
completion_tokens: resp.eval_count,
|
||||
done_reason: resp.done_reason,
|
||||
finish_reason: resp.done_reason.into(),
|
||||
total_duration: resp.total_duration,
|
||||
load_duration: resp.load_duration,
|
||||
},
|
||||
))
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
};
|
||||
|
||||
Ok(core::llm::chat::ChatCompletionResult::Stream(Box::pin(
|
||||
start_event.chain(mapped),
|
||||
)))
|
||||
Ok(core::llm::chat::ChatCompletionResult::Stream(Box::pin(out)))
|
||||
} else {
|
||||
let response = self.ollama.chat_completions(&request).await?;
|
||||
|
||||
@@ -223,7 +271,7 @@ impl ChatService {
|
||||
message: response.message.into(),
|
||||
prompt_tokens: response.prompt_eval_count,
|
||||
completion_tokens: response.eval_count,
|
||||
done_reason: response.done_reason,
|
||||
finish_reason: response.done_reason.into(),
|
||||
total_duration: response.total_duration,
|
||||
load_duration: response.load_duration,
|
||||
};
|
||||
|
||||
@@ -1,12 +1,13 @@
|
||||
use crate::databases::errors::DbError;
|
||||
use crate::providers::keycloak::errors::AuthError;
|
||||
use crate::providers::keycloak::errors::JwtValidationError;
|
||||
use crate::providers::ollama::errors::LlmError;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum ServiceError {
|
||||
Db(DbError),
|
||||
Llm(LlmError),
|
||||
Auth(AuthError),
|
||||
Auth(JwtValidationError),
|
||||
Internal(String),
|
||||
}
|
||||
|
||||
impl From<DbError> for ServiceError {
|
||||
@@ -21,8 +22,8 @@ impl From<LlmError> for ServiceError {
|
||||
}
|
||||
}
|
||||
|
||||
impl From<AuthError> for ServiceError {
|
||||
fn from(e: AuthError) -> Self {
|
||||
impl From<JwtValidationError> for ServiceError {
|
||||
fn from(e: JwtValidationError) -> Self {
|
||||
ServiceError::Auth(e)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user