diff --git a/Cargo.lock b/Cargo.lock index ce1fff3..f2bac97 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -257,7 +257,7 @@ dependencies = [ "thiserror 2.0.18", "tokio", "tokio-stream", - "tower-http", + "tower-http 0.7.0", "tracing", "tracing-subscriber", "utoipa", @@ -267,9 +267,9 @@ dependencies = [ [[package]] name = "chrono" -version = "0.4.44" +version = "0.4.45" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c673075a2e0e5f4a1dde27ce9dee1ea4558c7ffe648f576438a20ca1d2acc4b0" +checksum = "1aa79e62e7697b8e29b513a68abacf485adcd1fe8284a4316c5ae868e6633327" dependencies = [ "iana-time-zone", "js-sys", @@ -1222,14 +1222,14 @@ checksum = "b6d2cec3eae94f9f509c767b45932f1ada8350c4bdb85af2fcab4a3c14807981" [[package]] name = "libredox" -version = "0.1.16" +version = "0.1.18" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e02f3bb43d335493c96bf3fd3a321600bf6bd07ed34bc64118e9293bdffea46c" +checksum = "c943259e342f1e06ff2da7a83eabdfe7f92ce10262688dbf1895ff0b3e6e4652" dependencies = [ "bitflags", "libc", "plain", - "redox_syscall 0.7.5", + "redox_syscall 0.9.0", ] [[package]] @@ -1369,11 +1369,10 @@ dependencies = [ [[package]] name = "num-iter" -version = "0.1.45" +version = "0.1.46" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1429034a0490724d0075ebb2bc9e875d6503c3cf69e235a8941aa757d83ef5bf" +checksum = "c92800bd69a1eac91786bcfe9da64a897eb72911b8dc3095decbd07429e8048b" dependencies = [ - "autocfg", "num-integer", "num-traits", ] @@ -1693,9 +1692,9 @@ dependencies = [ [[package]] name = "redox_syscall" -version = "0.7.5" +version = "0.9.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4666a1a60d8412eab19d94f6d13dcc9cea0a5ef4fdf6a5db306537413c661b1b" +checksum = "c5102a6aaa05aa011a238e178e6bca86d2cb56fc9f586d37cb80f5bca6e07759" dependencies = [ "bitflags", ] @@ -1763,7 +1762,7 @@ dependencies = [ "tokio-rustls", "tokio-util", "tower", - "tower-http", + "tower-http 0.6.11", "tower-service", "url", "wasm-bindgen", @@ -2021,9 +2020,9 @@ dependencies = [ [[package]] name = "sha1" -version = "0.10.6" +version = "0.10.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" +checksum = "a978451301f4db1d02937a4ab3ccce137717b81826e79b7d49ffe3244a13c3b8" dependencies = [ "cfg-if", "cpufeatures 0.2.17", @@ -2617,6 +2616,21 @@ dependencies = [ "url", ] +[[package]] +name = "tower-http" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b11f75e912b0c2be01b63d8cf8057b8c3f97cf34abb3d431a3a4c8675498e233" +dependencies = [ + "bitflags", + "bytes", + "http", + "percent-encoding", + "pin-project-lite", + "tower-layer", + "tower-service", +] + [[package]] name = "tower-layer" version = "0.3.3" @@ -2793,9 +2807,9 @@ dependencies = [ [[package]] name = "uuid" -version = "1.23.2" +version = "1.23.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d258b83ceec21034727ecee8c382cfa6c3e133699b0742c64571814fb420c9f7" +checksum = "ea5fab0d6c3c01ae70085a09cb03d4c7a1d6314e2b3e075392783396d724ca0a" dependencies = [ "getrandom 0.4.2", "js-sys", @@ -3007,14 +3021,14 @@ version = "0.26.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "521bc38abb08001b01866da9f51eb7c5d647a19260e00054a8c7fd5f9e57f7a9" dependencies = [ - "webpki-roots 1.0.7", + "webpki-roots 1.0.8", ] [[package]] name = "webpki-roots" -version = "1.0.7" +version = "1.0.8" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "52f5ee44c96cf55f1b349600768e3ece3a8f26010c05265ab73f945bb1a2eb9d" +checksum = "bf85cb06032201fa7c6f829d7db5a7e5aa45bcc0655327713065f6f0576731bf" dependencies = [ "rustls-pki-types", ] diff --git a/Cargo.toml b/Cargo.toml index 38b7409..2f174a8 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -20,9 +20,9 @@ dotenvy = "0.15" thiserror = "2.0.18" tokio-stream = "0.1" futures = "0.3" -chrono = { version = "0.4.44", features = ["serde"] } -uuid = { version = "1.23.2", features = ["v4", "serde"] } -tower-http = { version = "0.6.11", features = ["cors"] } +chrono = { version = "0.4.45", features = ["serde"] } +uuid = { version = "1.23.5", features = ["v4", "serde"] } +tower-http = { version = "0.7.0", features = ["cors"] } tracing = "0.1.44" tracing-subscriber = { version = "0.3", features = ["env-filter"]} sqlx = { version = "0.8.6", features = ["runtime-tokio-rustls", "postgres", "uuid", "chrono", "macros"] } diff --git a/src/api/errors.rs b/src/api/errors.rs index 308cdf6..1905cab 100644 --- a/src/api/errors.rs +++ b/src/api/errors.rs @@ -1,3 +1,4 @@ +use crate::api::middlewares::errors::AuthMiddlewareError; use crate::databases::errors::DbError; use crate::providers::keycloak::errors::AuthError; use crate::providers::ollama::errors::LlmError; @@ -7,7 +8,6 @@ use axum::Json; use axum::http::StatusCode; use axum::response::{IntoResponse, Response}; use serde::Serialize; -use thiserror::Error; use utoipa::ToSchema; #[derive(Serialize, ToSchema)] @@ -25,18 +25,6 @@ pub struct ApiError { pub message: String, } -#[derive(Debug, Error)] -pub enum AuthMiddlewareError { - #[error("invalid authorization format")] - InvalidAuthorizationFormat, - - #[error("authentication required")] - AuthenticationRequired, - - #[error("insufficient permissions")] - Forbidden, -} - impl From for ApiError { fn from(err: AuthMiddlewareError) -> Self { match err { diff --git a/src/api/middlewares/auth.rs b/src/api/middlewares/auth.rs index 36ec3c8..277cdab 100644 --- a/src/api/middlewares/auth.rs +++ b/src/api/middlewares/auth.rs @@ -1,4 +1,5 @@ -use crate::api::errors::{ApiError, AuthMiddlewareError}; +use crate::api::errors::ApiError; +use crate::api::middlewares::errors::AuthMiddlewareError; use crate::api::state::SharedState; use axum::{ diff --git a/src/api/middlewares/errors.rs b/src/api/middlewares/errors.rs new file mode 100644 index 0000000..aceaa28 --- /dev/null +++ b/src/api/middlewares/errors.rs @@ -0,0 +1,13 @@ +use thiserror::Error; + +#[derive(Debug, Error)] +pub enum AuthMiddlewareError { + #[error("invalid authorization format")] + InvalidAuthorizationFormat, + + #[error("authentication required")] + AuthenticationRequired, + + #[error("insufficient permissions")] + Forbidden, +} diff --git a/src/api/middlewares/mod.rs b/src/api/middlewares/mod.rs index 0e4a05d..2c1d8be 100644 --- a/src/api/middlewares/mod.rs +++ b/src/api/middlewares/mod.rs @@ -1 +1,2 @@ pub mod auth; +pub mod errors; diff --git a/src/api/routes/v1/mod.rs b/src/api/routes/v1/mod.rs index e72836c..b10855a 100644 --- a/src/api/routes/v1/mod.rs +++ b/src/api/routes/v1/mod.rs @@ -16,22 +16,23 @@ async fn openapi_json() -> Json { Json(ApiDoc::openapi()) } -fn public_router() -> Router { - Router::new().route("/docs.json", get(openapi_json)) -} - -pub fn protected_router() -> Router { +fn llm_router() -> Router { Router::new() .route("/models", get(llm::list_models)) .route("/completions", post(llm::completions)) .route("/chat/completions", post(llm::chat_completions)) .route("/models/{model}/load", post(llm::load_model)) - // .route("/models/{model}/unload", post(models::unload_model)) - .route( - "/keys/generate", - post(apikey::create_api_key) // Usage - .route_layer(role_guard!(Some("admin"), None)), - ) +} + +fn keys_router() -> Router { + Router::new().route( + "/generate", + post(apikey::create_api_key).route_layer(role_guard!(Some("admin"), None)), + ) +} + +fn log_router() -> Router { + Router::new() .route("/conversations", get(llm::get_conversations)) .route( "/conversations/{conversation_id}/messages", @@ -40,13 +41,17 @@ pub fn protected_router() -> Router { } pub fn router(state: SharedState) -> Router { + let protected = Router::new() + .nest("/llm", llm_router()) + .nest("/keys", keys_router()) + .nest("/log", log_router()) + .route_layer(middleware::from_fn_with_state( + state.clone(), + auth_middleware, + )); + Router::new() - .merge(public_router()) - .merge( - protected_router().route_layer(middleware::from_fn_with_state( - state.clone(), - auth_middleware, - )), - ) + .route("/docs.json", get(openapi_json)) + .merge(protected) .with_state(state) }