diff --git a/Cargo.toml b/Cargo.toml index 3701e7c8db..b38693abcd 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -73,7 +73,7 @@ actix-web-httpauth = "0.8.1" actix-web-validator = "5.0.1" tonic = { workspace = true } tonic-reflection = { workspace = true } -tower = "0.4.13" +tower = { version = "0.4.13", features = ["filter"] } tower-layer = "0.3.2" tar = "0.4.40" reqwest = { version = "0.11", default-features = false, features = ["stream", "rustls-tls", "blocking"] } diff --git a/src/tonic/api_key.rs b/src/tonic/api_key.rs index 5b8b89e75b..8f5dd8ccfd 100644 --- a/src/tonic/api_key.rs +++ b/src/tonic/api_key.rs @@ -1,13 +1,6 @@ -use std::task::{Context, Poll}; - use actix_web_httpauth::headers::authorization::{Bearer, Scheme}; -use futures_util::future::BoxFuture; -use reqwest::header::HeaderValue; -use reqwest::StatusCode; -use tonic::body::BoxBody; -use tonic::Code; -use tower::Service; -use tower_layer::Layer; +use tonic::Status; +use tower::filter::{FilterLayer, Predicate}; use crate::common::auth::AuthKeys; use crate::common::strings::ct_eq; @@ -30,36 +23,20 @@ const READ_ONLY_RPC_PATHS: [&str; 14] = [ ]; #[derive(Clone)] -pub struct ApiKeyMiddleware { - service: T, +pub struct ApiKeyMiddleware { auth_keys: AuthKeys, } -#[derive(Clone)] -pub struct ApiKeyMiddlewareLayer { - auth_keys: AuthKeys, -} - -impl Service> for ApiKeyMiddleware -where - S: Service< - tonic::codegen::http::Request, - Response = tonic::codegen::http::Response, - >, - S::Future: Send + 'static, -{ - type Response = tonic::codegen::http::Response; - type Error = S::Error; - type Future = BoxFuture<'static, Result>; - - fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { - self.service.poll_ready(cx) +impl ApiKeyMiddleware { + pub fn new_layer(auth_keys: AuthKeys) -> FilterLayer { + FilterLayer::new(Self { auth_keys }) } +} - fn call( - &mut self, - request: tonic::codegen::http::Request, - ) -> Self::Future { +impl Predicate> for ApiKeyMiddleware { + type Request = tonic::codegen::http::Request; + + fn check(&mut self, request: Self::Request) -> Result { // Grab API key from request let key = // Request header @@ -76,38 +53,11 @@ where let is_allowed = self.auth_keys.can_write(&key) || (is_read_only(&request) && self.auth_keys.can_read(&key)); if is_allowed { - return Box::pin(self.service.call(request)); + return Ok(request); } } - let mut response = Self::Response::new(BoxBody::default()); - *response.status_mut() = StatusCode::FORBIDDEN; - response.headers_mut().append( - "grpc-status", - HeaderValue::from(Code::PermissionDenied as i32), - ); - response - .headers_mut() - .append("grpc-message", HeaderValue::from_static("Invalid api-key")); - - Box::pin(async move { Ok(response) }) - } -} - -impl ApiKeyMiddlewareLayer { - pub fn new(auth_keys: AuthKeys) -> Self { - Self { auth_keys } - } -} - -impl Layer for ApiKeyMiddlewareLayer { - type Service = ApiKeyMiddleware; - - fn layer(&self, service: S) -> Self::Service { - ApiKeyMiddleware { - service, - auth_keys: self.auth_keys.clone(), - } + Err(Box::new(Status::permission_denied("Invalid api-key"))) } } diff --git a/src/tonic/mod.rs b/src/tonic/mod.rs index 5e0cd9120d..bb79e1c7f9 100644 --- a/src/tonic/mod.rs +++ b/src/tonic/mod.rs @@ -194,7 +194,7 @@ pub fn init( telemetry_collector, )) .option_layer({ - AuthKeys::try_create(&settings.service).map(api_key::ApiKeyMiddlewareLayer::new) + AuthKeys::try_create(&settings.service).map(api_key::ApiKeyMiddleware::new_layer) }) .into_inner(); diff --git a/tests/api_key/test_grpc.py b/tests/api_key/test_grpc.py index bc05676fa3..68ba143e02 100644 --- a/tests/api_key/test_grpc.py +++ b/tests/api_key/test_grpc.py @@ -1,3 +1,4 @@ +from grpc import RpcError from qdrant_client import QdrantClient, grpc as qgrpc import pytest from qdrant_client.conversions.conversion import payload_to_grpc @@ -149,5 +150,5 @@ def assert_ro_token_failure(stub, request): try: stub(request, metadata=(("api-key", "my-ro-secret"),), timeout=1.0) pytest.fail("Request should have failed") - except: + except RpcError: return diff --git a/tests/integration-tests-api-key.sh b/tests/integration-tests-api-key.sh index a273018f57..e8aaa579ee 100755 --- a/tests/integration-tests-api-key.sh +++ b/tests/integration-tests-api-key.sh @@ -16,9 +16,6 @@ export QDRANT__SERVICE__READ_ONLY_API_KEY="my-ro-secret" #Capture PID of the process PID=$! -# Sleep to make sure the process has started (workaround for empty pidof) -sleep 5 - function clear_after_tests() { echo "server is going down" @@ -41,4 +38,3 @@ docker run --rm \ -e QDRANT_HOST=host.docker.internal \ --add-host host.docker.internal:host-gateway \ $IMAGE_NAME sh -c "pytest /tests" -