diff --git a/Cargo.lock b/Cargo.lock index 623ebf2..f5600d7 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2510,9 +2510,9 @@ checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" [[package]] name = "tokio" -version = "1.52.2" +version = "1.52.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "110a78583f19d5cdb2c5ccf321d1290344e71313c6c37d43520d386027d18386" +checksum = "8fc7f01b389ac15039e4dc9531aa973a135d7a4135281b12d7c1bc79fd57fffe" dependencies = [ "bytes", "libc", @@ -2775,6 +2775,7 @@ dependencies = [ "quote", "regex", "syn", + "uuid", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 2606fe3..4b4d1bb 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -5,11 +5,11 @@ edition = "2024" [dev-dependencies] wiremock = "0.6" -tokio = { version = "1.52.2", features = ["macros", "rt-multi-thread"] } +tokio = { version = "1.52.3", features = ["macros", "rt-multi-thread"] } [dependencies] axum = "0.8.9" -utoipa = { version = "5.5.0", features = ["axum_extras"] } +utoipa = { version = "5.5.0", features = ["axum_extras", "uuid"] } tokio = { version = "1", features = ["full"] } serde = { version = "1", features = ["derive"] } serde_json = "1" diff --git a/src/databases/postgres/chat.rs b/src/databases/postgres/chat.rs new file mode 100644 index 0000000..e69de29 diff --git a/src/databases/postgres/mod.rs b/src/databases/postgres/mod.rs index 35cb129..1aa7891 100644 --- a/src/databases/postgres/mod.rs +++ b/src/databases/postgres/mod.rs @@ -1,3 +1,3 @@ pub mod api_key; pub mod pool; -pub mod user_repository; +pub mod user; diff --git a/src/databases/postgres/user_repository.rs b/src/databases/postgres/user.rs similarity index 100% rename from src/databases/postgres/user_repository.rs rename to src/databases/postgres/user.rs diff --git a/src/dto/api.rs b/src/dto/api.rs index 5209edf..007d2ff 100644 --- a/src/dto/api.rs +++ b/src/dto/api.rs @@ -1,5 +1,6 @@ use serde::{Deserialize, Serialize}; use utoipa::ToSchema; +use uuid::Uuid; #[derive(Debug, Serialize, ToSchema)] pub struct ErrorResponse { @@ -87,6 +88,7 @@ pub enum FinishReason { Length, ContentFilter, ToolCalls, + Error, } #[derive(Debug, Serialize, Deserialize, ToSchema)] @@ -126,6 +128,9 @@ pub struct ChatRequest { pub base: BaseLLMRequest, pub messages: Vec, + + // Non standard Open AI + pub conversation_id: Option, } #[derive(Debug, Deserialize, Serialize, ToSchema)] diff --git a/src/middlewares/auth/middleware.rs b/src/middlewares/auth/middleware.rs index 700a9d9..e972dfb 100644 --- a/src/middlewares/auth/middleware.rs +++ b/src/middlewares/auth/middleware.rs @@ -5,9 +5,7 @@ use axum::{ response::{IntoResponse, Response}, }; -use crate::databases::postgres::{ - api_key::update_last_access, user_repository::ensure_user_exists, -}; +use crate::databases::postgres::{api_key::update_last_access, user::ensure_user_exists}; use crate::middlewares::auth::apikey::{ApiKeyClaims, ApiKeyClaimsRoles}; use crate::middlewares::auth::keycloak::{KeycloakClaims, get_jwks, refresh_jwks, validate_token}; use crate::state::app_state::AppState; diff --git a/tests/ollama_provider.rs b/tests/ollama_provider.rs index 2c8bd5d..ce7c819 100644 --- a/tests/ollama_provider.rs +++ b/tests/ollama_provider.rs @@ -183,6 +183,7 @@ async fn test_chat_completions_ok() { role: api::Role::User, content: "What is 2+2?".to_string(), }], + conversation_id: None, }; let res = provider.chat_completions(&req).await.unwrap(); @@ -217,6 +218,7 @@ async fn test_chat_completions_missing_messages() { ..Default::default() }, messages: vec![], + conversation_id: None, }; let err = provider.chat_completions(&req).await.unwrap_err(); @@ -243,6 +245,7 @@ async fn test_chat_completions_no_user_message() { role: api::Role::System, content: "be helpful".to_string(), }], + conversation_id: None, }; let err = provider.chat_completions(&req).await.unwrap_err(); @@ -269,6 +272,7 @@ async fn test_chat_completions_model_not_found() { role: api::Role::User, content: "hi".to_string(), }], + conversation_id: None, }; let err = provider.chat_completions(&req).await.unwrap_err();