Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1,110 changes: 0 additions & 1,110 deletions model_gateway/src/middleware.rs

This file was deleted.

58 changes: 58 additions & 0 deletions model_gateway/src/middleware/auth.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,58 @@
//! Bearer-token auth middleware backed by a precomputed SHA-256 hash.
//!
//! The hash is compared in constant time so the comparison cost does not
//! leak the configured key length.

use axum::{
body::Body,
extract::{Request, State},
http::{header, StatusCode},
middleware::Next,
response::{IntoResponse, Response},
};
use sha2::{Digest, Sha256};
use subtle::ConstantTimeEq;

#[derive(Clone)]
pub struct AuthConfig {
/// Precomputed SHA-256 hash of the API key, used for constant-time comparison
/// that doesn't leak key length via timing.
api_key_hash: Option<[u8; 32]>,
}

impl AuthConfig {
pub fn new(api_key: Option<String>) -> Self {
Self {
api_key_hash: api_key.map(|k| Sha256::digest(k.as_bytes()).into()),
}
}
}

/// Middleware to validate Bearer token against configured API key.
/// Only active when router has an API key configured.
pub async fn auth_middleware(
State(auth_config): State<AuthConfig>,
request: Request<Body>,
next: Next,
) -> Response {
if let Some(expected_hash) = &auth_config.api_key_hash {
let token = request
.headers()
.get(header::AUTHORIZATION)
.and_then(|h| h.to_str().ok())
.and_then(|h| h.strip_prefix("Bearer "));

let authorized = token.is_some_and(|t| {
Sha256::digest(t.as_bytes())
.as_slice()
.ct_eq(expected_hash)
.unwrap_u8()
== 1
});
Comment on lines +39 to +51

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🧹 Nitpick | 🔵 Trivial

Consider case-insensitive "Bearer" prefix matching.

The strip_prefix("Bearer ") at line 43 is case-sensitive. Per RFC 6750, the "Bearer" scheme is case-insensitive. While most clients send "Bearer", some might send "bearer" or "BEARER".

♻️ Optional: case-insensitive Bearer prefix
         let token = request
             .headers()
             .get(header::AUTHORIZATION)
             .and_then(|h| h.to_str().ok())
-            .and_then(|h| h.strip_prefix("Bearer "));
+            .and_then(|h| {
+                h.strip_prefix("Bearer ")
+                    .or_else(|| h.strip_prefix("bearer "))
+                    .or_else(|| h.strip_prefix("BEARER "))
+            });

Or more robustly:

.and_then(|h| {
    if h.len() > 7 && h[..7].eq_ignore_ascii_case("bearer ") {
        Some(&h[7..])
    } else {
        None
    }
})
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.

In `@model_gateway/src/middleware/auth.rs` around lines 39 - 51, The current
strip_prefix("Bearer ") call is case-sensitive and can miss valid tokens; update
the token extraction around request.headers().get(header::AUTHORIZATION) to
perform a case-insensitive "bearer " check (e.g., test h.len() > 7 and
h[..7].eq_ignore_ascii_case("bearer ") then take &h[7..]) instead of
strip_prefix, so the token variable, the authorized computation (which uses
Sha256::digest and expected_hash), and the existing authorized logic remain
unchanged but now accept any ASCII case variation of the Bearer scheme.

if !authorized {
return StatusCode::UNAUTHORIZED.into_response();
}
}

next.run(request).await
}
301 changes: 301 additions & 0 deletions model_gateway/src/middleware/concurrency.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,301 @@
//! Per-request concurrency limiting via a token bucket, with optional
//! queuing for backpressure.
//!
//! `ConcurrencyLimiter` wires a bounded `mpsc` channel that
//! `concurrency_limit_middleware` uses to enqueue requests when the
//! bucket is empty; `QueueProcessor` drains that channel and hands tokens
//! back to waiters. `TokenGuardBody` wraps the response body so the token
//! is only released after the entire stream has been delivered.

use std::{
pin::Pin,
sync::Arc,
task::{Context, Poll},
time::{Duration, Instant},
};

use axum::{
body::Body,
extract::{Request, State},
http::StatusCode,
middleware::Next,
response::{IntoResponse, Response},
Json,
};
use bytes::Bytes;
use http_body::Frame;
use serde_json::json;
use tokio::sync::{mpsc, oneshot};
use tracing::{debug, error, warn};

use super::token_bucket::TokenBucket;
use crate::{
observability::metrics::{metrics_labels, Metrics},
server::AppState,
};

/// A body wrapper that holds a token and returns it when the body is fully consumed or dropped.
/// This ensures that for streaming responses, the token is only returned after the entire
/// stream has been sent to the client.
pub struct TokenGuardBody {
inner: Body,
/// The token bucket to return tokens to. Uses Option so we can take() on drop.
token_bucket: Option<Arc<TokenBucket>>,
/// Number of tokens to return.
tokens: f64,
}

impl TokenGuardBody {
/// Create a new TokenGuardBody that will return tokens when dropped.
pub fn new(inner: Body, token_bucket: Arc<TokenBucket>, tokens: f64) -> Self {
Self {
inner,
token_bucket: Some(token_bucket),
tokens,
}
}
}

impl Drop for TokenGuardBody {
fn drop(&mut self) {
if let Some(bucket) = self.token_bucket.take() {
debug!(
"TokenGuardBody: stream ended, returning {} tokens to bucket",
self.tokens
);
// Use lock-free sync return - no runtime needed, guaranteed token return
bucket.return_tokens_sync(self.tokens);
}
}
}

impl http_body::Body for TokenGuardBody {
type Data = Bytes;
type Error = axum::Error;

fn poll_frame(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
// SAFETY: We never move the inner body, and Body is Unpin
// (it's a type alias for UnsyncBoxBody which is Unpin)
let this = self.get_mut();
Pin::new(&mut this.inner).poll_frame(cx)
}
Comment on lines +76 to +84

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🧹 Nitpick | 🔵 Trivial

Clarify the pin projection safety comment.

The comment at lines 80-81 states "SAFETY" but this isn't an unsafe block. Per repository conventions from learnings, SAFETY: should be reserved for unsafe blocks, while INVARIANT: documents assumptions in safe code. Since Body being Unpin is an invariant the code relies on, consider updating the comment.

♻️ Use INVARIANT marker
     fn poll_frame(
         self: Pin<&mut Self>,
         cx: &mut Context<'_>,
     ) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
-        // SAFETY: We never move the inner body, and Body is Unpin
-        // (it's a type alias for UnsyncBoxBody which is Unpin)
+        // INVARIANT: Body is Unpin (type alias for UnsyncBoxBody), so we can
+        // safely project through Pin without needing unsafe.
         let this = self.get_mut();
         Pin::new(&mut this.inner).poll_frame(cx)
     }

Based on learnings: "In Rust code across the repository, use the marker INVARIANT: to document assumptions in safe code. Reserve SAFETY: for explaining why unsafe blocks are sound."

📝 Committable suggestion

‼️ IMPORTANT
Carefully review the code before committing. Ensure that it accurately replaces the highlighted code, contains no missing lines, and has no issues with indentation. Thoroughly test & benchmark the code to ensure it meets the requirements.

Suggested change
fn poll_frame(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
// SAFETY: We never move the inner body, and Body is Unpin
// (it's a type alias for UnsyncBoxBody which is Unpin)
let this = self.get_mut();
Pin::new(&mut this.inner).poll_frame(cx)
}
fn poll_frame(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
// INVARIANT: Body is Unpin (type alias for UnsyncBoxBody), so we can
// safely project through Pin without needing unsafe.
let this = self.get_mut();
Pin::new(&mut this.inner).poll_frame(cx)
}
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.

In `@model_gateway/src/middleware/concurrency.rs` around lines 76 - 84, The
comment on poll_frame should use the repository convention for non-unsafe
assumptions: replace the "SAFETY:" marker with "INVARIANT:" and keep the
explanatory text that the inner body is never moved and that Body is Unpin (type
alias for UnsyncBoxBody which is Unpin); update the comment above Pin::new(&mut
this.inner).poll_frame(cx) to read "INVARIANT: We never move the inner body, and
Body is Unpin (it's a type alias for UnsyncBoxBody which is Unpin)".


fn is_end_stream(&self) -> bool {
self.inner.is_end_stream()
}

fn size_hint(&self) -> http_body::SizeHint {
self.inner.size_hint()
}
}

/// Request queue entry
pub struct QueuedRequest {
/// Time when the request was queued
queued_at: Instant,
/// Channel to send the permit back when acquired
permit_tx: oneshot::Sender<Result<(), StatusCode>>,
}

/// Queue processor that handles queued requests
pub struct QueueProcessor {
token_bucket: Arc<TokenBucket>,
queue_rx: mpsc::Receiver<QueuedRequest>,
queue_timeout: Duration,
}

impl QueueProcessor {
pub fn new(
token_bucket: Arc<TokenBucket>,
queue_rx: mpsc::Receiver<QueuedRequest>,
queue_timeout: Duration,
) -> Self {
Self {
token_bucket,
queue_rx,
queue_timeout,
}
}

pub async fn run(mut self) {
debug!("Starting concurrency queue processor");

// Process requests in a single task to reduce overhead
while let Some(queued) = self.queue_rx.recv().await {
// Check timeout immediately
let elapsed = queued.queued_at.elapsed();
if elapsed >= self.queue_timeout {
warn!("Request already timed out in queue");
let _ = queued.permit_tx.send(Err(StatusCode::REQUEST_TIMEOUT));
continue;
}

let remaining_timeout = self.queue_timeout - elapsed;

// Try to acquire token for this request
if self.token_bucket.try_acquire(1.0).is_ok() {
// Got token immediately
debug!("Queue: acquired token immediately for queued request");
let _ = queued.permit_tx.send(Ok(()));
} else {
// Need to wait for token
let token_bucket = self.token_bucket.clone();

// Spawn task only when we actually need to wait
#[expect(
clippy::disallowed_methods,
reason = "fire-and-forget permit acquisition: task is bounded by remaining_timeout and communicates via oneshot; dropping the JoinHandle detaches the task but it self-terminates"
)]
tokio::spawn(async move {
if token_bucket
.acquire_timeout(1.0, remaining_timeout)
.await
.is_ok()
{
debug!("Queue: acquired token after waiting");
let _ = queued.permit_tx.send(Ok(()));
} else {
warn!("Queue: request timed out waiting for token");
let _ = queued.permit_tx.send(Err(StatusCode::REQUEST_TIMEOUT));
}
});
Comment on lines +139 to +164

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

The current implementation of the queue processor can leak tokens if a request is cancelled while waiting in the queue or immediately after a token is acquired. If permit_tx.send() fails (because the receiver was dropped due to request cancellation), the acquired token is never returned to the bucket. Additionally, the spawned task should monitor for request cancellation while waiting for a token to avoid unnecessary acquisition. Note that data passed to the spawned task must have a 'static lifetime, so we use owned types or clones to ensure the data outlives the task.

            // Try to acquire token for this request
            if self.token_bucket.try_acquire(1.0).is_ok() {
                // Got token immediately
                debug!("Queue: acquired token immediately for queued request");
                if queued.permit_tx.send(Ok(())).is_err() {
                    debug!("Queue: request cancelled, returning token");
                    self.token_bucket.return_tokens_sync(1.0);
                }
            } else {
                // Need to wait for token
                let token_bucket = self.token_bucket.clone();
                let mut permit_tx = queued.permit_tx;

                // Spawn task only when we actually need to wait
                #[expect(
                    clippy::disallowed_methods,
                    reason = "fire-and-forget permit acquisition: task is bounded by remaining_timeout and communicates via oneshot; dropping the JoinHandle detaches the task but it self-terminates"
                )]
                tokio::spawn(async move {
                    tokio::select! {
                        _ = permit_tx.closed() => {
                            debug!("Queue: request cancelled while waiting for token");
                        }
                        res = token_bucket.acquire_timeout(1.0, remaining_timeout) => {
                            match res {
                                Ok(()) => {
                                    debug!("Queue: acquired token after waiting");
                                    if permit_tx.send(Ok(())).is_err() {
                                        debug!("Queue: request cancelled after acquisition, returning token");
                                        token_bucket.return_tokens_sync(1.0);
                                    }
                                }
                                Err(_) => {
                                    warn!("Queue: request timed out waiting for token");
                                    let _ = permit_tx.send(Err(StatusCode::REQUEST_TIMEOUT));
                                }
                            }
                        }
                    }
                });
            }
References
  1. Data passed to spawned background tasks must have a 'static lifetime. Use owned types or reference-counted pointers like Arc instead of passing references to ensure the data outlives the task.

}
}

warn!("Concurrency queue processor shutting down");
}
}

/// State for the concurrency limiter
pub struct ConcurrencyLimiter {
pub queue_tx: Option<mpsc::Sender<QueuedRequest>>,
}

impl ConcurrencyLimiter {
/// Create new concurrency limiter with optional queue
pub fn new(
token_bucket: Option<Arc<TokenBucket>>,
queue_size: usize,
queue_timeout: Duration,
) -> (Self, Option<QueueProcessor>) {
match (token_bucket, queue_size) {
(None, _) => (Self { queue_tx: None }, None),
(Some(bucket), size) if size > 0 => {
let (queue_tx, queue_rx) = mpsc::channel(size);
let processor = QueueProcessor::new(bucket, queue_rx, queue_timeout);
(
Self {
queue_tx: Some(queue_tx),
},
Some(processor),
)
}
(Some(_), _) => (Self { queue_tx: None }, None),
}
}
}

/// Middleware function for concurrency limiting with optional queuing
pub async fn concurrency_limit_middleware(
State(app_state): State<Arc<AppState>>,
request: Request<Body>,
next: Next,
) -> Response {
// Check mesh global rate limit first if mesh is enabled
// If mesh is not enabled, this check is skipped and local rate limiting is used
if let Some(mesh_handler) = &app_state.mesh_handler {
let (is_exceeded, current_count, limit) =
mesh_handler.sync_manager.check_global_rate_limit();
if is_exceeded {
debug!(
"Global rate limit exceeded: {}/{} req/s",
current_count, limit
);
return (
StatusCode::TOO_MANY_REQUESTS,
Json(json!({
"error": "Rate limit exceeded",
"current_count": current_count,
"limit": limit
})),
)
.into_response();
}
}

let token_bucket = match &app_state.context.rate_limiter {
Some(bucket) => bucket.clone(),
None => {
// Rate limiting disabled, pass through immediately
return next.run(request).await;
}
};

// Try to acquire token immediately
if token_bucket.try_acquire(1.0).is_ok() {
debug!("Acquired token immediately");
Metrics::record_http_rate_limit(metrics_labels::RATE_LIMIT_ALLOWED);
let response = next.run(request).await;

// Wrap the response body with TokenGuardBody to return token when stream ends
// This ensures that for streaming responses, the token is only returned
// after the entire stream has been sent to the client.
let (parts, body) = response.into_parts();
let guarded_body = TokenGuardBody::new(body, token_bucket, 1.0);
Response::from_parts(parts, Body::new(guarded_body))
Comment on lines +238 to +248

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

There is a potential token leak in concurrency_limit_middleware. If the request is cancelled (e.g., the client disconnects) during the execution of next.run(request).await, the TokenGuardBody is never created, and the token acquired at line 238 is leaked. To fix this, you should use a RAII guard that returns the token on drop unless it is explicitly disarmed after the response is successfully wrapped in TokenGuardBody.

} else {
// No tokens available, try to queue if enabled
if let Some(queue_tx) = &app_state.concurrency_queue_tx {
debug!("No tokens available, attempting to queue request");

// Create a channel for the token response
let (permit_tx, permit_rx) = oneshot::channel();

let queued = QueuedRequest {
queued_at: Instant::now(),
permit_tx,
};

// Try to send to queue
match queue_tx.try_send(queued) {
Ok(()) => {
// Wait for token from queue processor
match permit_rx.await {
Ok(Ok(())) => {
debug!("Acquired token from queue");
Metrics::record_http_rate_limit(metrics_labels::RATE_LIMIT_ALLOWED);
let response = next.run(request).await;

// Wrap the response body with TokenGuardBody to return token when stream ends
let (parts, body) = response.into_parts();
let guarded_body = TokenGuardBody::new(body, token_bucket, 1.0);
Response::from_parts(parts, Body::new(guarded_body))
}
Ok(Err(status)) => {
warn!("Queue returned error status: {}", status);
Metrics::record_http_rate_limit(metrics_labels::RATE_LIMIT_REJECTED);
status.into_response()
}
Err(_) => {
error!("Queue response channel closed");
Metrics::record_http_rate_limit(metrics_labels::RATE_LIMIT_REJECTED);
StatusCode::INTERNAL_SERVER_ERROR.into_response()
}
}
}
Err(_) => {
warn!("Request queue is full, returning 429");
Metrics::record_http_rate_limit(metrics_labels::RATE_LIMIT_REJECTED);
StatusCode::TOO_MANY_REQUESTS.into_response()
}
}
} else {
warn!("No tokens available and queuing is disabled, returning 429");
Metrics::record_http_rate_limit(metrics_labels::RATE_LIMIT_REJECTED);
StatusCode::TOO_MANY_REQUESTS.into_response()
}
}
}
Loading
Loading