fix: fix and check all llm endpoints
CI / Rust CI (push) Failing after 1m54s

This commit is contained in:
2026-07-14 23:11:20 +02:00
parent f1b8c310c4
commit bdbd9413fb
18 changed files with 909 additions and 209 deletions
Generated
+23
View File
@@ -42,6 +42,28 @@ dependencies = [
"serde_json", "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]] [[package]]
name = "atoi" name = "atoi"
version = "2.0.0" version = "2.0.0"
@@ -241,6 +263,7 @@ checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724"
name = "chat" name = "chat"
version = "0.1.0" version = "0.1.0"
dependencies = [ dependencies = [
"async-stream",
"axum", "axum",
"base64", "base64",
"chrono", "chrono",
+1
View File
@@ -29,3 +29,4 @@ sqlx = { version = "0.8.6", features = ["runtime-tokio-rustls", "postgres", "uui
rand = "0.8" rand = "0.8"
base64 = "0.22.1" base64 = "0.22.1"
sha2 = "0.11.0" sha2 = "0.11.0"
async-stream = "0.3"
+20 -14
View File
@@ -1,6 +1,6 @@
use crate::api::middlewares::errors::AuthMiddlewareError; use crate::api::middlewares::errors::AuthMiddlewareError;
use crate::databases::errors::DbError; 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::providers::ollama::errors::LlmError;
use crate::services::errors::ServiceError; use crate::services::errors::ServiceError;
@@ -92,35 +92,41 @@ impl From<ServiceError> for ApiError {
}, },
ServiceError::Auth(e) => match e { ServiceError::Auth(e) => match e {
AuthError::InvalidToken | AuthError::TokenValidationFailed => Self { JwtValidationError::InvalidToken | JwtValidationError::TokenValidationFailed => {
Self {
status: StatusCode::UNAUTHORIZED, status: StatusCode::UNAUTHORIZED,
code: "AUTH_INVALID_TOKEN", code: "AUTH_INVALID_TOKEN",
message: "invalid or expired token".into(), message: "invalid or expired token".into(),
}, }
AuthError::InvalidHeader }
// | AuthError::InvalidHeaderDecode JwtValidationError::InvalidHeader | JwtValidationError::MissingKid => Self {
| AuthError::MissingKid => Self {
status: StatusCode::UNAUTHORIZED, status: StatusCode::UNAUTHORIZED,
code: "AUTH_INVALID_HEADER", code: "AUTH_INVALID_HEADER",
message: "invalid authorization header".into(), message: "invalid authorization header".into(),
}, },
AuthError::JwkNotFound JwtValidationError::JwkNotFound
| AuthError::InvalidJwks | JwtValidationError::InvalidJwks
| AuthError::MissingModulus | JwtValidationError::MissingModulus
| AuthError::MissingExponent | JwtValidationError::MissingExponent
| AuthError::InvalidDecodingKey => Self { | JwtValidationError::InvalidDecodingKey => Self {
status: StatusCode::INTERNAL_SERVER_ERROR, status: StatusCode::INTERNAL_SERVER_ERROR,
code: "AUTH_JWKS_ERROR", code: "AUTH_JWKS_ERROR",
message: "key validation error".into(), message: "key validation error".into(),
}, },
AuthError::JwksFetchFailed JwtValidationError::JwksFetchFailed
| AuthError::JwksRefreshFailed | JwtValidationError::JwksRefreshFailed
| AuthError::Reqwest(_) => Self { | JwtValidationError::Reqwest(_) => Self {
status: StatusCode::SERVICE_UNAVAILABLE, status: StatusCode::SERVICE_UNAVAILABLE,
code: "AUTH_JWKS_FETCH", code: "AUTH_JWKS_FETCH",
message: "failed to fetch authorization keys".into(), message: "failed to fetch authorization keys".into(),
}, },
}, },
ServiceError::Internal(e) => Self {
status: StatusCode::INTERNAL_SERVER_ERROR,
code: "INTERNAL_SERVER_ERROR",
message: e,
},
} }
} }
} }
+51 -42
View File
@@ -20,9 +20,47 @@ enum ApiKeyError {
Service(ApiError), 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( pub async fn auth_middleware(
State(state): State<SharedState>, State(state): State<SharedState>,
req: Request, mut req: Request,
next: Next, next: Next,
) -> Response { ) -> Response {
let headers = req.headers(); let headers = req.headers();
@@ -32,48 +70,27 @@ pub async fn auth_middleware(
Err(jwt_err) => match try_api_key(&state, headers).await { Err(jwt_err) => match try_api_key(&state, headers).await {
Ok(auth) => auth, Ok(auth) => auth,
Err(api_key_err) => { Err(api_key_err) => {
let err = match (jwt_err, api_key_err) { return resolve_auth_error(jwt_err, api_key_err).into_response();
// 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();
} }
}, },
}; };
match handle_auth(&state, req, next, auth).await { if let Err(err) = record_auth(&state, &auth).await {
Ok(response) => {
tracing::debug!("User authentified");
response
}
Err(err) => {
tracing::debug!("Error during authentification {:?}", err); tracing::debug!("Error during authentification {:?}", err);
err.into_response() return err.into_response();
}
} }
tracing::debug!("User authentified");
req.extensions_mut().insert(auth);
next.run(req).await
} }
async fn try_jwt( async fn try_jwt(
state: &SharedState, state: &SharedState,
headers: &HeaderMap, headers: &HeaderMap,
) -> Result<crate::core::auth::Auth, JwtError> { ) -> Result<crate::core::auth::Auth, JwtError> {
println!("{:?}", headers);
let token = headers let token = headers
.get("authorization") .get("authorization")
.and_then(|v| v.to_str().ok()) .and_then(|v| v.to_str().ok())
@@ -108,13 +125,8 @@ async fn try_api_key(
Ok(crate::core::auth::Auth::ApiKey(auth)) Ok(crate::core::auth::Auth::ApiKey(auth))
} }
async fn handle_auth( async fn record_auth(state: &SharedState, auth: &crate::core::auth::Auth) -> Result<(), ApiError> {
state: &SharedState, match auth {
mut request: Request,
next: Next,
auth: crate::core::auth::Auth,
) -> Result<Response, ApiError> {
match &auth {
crate::core::auth::Auth::Jwt(_) => { crate::core::auth::Auth::Jwt(_) => {
state.auth_service.create_user(&auth.user_id()).await?; state.auth_service.create_user(&auth.user_id()).await?;
} }
@@ -125,10 +137,7 @@ async fn handle_auth(
.await?; .await?;
} }
} }
Ok(())
request.extensions_mut().insert(auth);
Ok(next.run(request).await)
} }
// ── Role guard ─────────────────────────────────────────────────────────────── // ── Role guard ───────────────────────────────────────────────────────────────
+13 -12
View File
@@ -136,14 +136,10 @@ pub async fn load_model(
pub async fn unload_model( pub async fn unload_model(
State(state): State<SharedState>, State(state): State<SharedState>,
Path(model): Path<String>, Path(model): Path<String>,
Json(body): Json<api::types::ApiUnloadModelRequest>,
) -> Result<Json<api::types::ApiUnloadModelResponse>, api::errors::ApiError> { ) -> Result<Json<api::types::ApiUnloadModelResponse>, api::errors::ApiError> {
let response = state let response = state
.chat_service .chat_service
.unload_model(crate::core::llm::models::UnloadModelRequest { .unload_model(crate::core::llm::models::UnloadModelRequest { model })
model,
keep_alive: body.keep_alive.clone(),
})
.await?; .await?;
Ok(Json(api::types::ApiUnloadModelResponse { Ok(Json(api::types::ApiUnloadModelResponse {
@@ -284,31 +280,36 @@ pub async fn chat_completions(
conversation_id, conversation_id,
message_id, message_id,
created_at, created_at,
model,
}) => { }) => {
println!("RECEIVED STRATTT");
let payload = serde_json::to_string(&api::types::StreamEvent::Start( let payload = serde_json::to_string(&api::types::StreamEvent::Start(
api::types::StartEventData { api::types::StartEventData {
conversation_id, conversation_id,
created: created_at, created: created_at,
id: message_id, id: message_id,
model,
}, },
)) ))
.unwrap_or_default(); .unwrap_or_default();
Ok(Event::default().data(payload)) 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 { let chunk = api::types::ChatCompletionChunk {
id: String::new(), id,
object: "chat.completion.chunk".to_string(), object: "chat.completion.chunk".to_string(),
choices: vec![api::types::ChatChunkChoice { choices: vec![api::types::ChatChunkChoice {
index: 0, index: 0,
delta: api::types::Delta { delta: api::types::Delta {
content: Some(tok), content: Some(content),
role: None, role: Some(api::types::ApiRole::Assistant),
}, },
finish_reason: None,
}], }],
usage: None,
}; };
let payload = serde_json::to_string(&api::types::StreamEvent::Delta(chunk)) let payload = serde_json::to_string(&api::types::StreamEvent::Delta(chunk))
.unwrap_or_default(); .unwrap_or_default();
@@ -318,13 +319,13 @@ pub async fn chat_completions(
Ok(crate::core::llm::chat::ChatCompletionStreamEvent::Final(res)) => { Ok(crate::core::llm::chat::ChatCompletionStreamEvent::Final(res)) => {
let payload = serde_json::to_string(&api::types::StreamEvent::End( let payload = serde_json::to_string(&api::types::StreamEvent::End(
api::types::EndEventData { api::types::EndEventData {
created: res.created_at.parse().unwrap_or(0),
id: res.id, id: res.id,
usage: api::types::Usage { usage: api::types::Usage {
prompt_tokens: res.prompt_tokens, prompt_tokens: res.prompt_tokens,
completion_tokens: res.completion_tokens, completion_tokens: res.completion_tokens,
total_tokens: res.prompt_tokens + res.completion_tokens, total_tokens: res.prompt_tokens + res.completion_tokens,
}, },
finish_reason: res.finish_reason.into(),
}, },
)) ))
.unwrap_or_default(); .unwrap_or_default();
+574
View File
@@ -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 -19
View File
@@ -47,11 +47,6 @@ pub struct ApiUnloadModelResponse {
pub status: String, pub status: String,
} }
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ApiUnloadModelRequest {
pub keep_alive: String,
}
// ------ Completions ------ // ------ Completions ------
#[derive(Debug, Deserialize, Serialize, ToSchema, Default)] #[derive(Debug, Deserialize, Serialize, ToSchema, Default)]
@@ -110,8 +105,6 @@ pub enum ApiCompletionObject {
pub enum ApiFinishReason { pub enum ApiFinishReason {
Stop, Stop,
Length, Length,
ContentFilter,
ToolCalls,
Error, Error,
} }
@@ -129,13 +122,6 @@ pub struct Usage {
pub total_tokens: u32, 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)] #[derive(Debug, Deserialize, Serialize, ToSchema)]
pub struct ApiChatRequest { pub struct ApiChatRequest {
#[serde(flatten)] #[serde(flatten)]
@@ -184,17 +170,15 @@ pub struct ApiChatChoice {
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ChatCompletionChunk { pub struct ChatCompletionChunk {
pub id: String, pub id: Uuid,
pub object: String, pub object: String,
pub choices: Vec<ChatChunkChoice>, pub choices: Vec<ChatChunkChoice>,
pub usage: Option<Usage>,
} }
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ChatChunkChoice { pub struct ChatChunkChoice {
pub index: u32, pub index: u32,
pub delta: Delta, pub delta: Delta,
pub finish_reason: Option<ApiFinishReason>,
} }
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
@@ -206,15 +190,16 @@ pub struct Delta {
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct StartEventData { pub struct StartEventData {
pub conversation_id: Uuid, pub conversation_id: Uuid,
pub created: u64, pub created: String,
pub id: Uuid, pub id: Uuid,
pub model: String,
} }
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct EndEventData { pub struct EndEventData {
pub created: u64,
pub id: Uuid, pub id: Uuid,
pub usage: Usage, pub usage: Usage,
pub finish_reason: ApiFinishReason,
} }
#[derive(Debug, Serialize, Deserialize, ToSchema)] #[derive(Debug, Serialize, Deserialize, ToSchema)]
+15 -3
View File
@@ -51,20 +51,32 @@ pub struct ChatCompletionResultNoStream {
pub prompt_tokens: u32, pub prompt_tokens: u32,
pub completion_tokens: u32, pub completion_tokens: u32,
pub done_reason: String, pub finish_reason: FinishReason,
pub total_duration: u64, pub total_duration: u64,
pub load_duration: u64, pub load_duration: u64,
} }
#[derive(Debug, Serialize, Clone)]
pub enum FinishReason {
Stop,
Length,
Error,
}
#[derive(Debug, Serialize)] #[derive(Debug, Serialize)]
pub enum ChatCompletionStreamEvent { pub enum ChatCompletionStreamEvent {
Start { Start {
conversation_id: Uuid, conversation_id: Uuid,
message_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), Final(ChatCompletionResultNoStream),
} }
-1
View File
@@ -38,7 +38,6 @@ pub struct LoadModelResponse {
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct UnloadModelRequest { pub struct UnloadModelRequest {
pub model: String, pub model: String,
pub keep_alive: String,
} }
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
+10
View File
@@ -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,
}
}
}
+10
View File
@@ -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 -1
View File
@@ -1,7 +1,7 @@
use thiserror::Error; use thiserror::Error;
#[derive(Debug, Error)] #[derive(Debug, Error)]
pub enum AuthError { pub enum JwtValidationError {
#[error("invalid authorization header")] #[error("invalid authorization header")]
InvalidHeader, InvalidHeader,
+4 -4
View File
@@ -1,4 +1,4 @@
use super::errors::AuthError; use super::errors::JwtValidationError;
use once_cell::sync::Lazy; use once_cell::sync::Lazy;
use serde_json::Value; 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")); 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()) let jwks = reqwest::get(JWKS_URL.as_str())
.await? .await?
.json::<Value>() .json::<Value>()
@@ -26,7 +26,7 @@ async fn fetch_jwks() -> Result<Value, AuthError> {
Ok(jwks) Ok(jwks)
} }
pub async fn refresh_jwks() -> Result<Value, AuthError> { pub async fn refresh_jwks() -> Result<Value, JwtValidationError> {
let jwks = fetch_jwks().await?; let jwks = fetch_jwks().await?;
let mut write = JWK_CACHE.write().await; let mut write = JWK_CACHE.write().await;
@@ -39,7 +39,7 @@ pub async fn refresh_jwks() -> Result<Value, AuthError> {
Ok(jwks) 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 let ttl = Duration::from_secs(3600); // 1 hour
{ {
+20 -14
View File
@@ -1,5 +1,5 @@
use super::claims::KeycloakClaims; use super::claims::KeycloakClaims;
use super::errors::AuthError; use super::errors::JwtValidationError;
use jsonwebtoken::{Algorithm, DecodingKey, Validation, decode, decode_header}; use jsonwebtoken::{Algorithm, DecodingKey, Validation, decode, decode_header};
use once_cell::sync::Lazy; 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")); 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 // 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 // 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 let key = keys
.iter() .iter()
.find(|k| k["kid"] == kid) .find(|k| k["kid"] == kid)
.ok_or(AuthError::InvalidDecodingKey)?; .ok_or(JwtValidationError::InvalidDecodingKey)?;
// 3. Extract RSA components // 3. Extract RSA components
let n = key["n"].as_str().ok_or(AuthError::MissingModulus)?; let n = key["n"]
let e = key["e"].as_str().ok_or(AuthError::MissingExponent)?; .as_str()
.ok_or(JwtValidationError::MissingModulus)?;
let e = key["e"]
.as_str()
.ok_or(JwtValidationError::MissingExponent)?;
let decoding_key = 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 // 4. Setup validation rules
let mut validation = Validation::new(Algorithm::RS256); let mut validation = Validation::new(Algorithm::RS256);
@@ -39,15 +45,15 @@ fn validate_token(token: &str, jwks: &Value) -> Result<KeycloakClaims, AuthError
// 5. Decode & verify // 5. Decode & verify
let token_data = decode::<KeycloakClaims>(token, &decoding_key, &validation) let token_data = decode::<KeycloakClaims>(token, &decoding_key, &validation)
.map_err(|_| AuthError::TokenValidationFailed)?; .map_err(|_| JwtValidationError::TokenValidationFailed)?;
Ok(token_data.claims) 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() let jwks = super::jwks::get_jwks()
.await .await
.map_err(|_| AuthError::JwksFetchFailed)?; .map_err(|_| JwtValidationError::JwksFetchFailed)?;
match validate_token(token, &jwks) { match validate_token(token, &jwks) {
Ok(claims) => Ok(claims), Ok(claims) => Ok(claims),
@@ -56,9 +62,9 @@ pub async fn authenticate_jwt(token: &str) -> Result<KeycloakClaims, AuthError>
// one retry with refresh // one retry with refresh
let fresh = super::jwks::refresh_jwks() let fresh = super::jwks::refresh_jwks()
.await .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)
} }
} }
} }
+42 -34
View File
@@ -285,61 +285,69 @@ impl OllamaProvider {
.await? .await?
.bytes_stream(); .bytes_stream();
let stream = byte_stream let stream = byte_stream.flat_map(|chunk_result| {
.flat_map(|chunk_result| { let mut events: Vec<Result<super::types::OllamaChatStreamEvent, LlmError>> = Vec::new();
let mut out: Vec<Result<super::types::OllamaChatStreamEvent, LlmError>> =
Vec::new();
let chunk = match chunk_result { let chunk = match chunk_result {
Ok(b) => b, Ok(b) => b,
Err(e) => { Err(e) => {
tracing::debug!("Error: {:?}", e); tracing::debug!("Error: {:?}", e);
out.push(Err(LlmError::Http(e))); events.push(Err(LlmError::Http(e)));
return futures::stream::iter(out); return futures::stream::iter(events);
} }
}; };
tracing::debug!("Raw line: {:?}", std::str::from_utf8(&chunk)); 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 chunk.split(|&b| b == b'\n') { for line in text.lines() {
if line.is_empty() { if line.trim().is_empty() {
continue; continue;
} }
let parsed: ollama::types::OllamaChatResponse = let parsed: ollama::types::OllamaChatStreamResponse =
match serde_json::from_slice(line) { match serde_json::from_str(line) {
Ok(v) => v, Ok(v) => v,
Err(_) => continue, Err(e) => {
tracing::debug!("Failed parsing Ollama line {:?}: {:?}", line, e);
continue;
}
}; };
tracing::debug!("Parsed: {:?}", parsed); tracing::debug!("Parsed: {:?}", parsed);
if !parsed.message.content.is_empty() { // End of generation
out.push(Ok(super::types::OllamaChatStreamEvent::Token( if parsed.done {
parsed.message.content.clone(), 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;
} }
if parsed.done { // Normal generated token
out.push(Ok(super::types::OllamaChatStreamEvent::Final(parsed))); if !parsed.message.content.is_empty() {
return futures::stream::iter(out); events.push(Ok(super::types::OllamaChatStreamEvent::Token(parsed)));
} }
} }
futures::stream::iter(out)
}) futures::stream::iter(events)
.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)
}); });
Ok(Box::pin(stream)) Ok(Box::pin(stream))
+27 -3
View File
@@ -116,8 +116,7 @@ pub struct OllamaChatResponse {
pub message: OllamaMessage, pub message: OllamaMessage,
pub done: bool, pub done_reason: OllamaFinishReason,
pub done_reason: String,
pub total_duration: u64, pub total_duration: u64,
pub load_duration: u64, pub load_duration: u64,
@@ -126,9 +125,34 @@ pub struct OllamaChatResponse {
pub eval_count: u32, 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)] #[derive(Debug)]
pub enum OllamaChatStreamEvent { pub enum OllamaChatStreamEvent {
Token(String), Token(OllamaChatStreamResponse),
Final(OllamaChatResponse), Final(OllamaChatResponse),
} }
+64 -33
View File
@@ -9,8 +9,11 @@ use crate::services::errors::ServiceError;
use super::ConversationService; use super::ConversationService;
use async_stream::try_stream;
use futures::StreamExt; use futures::StreamExt;
use std::boxed::Box; use std::boxed::Box;
use std::sync::Arc;
use tokio::sync::Mutex;
use uuid::Uuid; use uuid::Uuid;
#[derive(Clone)] #[derive(Clone)]
@@ -58,7 +61,7 @@ impl ChatService {
model: body.model.clone(), model: body.model.clone(),
prompt: "unload".to_string(), prompt: "unload".to_string(),
stream: false, stream: false,
keep_alive: body.keep_alive, keep_alive: "0s".to_string(),
options: None, options: None,
}; };
@@ -153,47 +156,77 @@ impl ChatService {
request.messages = history.into_iter().map(Into::into).collect(); request.messages = history.into_iter().map(Into::into).collect();
if stream { if stream {
let created_at = 3; let mut ollama_stream = Box::pin(self.ollama.chat_completions_stream(&request).await?);
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 conversation_svc = self.conversation.clone(); let conversation_svc = self.conversation.clone();
let user_id = auth.user_id(); let user_id = auth.user_id();
let mapped = ollama_stream.then(move |item| { let out = try_stream! {
let conversation_svc = conversation_svc.clone(); let accumulated = Arc::new(Mutex::new(String::new()));
async move {
// 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 { match item {
Err(e) => Err(e.into()), OllamaChatStreamEvent::Token(tok) => {
Ok(OllamaChatStreamEvent::Token(tok)) => { accumulated.lock().await.push_str(&tok.message.content);
Ok(core::llm::chat::ChatCompletionStreamEvent::Token(tok))
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 let assistant_message_id = conversation_svc
.log_assistant_message( .log_assistant_message(
user_id, user_id,
conversation_id, conversation_id,
user_msg_id, user_msg_id,
&resp.message.content, &content,
resp.eval_count, resp.eval_count,
) )
.await?; .await?;
conversation_svc yield core::llm::chat::ChatCompletionStreamEvent::Final(
.update_message_tokens(user_id, user_msg_id, resp.prompt_eval_count)
.await?;
Ok(core::llm::chat::ChatCompletionStreamEvent::Final(
core::llm::chat::ChatCompletionResultNoStream { core::llm::chat::ChatCompletionResultNoStream {
id: assistant_message_id, id: assistant_message_id,
conversation_id, conversation_id,
@@ -202,19 +235,17 @@ impl ChatService {
message: resp.message.into(), message: resp.message.into(),
prompt_tokens: resp.prompt_eval_count, prompt_tokens: resp.prompt_eval_count,
completion_tokens: resp.eval_count, completion_tokens: resp.eval_count,
done_reason: resp.done_reason, finish_reason: resp.done_reason.into(),
total_duration: resp.total_duration, total_duration: resp.total_duration,
load_duration: resp.load_duration, load_duration: resp.load_duration,
}, },
)) );
} }
} }
} }
}); };
Ok(core::llm::chat::ChatCompletionResult::Stream(Box::pin( Ok(core::llm::chat::ChatCompletionResult::Stream(Box::pin(out)))
start_event.chain(mapped),
)))
} else { } else {
let response = self.ollama.chat_completions(&request).await?; let response = self.ollama.chat_completions(&request).await?;
@@ -240,7 +271,7 @@ impl ChatService {
message: response.message.into(), message: response.message.into(),
prompt_tokens: response.prompt_eval_count, prompt_tokens: response.prompt_eval_count,
completion_tokens: response.eval_count, completion_tokens: response.eval_count,
done_reason: response.done_reason, finish_reason: response.done_reason.into(),
total_duration: response.total_duration, total_duration: response.total_duration,
load_duration: response.load_duration, load_duration: response.load_duration,
}; };
+5 -4
View File
@@ -1,12 +1,13 @@
use crate::databases::errors::DbError; 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::providers::ollama::errors::LlmError;
#[derive(Debug)] #[derive(Debug)]
pub enum ServiceError { pub enum ServiceError {
Db(DbError), Db(DbError),
Llm(LlmError), Llm(LlmError),
Auth(AuthError), Auth(JwtValidationError),
Internal(String),
} }
impl From<DbError> for ServiceError { impl From<DbError> for ServiceError {
@@ -21,8 +22,8 @@ impl From<LlmError> for ServiceError {
} }
} }
impl From<AuthError> for ServiceError { impl From<JwtValidationError> for ServiceError {
fn from(e: AuthError) -> Self { fn from(e: JwtValidationError) -> Self {
ServiceError::Auth(e) ServiceError::Auth(e)
} }
} }