From bdbd9413fb96f74d19e1fa7988ff899c17c2c63f Mon Sep 17 00:00:00 2001 From: LucasDLTG Date: Tue, 14 Jul 2026 23:11:20 +0200 Subject: [PATCH] fix: fix and check all llm endpoints --- Cargo.lock | 23 ++ Cargo.toml | 3 +- src/api/errors.rs | 40 +- src/api/middlewares/auth.rs | 95 ++--- src/api/routes/v1/llm.rs | 25 +- src/api/routes/v1/tts.rs | 574 ++++++++++++++++++++++++++++ src/api/types.rs | 23 +- src/core/llm/chat.rs | 18 +- src/core/llm/models.rs | 1 - src/mappers/core_to_api.rs | 10 + src/mappers/ollama_to_core.rs | 10 + src/providers/keycloak/errors.rs | 2 +- src/providers/keycloak/jwks.rs | 8 +- src/providers/keycloak/validator.rs | 34 +- src/providers/ollama/client.rs | 116 +++--- src/providers/ollama/types.rs | 30 +- src/services/chat_service.rs | 97 +++-- src/services/errors.rs | 9 +- 18 files changed, 909 insertions(+), 209 deletions(-) create mode 100644 src/api/routes/v1/tts.rs diff --git a/Cargo.lock b/Cargo.lock index f2bac97..54f1753 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -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", diff --git a/Cargo.toml b/Cargo.toml index 2f174a8..b18cc72 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -28,4 +28,5 @@ 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" \ No newline at end of file +sha2 = "0.11.0" +async-stream = "0.3" \ No newline at end of file diff --git a/src/api/errors.rs b/src/api/errors.rs index 1905cab..b1f1afb 100644 --- a/src/api/errors.rs +++ b/src/api/errors.rs @@ -1,6 +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; @@ -92,35 +92,41 @@ impl From 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, + }, } } } diff --git a/src/api/middlewares/auth.rs b/src/api/middlewares/auth.rs index 277cdab..08785f3 100644 --- a/src/api/middlewares/auth.rs +++ b/src/api/middlewares/auth.rs @@ -20,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, - req: Request, + mut req: Request, next: Next, ) -> Response { let headers = req.headers(); @@ -32,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 { + println!("{:?}", headers); + let token = headers .get("authorization") .and_then(|v| v.to_str().ok()) @@ -108,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 { - 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?; } @@ -125,10 +137,7 @@ async fn handle_auth( .await?; } } - - request.extensions_mut().insert(auth); - - Ok(next.run(request).await) + Ok(()) } // ── Role guard ─────────────────────────────────────────────────────────────── diff --git a/src/api/routes/v1/llm.rs b/src/api/routes/v1/llm.rs index d456fb0..b8a52de 100644 --- a/src/api/routes/v1/llm.rs +++ b/src/api/routes/v1/llm.rs @@ -136,14 +136,10 @@ pub async fn load_model( pub async fn unload_model( State(state): State, Path(model): Path, - Json(body): Json, ) -> Result, api::errors::ApiError> { let response = state .chat_service - .unload_model(crate::core::llm::models::UnloadModelRequest { - model, - keep_alive: body.keep_alive.clone(), - }) + .unload_model(crate::core::llm::models::UnloadModelRequest { model }) .await?; Ok(Json(api::types::ApiUnloadModelResponse { @@ -284,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(); @@ -318,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(); diff --git a/src/api/routes/v1/tts.rs b/src/api/routes/v1/tts.rs new file mode 100644 index 0000000..51893bd --- /dev/null +++ b/src/api/routes/v1/tts.rs @@ -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, + /// BCP-47 language code. Defaults to the voice's native language. + #[schema(example = "en-US")] + pub language: Option, + /// Output audio format. Defaults to mp3. + pub format: Option, + /// Playback speed multiplier (0.5-2.0). Defaults to 1.0. + #[schema(example = 1.0)] + pub speed: Option, +} + +#[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, + pub is_cloned: bool, + pub created_at: chrono::DateTime, +} + +#[derive(Debug, Serialize, ToSchema)] +pub struct ApiVoiceListResponse { + pub voices: Vec, +} + +#[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, +} + +#[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, +} + +#[derive(Debug, Serialize, ToSchema)] +pub struct ApiTtsModel { + pub id: String, + pub name: String, + pub description: Option, + pub supported_languages: Vec, +} + +#[derive(Debug, Serialize, ToSchema)] +pub struct ApiTtsModelListResponse { + pub models: Vec, +} + +// --------------------------------------------------------------------- +// 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, + Json(body): Json, +) -> Result { + 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, + Json(body): Json, +) -> Result { + 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, +} + +#[utoipa::path( + get, + path = "/v1/voices", + tag = "voices", + params( + ("language" = Option, 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, + Query(query): Query, +) -> Result, 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, + mut multipart: Multipart, +) -> Result, ApiError> { + let mut name: Option = None; + let mut language: Option = None; + let mut audio_sample: Option> = 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, + Path(id): Path, +) -> Result, 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, + Path(id): Path, +) -> Result, 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, +) -> Result, 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, +) -> Result, 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 { + 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)) +} \ No newline at end of file diff --git a/src/api/types.rs b/src/api/types.rs index 7a9d1fb..0af5ffc 100644 --- a/src/api/types.rs +++ b/src/api/types.rs @@ -47,11 +47,6 @@ pub struct ApiUnloadModelResponse { pub status: String, } -#[derive(Debug, Serialize, Deserialize, ToSchema)] -pub struct ApiUnloadModelRequest { - pub keep_alive: String, -} - // ------ Completions ------ #[derive(Debug, Deserialize, Serialize, ToSchema, Default)] @@ -110,8 +105,6 @@ pub enum ApiCompletionObject { pub enum ApiFinishReason { Stop, Length, - ContentFilter, - ToolCalls, Error, } @@ -129,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, -// } - #[derive(Debug, Deserialize, Serialize, ToSchema)] pub struct ApiChatRequest { #[serde(flatten)] @@ -184,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, - pub usage: Option, } #[derive(Debug, Serialize, Deserialize, ToSchema)] pub struct ChatChunkChoice { pub index: u32, pub delta: Delta, - pub finish_reason: Option, } #[derive(Debug, Serialize, Deserialize, ToSchema)] @@ -206,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)] diff --git a/src/core/llm/chat.rs b/src/core/llm/chat.rs index 3cf389d..bfce257 100644 --- a/src/core/llm/chat.rs +++ b/src/core/llm/chat.rs @@ -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), } diff --git a/src/core/llm/models.rs b/src/core/llm/models.rs index 1029f5a..a7e29e7 100644 --- a/src/core/llm/models.rs +++ b/src/core/llm/models.rs @@ -38,7 +38,6 @@ pub struct LoadModelResponse { #[derive(Debug, Clone)] pub struct UnloadModelRequest { pub model: String, - pub keep_alive: String, } #[derive(Debug, Clone)] diff --git a/src/mappers/core_to_api.rs b/src/mappers/core_to_api.rs index 8c73ea0..bddb23e 100644 --- a/src/mappers/core_to_api.rs +++ b/src/mappers/core_to_api.rs @@ -134,3 +134,13 @@ impl From for api::types::ApiChat } } } + +impl From 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, + } + } +} diff --git a/src/mappers/ollama_to_core.rs b/src/mappers/ollama_to_core.rs index 0a2eb83..8e849eb 100644 --- a/src/mappers/ollama_to_core.rs +++ b/src/mappers/ollama_to_core.rs @@ -48,3 +48,13 @@ impl From for core::llm::chat::Message { } } } + +impl From 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, + } + } +} diff --git a/src/providers/keycloak/errors.rs b/src/providers/keycloak/errors.rs index ba3bb2f..f721a6e 100644 --- a/src/providers/keycloak/errors.rs +++ b/src/providers/keycloak/errors.rs @@ -1,7 +1,7 @@ use thiserror::Error; #[derive(Debug, Error)] -pub enum AuthError { +pub enum JwtValidationError { #[error("invalid authorization header")] InvalidHeader, diff --git a/src/providers/keycloak/jwks.rs b/src/providers/keycloak/jwks.rs index 144803e..0f737a1 100644 --- a/src/providers/keycloak/jwks.rs +++ b/src/providers/keycloak/jwks.rs @@ -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>>> = Lazy::new(|| Arc::new(R static JWKS_URL: Lazy = Lazy::new(|| env::var("JWKS_URL").expect("JWKS_URL not set")); -async fn fetch_jwks() -> Result { +async fn fetch_jwks() -> Result { let jwks = reqwest::get(JWKS_URL.as_str()) .await? .json::() @@ -26,7 +26,7 @@ async fn fetch_jwks() -> Result { Ok(jwks) } -pub async fn refresh_jwks() -> Result { +pub async fn refresh_jwks() -> Result { let jwks = fetch_jwks().await?; let mut write = JWK_CACHE.write().await; @@ -39,7 +39,7 @@ pub async fn refresh_jwks() -> Result { Ok(jwks) } -pub async fn get_jwks() -> Result { +pub async fn get_jwks() -> Result { let ttl = Duration::from_secs(3600); // 1 hour { diff --git a/src/providers/keycloak/validator.rs b/src/providers/keycloak/validator.rs index 295bb9f..362cfab 100644 --- a/src/providers/keycloak/validator.rs +++ b/src/providers/keycloak/validator.rs @@ -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 = Lazy::new(|| env::var("ISSUER").expect("ISSUER not set")); -fn validate_token(token: &str, jwks: &Value) -> Result { +fn validate_token(token: &str, jwks: &Value) -> Result { // 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(token, &decoding_key, &validation) - .map_err(|_| AuthError::TokenValidationFailed)?; + .map_err(|_| JwtValidationError::TokenValidationFailed)?; Ok(token_data.claims) } -pub async fn authenticate_jwt(token: &str) -> Result { +pub async fn authenticate_jwt(token: &str) -> Result { 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 // 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) } } } diff --git a/src/providers/ollama/client.rs b/src/providers/ollama/client.rs index 5c319a6..5d809a7 100644 --- a/src/providers/ollama/client.rs +++ b/src/providers/ollama/client.rs @@ -285,62 +285,70 @@ impl OllamaProvider { .await? .bytes_stream(); - let stream = byte_stream - .flat_map(|chunk_result| { - let mut out: Vec> = - 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> = 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)) } diff --git a/src/providers/ollama/types.rs b/src/providers/ollama/types.rs index 4d4511d..388ee8c 100644 --- a/src/providers/ollama/types.rs +++ b/src/providers/ollama/types.rs @@ -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, + + pub total_duration: Option, + pub load_duration: Option, + + pub prompt_eval_count: Option, + pub eval_count: Option, +} + +#[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), } diff --git a/src/services/chat_service.rs b/src/services/chat_service.rs index 237b496..96d536e 100644 --- a/src/services/chat_service.rs +++ b/src/services/chat_service.rs @@ -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)] @@ -58,7 +61,7 @@ impl ChatService { model: body.model.clone(), prompt: "unload".to_string(), stream: false, - keep_alive: body.keep_alive, + keep_alive: "0s".to_string(), options: None, }; @@ -153,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, @@ -202,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?; @@ -240,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, }; diff --git a/src/services/errors.rs b/src/services/errors.rs index c3bac93..e864db6 100644 --- a/src/services/errors.rs +++ b/src/services/errors.rs @@ -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 for ServiceError { @@ -21,8 +22,8 @@ impl From for ServiceError { } } -impl From for ServiceError { - fn from(e: AuthError) -> Self { +impl From for ServiceError { + fn from(e: JwtValidationError) -> Self { ServiceError::Auth(e) } }