Skip to main content

objectstore_server/extractors/
id.rs

1use std::borrow::Cow;
2
3use axum::extract::rejection::PathRejection;
4use axum::extract::{FromRequestParts, Path};
5use axum::http::request::Parts;
6use axum::response::{IntoResponse, Response};
7use objectstore_service::id::{ObjectContext, ObjectId};
8use objectstore_types::scope::{EMPTY_SCOPES, Scope, Scopes};
9use serde::{Deserialize, de};
10
11use crate::extractors::Xt;
12use crate::extractors::downstream_service::DownstreamService;
13use crate::state::ServiceState;
14
15#[derive(Debug)]
16pub enum ObjectRejection {
17    Path(PathRejection),
18    Killswitched,
19    RateLimited,
20}
21
22impl IntoResponse for ObjectRejection {
23    fn into_response(self) -> Response {
24        match self {
25            ObjectRejection::Path(rejection) => rejection.into_response(),
26            ObjectRejection::Killswitched => (
27                http::StatusCode::FORBIDDEN,
28                "Object access is disabled for this scope through killswitches",
29            )
30                .into_response(),
31            ObjectRejection::RateLimited => (
32                http::StatusCode::TOO_MANY_REQUESTS,
33                "Object access is rate limited",
34            )
35                .into_response(),
36        }
37    }
38}
39
40impl From<PathRejection> for ObjectRejection {
41    fn from(rejection: PathRejection) -> Self {
42        ObjectRejection::Path(rejection)
43    }
44}
45
46impl FromRequestParts<ServiceState> for Xt<ObjectId> {
47    type Rejection = ObjectRejection;
48
49    async fn from_request_parts(
50        parts: &mut Parts,
51        state: &ServiceState,
52    ) -> Result<Self, Self::Rejection> {
53        let Path(params) = Path::<ObjectParams>::from_request_parts(parts, state).await?;
54        let id = ObjectId::from_parts(params.usecase, params.scopes, params.key);
55
56        populate_sentry_context(id.context());
57        sentry::configure_scope(|s| s.set_extra("key", id.key().into()));
58
59        let service = DownstreamService::from_request_parts(parts, state)
60            .await
61            .unwrap();
62
63        if state
64            .config
65            .killswitches
66            .matches(id.context(), service.as_str())
67        {
68            return Err(ObjectRejection::Killswitched);
69        }
70
71        if !state.rate_limiter.check(id.context(), Some(id.key())) {
72            return Err(ObjectRejection::RateLimited);
73        }
74
75        Ok(Xt(id))
76    }
77}
78
79/// Path parameters used for object-level endpoints.
80///
81/// This is meant to be used with the axum `Path` extractor.
82#[derive(Clone, Debug, Deserialize)]
83struct ObjectParams {
84    usecase: String,
85    #[serde(deserialize_with = "deserialize_scopes")]
86    scopes: Scopes,
87    key: String,
88}
89
90/// Deserializes a `Scopes` instance from a string representation.
91///
92/// The string representation is a semicolon-separated list of `key=value` pairs, following the
93/// Matrix URIs proposal. An empty scopes string (`"_"`) represents no scopes.
94fn deserialize_scopes<'de, D>(deserializer: D) -> Result<Scopes, D::Error>
95where
96    D: de::Deserializer<'de>,
97{
98    let s = Cow::<str>::deserialize(deserializer)?;
99    if s == EMPTY_SCOPES {
100        return Ok(Scopes::empty());
101    }
102
103    let scopes = s
104        .split(';')
105        .map(|s| {
106            let (key, value) = s
107                .split_once('=')
108                .ok_or_else(|| de::Error::custom("scope must be 'key=value'"))?;
109
110            Scope::create(key, value).map_err(de::Error::custom)
111        })
112        .collect::<Result<_, _>>()?;
113
114    Ok(scopes)
115}
116
117impl FromRequestParts<ServiceState> for Xt<ObjectContext> {
118    type Rejection = ObjectRejection;
119
120    async fn from_request_parts(
121        parts: &mut Parts,
122        state: &ServiceState,
123    ) -> Result<Self, Self::Rejection> {
124        let Path(params) = Path::<ContextParams>::from_request_parts(parts, state).await?;
125        let context = ObjectContext {
126            usecase: params.usecase,
127            scopes: params.scopes,
128        };
129
130        populate_sentry_context(&context);
131
132        let service = DownstreamService::from_request_parts(parts, state)
133            .await
134            .unwrap();
135
136        if state
137            .config
138            .killswitches
139            .matches(&context, service.as_str())
140        {
141            return Err(ObjectRejection::Killswitched);
142        }
143
144        if !state.rate_limiter.check(&context, None) {
145            return Err(ObjectRejection::RateLimited);
146        }
147
148        Ok(Xt(context))
149    }
150}
151
152/// Path parameters for extracting an [`ObjectContext`] from a request path.
153///
154/// Works on both collection-level (`/objects/{usecase}/{scopes}`) and object-level
155/// (`/objects/{usecase}/{scopes}/{*key}`) routes — the extra `key` parameter is ignored.
156#[derive(Clone, Debug, Deserialize)]
157pub(super) struct ContextParams {
158    pub usecase: String,
159    #[serde(deserialize_with = "deserialize_scopes")]
160    pub scopes: Scopes,
161}
162
163fn populate_sentry_context(context: &ObjectContext) {
164    sentry::configure_scope(|s| {
165        s.set_tag("usecase", &context.usecase);
166        for scope in &context.scopes {
167            s.set_tag(&format!("scope.{}", scope.name()), scope.value());
168        }
169    });
170}
171
172#[cfg(test)]
173mod tests {
174    use super::*;
175    use serde::de::IntoDeserializer;
176    use serde::de::value::{CowStrDeserializer, Error as DeError};
177    use std::borrow::Cow;
178
179    fn deser_scopes(input: &str) -> Result<Scopes, DeError> {
180        let deserializer: CowStrDeserializer<DeError> = Cow::Borrowed(input).into_deserializer();
181        deserialize_scopes(deserializer)
182    }
183
184    #[test]
185    fn parse_single_scope() {
186        let scopes = deser_scopes("org=123").unwrap();
187        assert_eq!(scopes.get_value("org"), Some("123"));
188    }
189
190    #[test]
191    fn parse_multiple_scopes() {
192        let scopes = deser_scopes("org=123;project=456").unwrap();
193        assert_eq!(scopes.get_value("org"), Some("123"));
194        assert_eq!(scopes.get_value("project"), Some("456"));
195    }
196
197    #[test]
198    fn parse_empty_scopes() {
199        let scopes = deser_scopes("_").unwrap();
200        assert!(scopes.is_empty());
201    }
202
203    #[test]
204    fn parse_missing_equals() {
205        let result = deser_scopes("org123");
206        assert!(result.is_err());
207    }
208
209    #[test]
210    fn parse_invalid_scope_chars() {
211        let result = deser_scopes("org=hello world");
212        assert!(result.is_err());
213    }
214
215    #[test]
216    fn parse_empty_key_or_value() {
217        assert!(deser_scopes("=value").is_err());
218        assert!(deser_scopes("key=").is_err());
219    }
220
221    // --- Extractor integration tests ---
222
223    use std::collections::BTreeMap;
224    use std::sync::Arc;
225
226    use axum::Router;
227    use axum::body::Body;
228    use axum::http::{Request, StatusCode};
229    use axum::routing::{get, post};
230    use objectstore_service::StorageService;
231    use objectstore_service::backend::in_memory::InMemoryBackend;
232    use objectstore_service::encryption::Cipher;
233    use tower::ServiceExt;
234
235    use crate::auth::PublicKeyDirectory;
236    use crate::config::Config;
237    use crate::killswitches::{Killswitch, Killswitches};
238    use crate::rate_limits::{RateLimiter, RateLimits, ThroughputLimits};
239    use crate::state::{ServiceState, Services};
240    use crate::web::RequestCounter;
241
242    async fn test_state(config: Config) -> ServiceState {
243        let service = StorageService::new(
244            Box::new(InMemoryBackend::new("in-memory")),
245            Cipher::ephemeral().unwrap(),
246        );
247        let key_directory = Arc::new(PublicKeyDirectory::from_config(&config.auth).await.unwrap());
248        let rate_limiter = RateLimiter::new(config.rate_limits.clone());
249
250        Arc::new(Services {
251            config,
252            service,
253            key_directory,
254            rate_limiter,
255            request_counter: RequestCounter::new(0),
256        })
257    }
258
259    async fn handle_object_id(Xt(id): Xt<ObjectId>) -> String {
260        format!(
261            "usecase={} key={} scopes_empty={}",
262            id.context().usecase,
263            id.key(),
264            id.context().scopes.is_empty(),
265        )
266    }
267
268    async fn handle_object_context(Xt(ctx): Xt<ObjectContext>) -> String {
269        format!(
270            "usecase={} scopes_empty={}",
271            ctx.usecase,
272            ctx.scopes.is_empty(),
273        )
274    }
275
276    fn test_router(state: ServiceState) -> Router {
277        Router::new()
278            .route("/objects/{usecase}/{scopes}/{*key}", get(handle_object_id))
279            .route("/objects/{usecase}/{scopes}/", post(handle_object_context))
280            .with_state(state)
281    }
282
283    async fn response_body(response: http::Response<Body>) -> String {
284        let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
285            .await
286            .unwrap();
287        String::from_utf8(bytes.to_vec()).unwrap()
288    }
289
290    // Extraction tests
291
292    #[tokio::test]
293    async fn extract_object_id_parses_path() {
294        let state = test_state(Config::default()).await;
295        let app = test_router(state);
296
297        let request = Request::builder()
298            .uri("/objects/myusecase/org=123;project=456/my-key")
299            .body(Body::empty())
300            .unwrap();
301        let response = app.oneshot(request).await.unwrap();
302
303        assert_eq!(response.status(), StatusCode::OK);
304        let body = response_body(response).await;
305        assert!(body.contains("usecase=myusecase"));
306        assert!(body.contains("key=my-key"));
307        assert!(body.contains("scopes_empty=false"));
308    }
309
310    #[tokio::test]
311    async fn extract_object_id_with_empty_scopes() {
312        let state = test_state(Config::default()).await;
313        let app = test_router(state);
314
315        let request = Request::builder()
316            .uri("/objects/myusecase/_/my-key")
317            .body(Body::empty())
318            .unwrap();
319        let response = app.oneshot(request).await.unwrap();
320
321        assert_eq!(response.status(), StatusCode::OK);
322        let body = response_body(response).await;
323        assert!(body.contains("scopes_empty=true"));
324    }
325
326    #[tokio::test]
327    async fn extract_object_context_parses_path() {
328        let state = test_state(Config::default()).await;
329        let app = test_router(state);
330
331        let request = Request::builder()
332            .method("POST")
333            .uri("/objects/myusecase/org=123;project=456/")
334            .body(Body::empty())
335            .unwrap();
336        let response = app.oneshot(request).await.unwrap();
337
338        assert_eq!(response.status(), StatusCode::OK);
339        let body = response_body(response).await;
340        assert!(body.contains("usecase=myusecase"));
341        assert!(body.contains("scopes_empty=false"));
342    }
343
344    #[tokio::test]
345    async fn extract_object_context_with_empty_scopes() {
346        let state = test_state(Config::default()).await;
347        let app = test_router(state);
348
349        let request = Request::builder()
350            .method("POST")
351            .uri("/objects/myusecase/_/")
352            .body(Body::empty())
353            .unwrap();
354        let response = app.oneshot(request).await.unwrap();
355
356        assert_eq!(response.status(), StatusCode::OK);
357        let body = response_body(response).await;
358        assert!(body.contains("scopes_empty=true"));
359    }
360
361    #[tokio::test]
362    async fn extract_object_id_invalid_scopes() {
363        let state = test_state(Config::default()).await;
364        let app = test_router(state);
365
366        let request = Request::builder()
367            .uri("/objects/myusecase/invalid-no-equals/key")
368            .body(Body::empty())
369            .unwrap();
370        let response = app.oneshot(request).await.unwrap();
371
372        assert_eq!(response.status(), StatusCode::BAD_REQUEST);
373    }
374
375    // Killswitch tests
376
377    #[tokio::test]
378    async fn extract_object_id_killswitched() {
379        let config = Config {
380            killswitches: Killswitches::new(vec![Killswitch {
381                usecase: Some("blocked".into()),
382                scopes: BTreeMap::new(),
383                service: None,
384            }]),
385            ..Config::default()
386        };
387        let state = test_state(config).await;
388        let app = test_router(state);
389
390        let request = Request::builder()
391            .uri("/objects/blocked/org=1/key")
392            .body(Body::empty())
393            .unwrap();
394        let response = app.clone().oneshot(request).await.unwrap();
395        assert_eq!(response.status(), StatusCode::FORBIDDEN);
396
397        let request = Request::builder()
398            .uri("/objects/allowed/org=1/key")
399            .body(Body::empty())
400            .unwrap();
401        let response = app.oneshot(request).await.unwrap();
402        assert_eq!(response.status(), StatusCode::OK);
403    }
404
405    #[tokio::test]
406    async fn extract_object_context_killswitched() {
407        let config = Config {
408            killswitches: Killswitches::new(vec![Killswitch {
409                usecase: Some("blocked".into()),
410                scopes: BTreeMap::new(),
411                service: None,
412            }]),
413            ..Config::default()
414        };
415        let state = test_state(config).await;
416        let app = test_router(state);
417
418        let request = Request::builder()
419            .method("POST")
420            .uri("/objects/blocked/org=1/")
421            .body(Body::empty())
422            .unwrap();
423        let response = app.clone().oneshot(request).await.unwrap();
424        assert_eq!(response.status(), StatusCode::FORBIDDEN);
425
426        let request = Request::builder()
427            .method("POST")
428            .uri("/objects/allowed/org=1/")
429            .body(Body::empty())
430            .unwrap();
431        let response = app.oneshot(request).await.unwrap();
432        assert_eq!(response.status(), StatusCode::OK);
433    }
434
435    #[tokio::test]
436    async fn extract_object_id_killswitched_with_service() {
437        let config = Config {
438            killswitches: Killswitches::new(vec![Killswitch {
439                usecase: None,
440                scopes: BTreeMap::new(),
441                service: Some("test-*".into()),
442            }]),
443            ..Config::default()
444        };
445        let state = test_state(config).await;
446        let app = test_router(state);
447
448        // Matching service header → 403
449        let request = Request::builder()
450            .uri("/objects/any/org=1/key")
451            .header("x-downstream-service", "test-service")
452            .body(Body::empty())
453            .unwrap();
454        let response = app.clone().oneshot(request).await.unwrap();
455        assert_eq!(response.status(), StatusCode::FORBIDDEN);
456
457        // Non-matching service header → 200
458        let request = Request::builder()
459            .uri("/objects/any/org=1/key")
460            .header("x-downstream-service", "other-service")
461            .body(Body::empty())
462            .unwrap();
463        let response = app.clone().oneshot(request).await.unwrap();
464        assert_eq!(response.status(), StatusCode::OK);
465
466        // No service header → 200
467        let request = Request::builder()
468            .uri("/objects/any/org=1/key")
469            .body(Body::empty())
470            .unwrap();
471        let response = app.oneshot(request).await.unwrap();
472        assert_eq!(response.status(), StatusCode::OK);
473    }
474
475    // Rate limiter tests
476
477    #[tokio::test]
478    async fn extract_object_id_rate_limited() {
479        let config = Config {
480            rate_limits: RateLimits {
481                throughput: ThroughputLimits {
482                    global_rps: Some(1),
483                    burst: 0,
484                    ..ThroughputLimits::default()
485                },
486                ..RateLimits::default()
487            },
488            ..Config::default()
489        };
490        let state = test_state(config).await;
491        let app = test_router(state);
492
493        let request = Request::builder()
494            .uri("/objects/test/org=1/key")
495            .body(Body::empty())
496            .unwrap();
497        let response = app.clone().oneshot(request).await.unwrap();
498        assert_eq!(response.status(), StatusCode::OK);
499
500        let request = Request::builder()
501            .uri("/objects/test/org=1/key")
502            .body(Body::empty())
503            .unwrap();
504        let response = app.oneshot(request).await.unwrap();
505        assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS);
506    }
507
508    #[tokio::test]
509    async fn extract_object_context_rate_limited() {
510        let config = Config {
511            rate_limits: RateLimits {
512                throughput: ThroughputLimits {
513                    global_rps: Some(1),
514                    burst: 0,
515                    ..ThroughputLimits::default()
516                },
517                ..RateLimits::default()
518            },
519            ..Config::default()
520        };
521        let state = test_state(config).await;
522        let app = test_router(state);
523
524        let request = Request::builder()
525            .method("POST")
526            .uri("/objects/test/org=1/")
527            .body(Body::empty())
528            .unwrap();
529        let response = app.clone().oneshot(request).await.unwrap();
530        assert_eq!(response.status(), StatusCode::OK);
531
532        let request = Request::builder()
533            .method("POST")
534            .uri("/objects/test/org=1/")
535            .body(Body::empty())
536            .unwrap();
537        let response = app.oneshot(request).await.unwrap();
538        assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS);
539    }
540}