Skip to main content

greentic_runner_host/http/
auth.rs

1use std::net::SocketAddr;
2
3use axum::Json;
4use axum::extract::connect_info::ConnectInfo;
5use axum::extract::{FromRef, FromRequestParts};
6use axum::http::StatusCode;
7use axum::http::header::AUTHORIZATION;
8use axum::http::request::Parts;
9use serde_json::json;
10
11use crate::runner::ServerState;
12
13#[derive(Clone, Default)]
14pub struct AdminAuth {
15    token: Option<String>,
16}
17
18impl AdminAuth {
19    pub fn new(token: Option<String>) -> Self {
20        Self {
21            token: token.filter(|v| !v.is_empty()),
22        }
23    }
24
25    fn authorize(&self, addr: SocketAddr, bearer: Option<&str>) -> Result<(), StatusCode> {
26        if let Some(expected) = &self.token {
27            let token = bearer.ok_or(StatusCode::UNAUTHORIZED)?;
28            if constant_time_eq(token.as_bytes(), expected.as_bytes()) {
29                Ok(())
30            } else {
31                Err(StatusCode::UNAUTHORIZED)
32            }
33        } else if addr.ip().is_loopback() {
34            Ok(())
35        } else {
36            Err(StatusCode::FORBIDDEN)
37        }
38    }
39}
40
41pub struct AdminGuard;
42
43impl<S> FromRequestParts<S> for AdminGuard
44where
45    ServerState: FromRef<S>,
46    S: Send + Sync,
47{
48    type Rejection = (StatusCode, Json<serde_json::Value>);
49
50    fn from_request_parts(
51        parts: &mut Parts,
52        state: &S,
53    ) -> impl std::future::Future<Output = Result<Self, Self::Rejection>> + Send {
54        let server_state = ServerState::from_ref(state);
55        let admin = server_state.admin.clone();
56        let addr = parts
57            .extensions
58            .get::<ConnectInfo<SocketAddr>>()
59            .map(|info| info.0);
60        let bearer = extract_bearer(parts);
61
62        async move {
63            let addr = addr.ok_or((
64                StatusCode::INTERNAL_SERVER_ERROR,
65                Json(json!({ "error": "connect info unavailable" })),
66            ))?;
67            admin.authorize(addr, bearer.as_deref()).map_err(|status| {
68                (
69                    status,
70                    Json(json!({
71                        "error": if status == StatusCode::UNAUTHORIZED {
72                            "admin token required"
73                        } else {
74                            "admin access restricted"
75                        }
76                    })),
77                )
78            })?;
79            Ok(AdminGuard)
80        }
81    }
82}
83
84fn extract_bearer(parts: &Parts) -> Option<String> {
85    let header = parts.headers.get(AUTHORIZATION)?.to_str().ok()?;
86    let (scheme, value) = header.split_once(' ')?;
87    if !scheme.eq_ignore_ascii_case("Bearer") {
88        return None;
89    }
90    Some(value.trim().to_string())
91}
92
93fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
94    if a.len() != b.len() {
95        return false;
96    }
97    let mut diff = 0u8;
98    for (&left, &right) in a.iter().zip(b.iter()) {
99        diff |= left ^ right;
100    }
101    diff == 0
102}
103
104#[cfg(test)]
105mod tests {
106    use super::*;
107    use axum::extract::FromRef;
108    use axum::extract::connect_info::ConnectInfo;
109    use axum::http::Request;
110    use std::sync::Arc;
111
112    use crate::http::health::HealthState;
113    use crate::routing::{RoutingConfig, TenantRouting};
114    use crate::runner::ServerState;
115    use crate::runtime::ActivePacks;
116
117    #[derive(Clone)]
118    struct AppState {
119        server: ServerState,
120    }
121
122    impl FromRef<AppState> for ServerState {
123        fn from_ref(input: &AppState) -> Self {
124            input.server.clone()
125        }
126    }
127
128    fn server_state(admin: AdminAuth) -> AppState {
129        AppState {
130            server: ServerState {
131                active: Arc::new(ActivePacks::new()),
132                routing: TenantRouting::new(RoutingConfig::default()),
133                health: Arc::new(HealthState::new()),
134                reload: None,
135                admin,
136                host: crate::host::RunnerHost::for_test(),
137                sql: crate::sql::SqlGateway::new(std::collections::HashMap::new(), String::new()),
138            },
139        }
140    }
141
142    #[test]
143    fn loopback_without_token_is_allowed() {
144        let auth = AdminAuth::new(None);
145        assert!(auth.authorize("127.0.0.1:0".parse().unwrap(), None).is_ok());
146    }
147
148    #[test]
149    fn remote_without_token_is_forbidden() {
150        let auth = AdminAuth::new(None);
151        assert_eq!(
152            auth.authorize("10.0.0.1:0".parse().unwrap(), None),
153            Err(StatusCode::FORBIDDEN)
154        );
155    }
156
157    #[test]
158    fn token_requires_bearer() {
159        let auth = AdminAuth {
160            token: Some("secret".into()),
161        };
162        assert_eq!(
163            auth.authorize("127.0.0.1:0".parse().unwrap(), None),
164            Err(StatusCode::UNAUTHORIZED)
165        );
166        assert!(
167            auth.authorize("127.0.0.1:0".parse().unwrap(), Some("secret"))
168                .is_ok()
169        );
170    }
171
172    #[test]
173    fn bearer_scheme_is_case_insensitive() {
174        let (parts, _) = axum::http::Request::builder()
175            .header(AUTHORIZATION, "bearer secret")
176            .body(())
177            .expect("request")
178            .into_parts();
179
180        assert_eq!(extract_bearer(&parts).as_deref(), Some("secret"));
181    }
182
183    #[test]
184    fn non_bearer_authorization_header_is_rejected() {
185        let (parts, _) = axum::http::Request::builder()
186            .header(AUTHORIZATION, "Basic dXNlcjpzZWNyZXQ=")
187            .body(())
188            .expect("request")
189            .into_parts();
190
191        assert_eq!(extract_bearer(&parts), None);
192    }
193
194    #[test]
195    fn empty_admin_token_is_treated_as_disabled() {
196        let auth = AdminAuth::new(Some(String::new()));
197        assert!(auth.authorize("127.0.0.1:0".parse().unwrap(), None).is_ok());
198    }
199
200    #[test]
201    fn wrong_bearer_token_is_rejected() {
202        let auth = AdminAuth::new(Some("secret".into()));
203        assert_eq!(
204            auth.authorize("127.0.0.1:0".parse().unwrap(), Some("wrong")),
205            Err(StatusCode::UNAUTHORIZED)
206        );
207    }
208
209    #[test]
210    fn constant_time_eq_rejects_length_mismatch() {
211        assert!(!constant_time_eq(b"short", b"longer"));
212    }
213
214    #[test]
215    fn malformed_authorization_header_is_rejected() {
216        let (parts, _) = axum::http::Request::builder()
217            .header(AUTHORIZATION, "Bearer")
218            .body(())
219            .expect("request")
220            .into_parts();
221
222        assert_eq!(extract_bearer(&parts), None);
223    }
224
225    #[tokio::test]
226    async fn admin_guard_rejects_missing_connect_info() {
227        let (mut parts, _) = Request::builder().body(()).expect("request").into_parts();
228        let state = server_state(AdminAuth::default());
229
230        let rejection = match AdminGuard::from_request_parts(&mut parts, &state).await {
231            Ok(_) => panic!("missing connect info should reject"),
232            Err(rejection) => rejection,
233        };
234
235        assert_eq!(rejection.0, StatusCode::INTERNAL_SERVER_ERROR);
236        assert_eq!(rejection.1.0["error"], "connect info unavailable");
237    }
238
239    #[tokio::test]
240    async fn admin_guard_rejects_wrong_remote_token() {
241        let (mut parts, _) = Request::builder()
242            .header(AUTHORIZATION, "Bearer wrong")
243            .body(())
244            .expect("request")
245            .into_parts();
246        parts.extensions.insert(ConnectInfo(
247            "10.0.0.2:8080".parse::<std::net::SocketAddr>().unwrap(),
248        ));
249        let state = server_state(AdminAuth::new(Some("secret".into())));
250
251        let rejection = match AdminGuard::from_request_parts(&mut parts, &state).await {
252            Ok(_) => panic!("wrong token should reject"),
253            Err(rejection) => rejection,
254        };
255
256        assert_eq!(rejection.0, StatusCode::UNAUTHORIZED);
257        assert_eq!(rejection.1.0["error"], "admin token required");
258    }
259
260    #[tokio::test]
261    async fn admin_guard_allows_loopback_without_token_when_disabled() {
262        let (mut parts, _) = Request::builder().body(()).expect("request").into_parts();
263        parts.extensions.insert(ConnectInfo(
264            "127.0.0.1:8080".parse::<std::net::SocketAddr>().unwrap(),
265        ));
266        let state = server_state(AdminAuth::default());
267
268        AdminGuard::from_request_parts(&mut parts, &state)
269            .await
270            .expect("loopback should pass without token");
271    }
272}