Skip to main content

openkind_api/middleware/
auth.rs

1//! Bearer token authentication middleware and timing-safe verification.
2
3use std::sync::Arc;
4
5use axum::{
6    body::Body,
7    extract::State,
8    http::{HeaderName, HeaderValue, Request, StatusCode},
9    middleware::Next,
10    response::{IntoResponse, Response},
11    Json,
12};
13
14use super::request_id::{RequestId, REQUEST_ID_HEADER};
15
16/// Authorization header.
17pub const AUTH_HEADER: HeaderName = HeaderName::from_static("authorization");
18
19/// Optional bearer-auth state. `None` ⇒ no auth required.
20#[derive(Clone, Default)]
21pub struct AuthConfig {
22    /// Expected Bearer API key token wrapped in an `Arc`. If `None`, authentication is disabled.
23    pub expected: Arc<Option<String>>,
24    /// SHA-256 digest of the expected token, computed once at construction so
25    /// the per-request comparison only hashes the supplied token.
26    expected_digest: Option<(Arc<str>, [u8; 32])>,
27}
28
29impl AuthConfig {
30    /// Construct a new `AuthConfig` with the specified optional expected API key token.
31    pub fn new(expected: Option<String>) -> Self {
32        let expected_digest = expected
33            .as_deref()
34            .map(|token| (Arc::from(token), digest_of(token)));
35        Self {
36            expected: Arc::new(expected),
37            expected_digest,
38        }
39    }
40
41    /// Resolve an API key by consulting environment lookup closure.
42    ///
43    /// `OPENDECISION_API_KEY` and `OPENPICK_API_KEY` remain deprecated fallbacks
44    /// so that upgrading a deployment cannot silently disable authentication.
45    pub fn resolve_api_key_with<F>(get_env: F) -> Option<String>
46    where
47        F: Fn(&str) -> Result<String, std::env::VarError>,
48    {
49        get_env("OPENKIND_API_KEY")
50            .ok()
51            .filter(|s| !s.is_empty())
52            .or_else(|| {
53                get_env("OPENDECISION_API_KEY")
54                    .ok()
55                    .filter(|s| !s.is_empty())
56            })
57            .or_else(|| get_env("TYPESAFE_API_KEY").ok().filter(|s| !s.is_empty()))
58            .or_else(|| get_env("OPENPICK_API_KEY").ok().filter(|s| !s.is_empty()))
59    }
60
61    /// Construct `AuthConfig` by resolving from the supported API-key environment variables.
62    pub fn from_env() -> Self {
63        Self::new(Self::resolve_api_key_with(|k| std::env::var(k)))
64    }
65
66    /// Returns `true` if authentication is required (an expected API key is configured).
67    pub fn is_required(&self) -> bool {
68        self.expected.is_some()
69    }
70
71    pub(crate) fn token_matches(&self, supplied: &str) -> bool {
72        use subtle::ConstantTimeEq;
73        let Some(expected) = self.expected.as_deref() else {
74            return false;
75        };
76        // `expected` is public and can be replaced or edited through its Arc.
77        // Reuse the digest only while its configuration snapshot still matches.
78        let expected_digest = self
79            .expected_digest
80            .as_ref()
81            .filter(|(cached, _)| cached.as_ref() == expected)
82            .map(|(_, digest)| *digest)
83            .unwrap_or_else(|| digest_of(expected));
84        digest_of(supplied).ct_eq(&expected_digest).into()
85    }
86}
87
88/// Stackable middleware function: gate `/v1/*` requests on a bearer
89/// token when one is configured. `/health` and `/metrics` are always
90/// open so probes and scrapers don't need credentials. `/playground` is
91/// likewise open when the route is enabled: it serves an inert HTML shell,
92/// and evaluation plus `/playground/api/*` model controls remain gated
93/// (the UI collects an optional API key for those calls).
94pub async fn auth_layer(
95    State(auth): State<AuthConfig>,
96    req: Request<Body>,
97    next: Next,
98) -> Response {
99    auth_layer_with_rate_limit(State((auth, super::RateLimiter::disabled())), req, next).await
100}
101
102pub(crate) async fn auth_layer_with_rate_limit(
103    State((auth, failed_auth)): State<(AuthConfig, super::RateLimiter)>,
104    req: Request<Body>,
105    next: Next,
106) -> Response {
107    let path = req.uri().path();
108    if !auth.is_required()
109        || path == "/health"
110        || path == "/metrics"
111        || path == "/playground"
112        || req.method() == axum::http::Method::OPTIONS
113    {
114        return next.run(req).await;
115    }
116
117    let supplied = req
118        .headers()
119        .get(&AUTH_HEADER)
120        .and_then(|v| v.to_str().ok())
121        .and_then(|s| {
122            let (scheme, token) = s.split_once(' ')?;
123            scheme.eq_ignore_ascii_case("Bearer").then_some(token)
124        });
125
126    let ok = supplied.is_some_and(|token| auth.token_matches(token));
127
128    if !ok {
129        metrics::counter!("openkind_auth_failures_total", "transport" => "http").increment(1);
130        if let Some(peer) = req
131            .extensions()
132            .get::<axum::extract::ConnectInfo<std::net::SocketAddr>>()
133        {
134            if let Err(retry_after_ms) = failed_auth.check(peer.0.ip()) {
135                return crate::ApiError::RateLimited { retry_after_ms }.into_response();
136            }
137        }
138        let body = Json(serde_json::json!({
139            "error": {
140                "code": "unauthorized",
141                "message": "missing or invalid API key",
142            }
143        }));
144        let mut resp = (StatusCode::UNAUTHORIZED, body).into_response();
145        if let Ok(v) = HeaderValue::from_str("Bearer") {
146            resp.headers_mut()
147                .insert(axum::http::header::WWW_AUTHENTICATE, v);
148        }
149        // Also stamp the request id on the 401.
150        if let Some(req_id) = req.extensions().get::<RequestId>() {
151            if let Ok(v) = HeaderValue::from_str(&req_id.0) {
152                resp.headers_mut().insert(REQUEST_ID_HEADER.clone(), v);
153            }
154        }
155        return resp;
156    }
157
158    next.run(req).await
159}
160
161/// SHA-256 digest of one token as a fixed 32-byte array.
162fn digest_of(token: &str) -> [u8; 32] {
163    let digest = ring::digest::digest(&ring::digest::SHA256, token.as_bytes());
164    let mut out = [0u8; 32];
165    out.copy_from_slice(digest.as_ref());
166    out
167}
168
169/// Secure constant-time token comparison.
170///
171/// To completely eliminate timing side-channels (including length-leakage attacks),
172/// both inputs are hashed using SHA-256 into fixed 32-byte digests, and the digests
173/// are compared in constant time using `subtle::ConstantTimeEq`.
174pub fn secure_token_eq(a: &str, b: &str) -> bool {
175    use subtle::ConstantTimeEq;
176    digest_of(a).ct_eq(&digest_of(b)).into()
177}
178
179/// Dummy route handler used when attaching authentication middleware as an independent router layer.
180async fn auth_layer_dummy_handler() -> StatusCode {
181    StatusCode::OK
182}
183
184/// Build the auth middleware as a Layer for use with `.layer()`.
185pub fn auth_layer_for(auth: AuthConfig) -> axum::Router {
186    axum::Router::new()
187        // Dummy root route to attach auth middleware layer.
188        .route("/", axum::routing::get(auth_layer_dummy_handler))
189        .layer(axum::middleware::from_fn_with_state(auth, auth_layer))
190}