greentic_runner_host/http/
auth.rs1use 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}