Skip to main content

reqsign_core/
signer.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18use crate::Context;
19use crate::Error;
20use crate::ProvideCredential;
21use crate::ProvideCredentialDyn;
22use crate::Result;
23use crate::SignRequest;
24use crate::SignRequestDyn;
25use crate::SigningCredential;
26use std::any::type_name;
27use std::fmt::{Debug, Formatter};
28use std::sync::{Arc, Mutex};
29use std::time::Duration;
30
31/// Loads credentials and atomically signs request heads.
32///
33/// The service-specific [`SignRequest`] runs against a private candidate. Only the
34/// candidate URI and headers are committed after successful signing.
35#[derive(Clone)]
36pub struct Signer<K: SigningCredential> {
37    ctx: Context,
38    loader: Arc<dyn ProvideCredentialDyn<Credential = K>>,
39    builder: Arc<dyn SignRequestDyn<Credential = K>>,
40    credential: Arc<Mutex<Option<K>>>,
41}
42
43impl<K: SigningCredential> Debug for Signer<K> {
44    fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
45        f.debug_struct("Signer")
46            .field("credential_type", &type_name::<K>())
47            .finish_non_exhaustive()
48    }
49}
50
51impl<K: SigningCredential> Signer<K> {
52    /// Create a new signer.
53    pub fn new(
54        ctx: Context,
55        loader: impl ProvideCredential<Credential = K>,
56        builder: impl SignRequest<Credential = K>,
57    ) -> Self {
58        Self {
59            ctx,
60
61            loader: Arc::new(loader),
62            builder: Arc::new(builder),
63            credential: Arc::new(Mutex::new(None)),
64        }
65    }
66
67    /// Replace the context while keeping credential provider and request signer.
68    pub fn with_context(mut self, ctx: Context) -> Self {
69        self.ctx = ctx;
70        self
71    }
72
73    /// Replace the credential provider while keeping context and request signer.
74    pub fn with_credential_provider(
75        mut self,
76        provider: impl ProvideCredential<Credential = K>,
77    ) -> Self {
78        self.loader = Arc::new(provider);
79        self.credential = Arc::new(Mutex::new(None)); // Clear cached credential
80        self
81    }
82
83    /// Replace the request signer while keeping context and credential provider.
84    pub fn with_request_signer(mut self, signer: impl SignRequest<Credential = K>) -> Self {
85        self.builder = Arc::new(signer);
86        self
87    }
88
89    /// Sign a wire-ready request head.
90    ///
91    /// The request URI must satisfy the input contract of the configured
92    /// [`SignRequest`]. Built-in signers require an authority and expect path and query
93    /// data to be percent-encoded exactly once before this call. Signing does not
94    /// perform general-purpose URI encoding for the caller.
95    ///
96    /// If credential loading or request signing returns an error, `req` is unchanged.
97    /// On success, only `req.uri` and `req.headers` may change; the method, version, and
98    /// extensions retain their input values.
99    ///
100    /// `expires_in` is a service-specific validity input and does not universally
101    /// select query authentication. The configured service signer and credential type
102    /// determine how it is interpreted.
103    ///
104    /// Cached credentials must be fresh according to [`SigningCredential::is_valid`]
105    /// and usable through [`SignRequest::required_valid_until`]. A refreshed credential
106    /// only needs to satisfy the exact operation deadline. Provider errors are returned
107    /// without internal retry or fallback to the previous cached credential.
108    pub async fn sign(
109        &self,
110        req: &mut http::request::Parts,
111        expires_in: Option<Duration>,
112    ) -> Result<()> {
113        let credential = self.credential.lock().expect("lock poisoned").clone();
114        let credential = match credential {
115            Some(credential)
116                if credential.is_valid()
117                    && credential.is_valid_at(
118                        self.builder
119                            .required_valid_until_dyn(&credential, expires_in),
120                    ) =>
121            {
122                credential
123            }
124            _ => {
125                let credential = self
126                    .loader
127                    .provide_credential_dyn(&self.ctx)
128                    .await?
129                    .ok_or_else(|| {
130                        Error::credential_invalid("failed to load signing credential")
131                            .with_context(format!("credential_type: {}", type_name::<K>()))
132                    })?;
133
134                *self.credential.lock().expect("lock poisoned") = Some(credential.clone());
135
136                let required_until = self
137                    .builder
138                    .required_valid_until_dyn(&credential, expires_in);
139                if !credential.is_valid_at(required_until) {
140                    return Err(Error::credential_invalid(
141                        "refreshed signing credential expires before the requested operation deadline",
142                    )
143                    .with_context(format!("credential_type: {}", type_name::<K>()))
144                    .with_context(format!("required_valid_until: {required_until}")));
145                }
146
147                credential
148            }
149        };
150
151        let mut candidate = req.clone();
152        self.builder
153            .sign_request_dyn(&self.ctx, &mut candidate, Some(&credential), expires_in)
154            .await?;
155
156        req.uri = candidate.uri;
157        req.headers = candidate.headers;
158        Ok(())
159    }
160}
161
162#[cfg(test)]
163mod tests {
164    use super::*;
165    use crate::time::Timestamp;
166    use crate::{ErrorKind, ProvideCredential, SignRequest};
167    use http::{HeaderValue, Method, Request, Version};
168    use std::collections::VecDeque;
169    use std::sync::atomic::{AtomicUsize, Ordering};
170
171    #[derive(Clone, Debug)]
172    struct TestCredential;
173
174    impl SigningCredential for TestCredential {
175        fn is_valid(&self) -> bool {
176            true
177        }
178    }
179
180    #[derive(Debug)]
181    struct StaticProvider;
182
183    impl ProvideCredential for StaticProvider {
184        type Credential = TestCredential;
185
186        async fn provide_credential(&self, _ctx: &Context) -> Result<Option<Self::Credential>> {
187            Ok(Some(TestCredential))
188        }
189    }
190
191    #[derive(Clone, Debug, PartialEq, Eq)]
192    struct Extension(&'static str);
193
194    #[derive(Debug)]
195    struct MutatingSigner {
196        fail: bool,
197    }
198
199    impl SignRequest for MutatingSigner {
200        type Credential = TestCredential;
201
202        async fn sign_request(
203            &self,
204            _ctx: &Context,
205            req: &mut http::request::Parts,
206            _credential: Option<&Self::Credential>,
207            _expires_in: Option<Duration>,
208        ) -> Result<()> {
209            req.method = Method::POST;
210            req.uri = "https://signed.example.com/result?auth=1"
211                .parse()
212                .expect("URI must parse");
213            req.version = Version::HTTP_2;
214            req.headers.clear();
215            req.headers
216                .insert("authorization", HeaderValue::from_static("signed"));
217            req.extensions.insert(Extension("candidate"));
218
219            if self.fail {
220                Err(Error::unexpected("injected signing failure"))
221            } else {
222                Ok(())
223            }
224        }
225    }
226
227    #[derive(Clone, Debug)]
228    struct ExpiringCredential {
229        generation: u8,
230        fresh: bool,
231        expires_at: Timestamp,
232        required_until: Timestamp,
233    }
234
235    impl SigningCredential for ExpiringCredential {
236        fn is_valid(&self) -> bool {
237            self.fresh
238        }
239
240        fn is_valid_at(&self, timestamp: Timestamp) -> bool {
241            self.expires_at > timestamp
242        }
243    }
244
245    #[derive(Debug)]
246    struct SequenceProvider {
247        responses: Mutex<VecDeque<Result<Option<ExpiringCredential>>>>,
248        calls: Arc<AtomicUsize>,
249    }
250
251    impl SequenceProvider {
252        fn new(
253            responses: impl IntoIterator<Item = Result<Option<ExpiringCredential>>>,
254        ) -> (Self, Arc<AtomicUsize>) {
255            let calls = Arc::new(AtomicUsize::new(0));
256            (
257                Self {
258                    responses: Mutex::new(responses.into_iter().collect()),
259                    calls: calls.clone(),
260                },
261                calls,
262            )
263        }
264    }
265
266    impl ProvideCredential for SequenceProvider {
267        type Credential = ExpiringCredential;
268
269        async fn provide_credential(&self, _ctx: &Context) -> Result<Option<Self::Credential>> {
270            self.calls.fetch_add(1, Ordering::SeqCst);
271            self.responses
272                .lock()
273                .expect("lock poisoned")
274                .pop_front()
275                .unwrap_or(Ok(None))
276        }
277    }
278
279    #[derive(Debug)]
280    struct OperationSigner;
281
282    impl SignRequest for OperationSigner {
283        type Credential = ExpiringCredential;
284
285        fn required_valid_until(
286            &self,
287            credential: &Self::Credential,
288            _expires_in: Option<Duration>,
289        ) -> Timestamp {
290            credential.required_until
291        }
292
293        async fn sign_request(
294            &self,
295            _ctx: &Context,
296            req: &mut http::request::Parts,
297            credential: Option<&Self::Credential>,
298            expires_in: Option<Duration>,
299        ) -> Result<()> {
300            let credential = credential.expect("credential must be present");
301            if !credential.is_valid_at(self.required_valid_until(credential, expires_in)) {
302                return Err(Error::credential_invalid(
303                    "credential is not valid for operation",
304                ));
305            }
306            req.headers.insert(
307                "x-credential-generation",
308                credential.generation.to_string().parse()?,
309            );
310            Ok(())
311        }
312    }
313
314    fn request_parts() -> http::request::Parts {
315        let mut parts = Request::get("https://example.com/original?x=%2F")
316            .version(Version::HTTP_11)
317            .header("x-original", "value")
318            .body(())
319            .expect("request must build")
320            .into_parts()
321            .0;
322        parts.extensions.insert(Extension("caller"));
323        parts
324    }
325
326    #[test]
327    fn failure_leaves_entire_request_head_unchanged() {
328        let signer = Signer::new(
329            Context::new(),
330            StaticProvider,
331            MutatingSigner { fail: true },
332        );
333        let mut parts = request_parts();
334        let original = parts.clone();
335
336        let result = futures::executor::block_on(signer.sign(&mut parts, None));
337
338        assert!(result.is_err());
339        assert_eq!(parts.method, original.method);
340        assert_eq!(parts.uri, original.uri);
341        assert_eq!(parts.version, original.version);
342        assert_eq!(parts.headers, original.headers);
343        assert_eq!(
344            parts.extensions.get::<Extension>(),
345            original.extensions.get::<Extension>()
346        );
347    }
348
349    #[test]
350    fn success_commits_only_uri_and_headers() {
351        let signer = Signer::new(
352            Context::new(),
353            StaticProvider,
354            MutatingSigner { fail: false },
355        );
356        let mut parts = request_parts();
357        let original = parts.clone();
358
359        futures::executor::block_on(signer.sign(&mut parts, None)).expect("signing must succeed");
360
361        assert_eq!(parts.method, original.method);
362        assert_eq!(parts.version, original.version);
363        assert_eq!(
364            parts.extensions.get::<Extension>(),
365            original.extensions.get::<Extension>()
366        );
367        assert_eq!(
368            parts.uri,
369            "https://signed.example.com/result?auth=1"
370                .parse::<http::Uri>()
371                .expect("URI must parse")
372        );
373        assert_eq!(
374            parts.headers.get("authorization"),
375            Some(&HeaderValue::from_static("signed"))
376        );
377        assert!(!parts.headers.contains_key("x-original"));
378    }
379
380    #[test]
381    fn refreshes_cached_credential_for_operation_requirement() {
382        let base = Timestamp::from_second(1_000).expect("timestamp must be valid");
383        let cached = ExpiringCredential {
384            generation: 1,
385            fresh: true,
386            expires_at: base + Duration::from_secs(20),
387            required_until: base + Duration::from_secs(30),
388        };
389        let refreshed = ExpiringCredential {
390            generation: 2,
391            fresh: true,
392            expires_at: base + Duration::from_secs(20),
393            required_until: base + Duration::from_secs(10),
394        };
395        let (provider, calls) = SequenceProvider::new([Ok(Some(refreshed))]);
396        let signer = Signer::new(Context::new(), provider, OperationSigner);
397        *signer.credential.lock().expect("lock poisoned") = Some(cached);
398
399        let mut parts = request_parts();
400        futures::executor::block_on(signer.sign(&mut parts, None))
401            .expect("refreshed credential must satisfy the recomputed requirement");
402
403        assert_eq!(calls.load(Ordering::SeqCst), 1);
404        assert_eq!(
405            parts.headers.get("x-credential-generation"),
406            Some(&HeaderValue::from_static("2"))
407        );
408    }
409
410    #[test]
411    fn uses_refreshed_credential_that_is_usable_but_not_fresh() {
412        let base = Timestamp::from_second(2_000).expect("timestamp must be valid");
413        let credential = ExpiringCredential {
414            generation: 1,
415            fresh: false,
416            expires_at: base + Duration::from_secs(30),
417            required_until: base + Duration::from_secs(10),
418        };
419        let (provider, calls) =
420            SequenceProvider::new([Ok(Some(credential.clone())), Ok(Some(credential))]);
421        let signer = Signer::new(Context::new(), provider, OperationSigner);
422
423        for _ in 0..2 {
424            let mut parts = request_parts();
425            futures::executor::block_on(signer.sign(&mut parts, None))
426                .expect("usable refreshed credential must be accepted");
427        }
428
429        assert_eq!(calls.load(Ordering::SeqCst), 2);
430    }
431
432    #[test]
433    fn refresh_error_does_not_fall_back_and_caller_can_retry() {
434        let base = Timestamp::from_second(3_000).expect("timestamp must be valid");
435        let cached = ExpiringCredential {
436            generation: 1,
437            fresh: false,
438            expires_at: base + Duration::from_secs(30),
439            required_until: base + Duration::from_secs(10),
440        };
441        let refreshed = ExpiringCredential {
442            generation: 2,
443            fresh: true,
444            expires_at: base + Duration::from_secs(30),
445            required_until: base + Duration::from_secs(10),
446        };
447        let (provider, calls) = SequenceProvider::new([
448            Err(Error::unexpected("injected refresh failure")),
449            Ok(Some(refreshed)),
450        ]);
451        let signer = Signer::new(Context::new(), provider, OperationSigner);
452        *signer.credential.lock().expect("lock poisoned") = Some(cached);
453
454        let mut parts = request_parts();
455        let original = parts.clone();
456        let err = futures::executor::block_on(signer.sign(&mut parts, None))
457            .expect_err("refresh error must be returned");
458        assert_eq!(err.kind(), ErrorKind::Unexpected);
459        assert_eq!(parts.uri, original.uri);
460        assert_eq!(parts.headers, original.headers);
461        assert_eq!(calls.load(Ordering::SeqCst), 1);
462
463        futures::executor::block_on(signer.sign(&mut parts, None))
464            .expect("caller retry must attempt refresh again");
465        assert_eq!(calls.load(Ordering::SeqCst), 2);
466        assert_eq!(
467            parts.headers.get("x-credential-generation"),
468            Some(&HeaderValue::from_static("2"))
469        );
470    }
471
472    #[test]
473    fn missing_refresh_does_not_fall_back_and_caller_can_retry() {
474        let base = Timestamp::from_second(4_000).expect("timestamp must be valid");
475        let cached = ExpiringCredential {
476            generation: 1,
477            fresh: false,
478            expires_at: base + Duration::from_secs(30),
479            required_until: base + Duration::from_secs(10),
480        };
481        let refreshed = ExpiringCredential {
482            generation: 2,
483            fresh: true,
484            expires_at: base + Duration::from_secs(30),
485            required_until: base + Duration::from_secs(10),
486        };
487        let (provider, calls) = SequenceProvider::new([Ok(None), Ok(Some(refreshed))]);
488        let signer = Signer::new(Context::new(), provider, OperationSigner);
489        *signer.credential.lock().expect("lock poisoned") = Some(cached);
490        let mut parts = request_parts();
491        let original = parts.clone();
492
493        let err = futures::executor::block_on(signer.sign(&mut parts, None))
494            .expect_err("missing credential must fail");
495
496        assert_eq!(err.kind(), ErrorKind::CredentialInvalid);
497        assert_eq!(calls.load(Ordering::SeqCst), 1);
498        assert_eq!(parts.uri, original.uri);
499        assert_eq!(parts.headers, original.headers);
500
501        futures::executor::block_on(signer.sign(&mut parts, None))
502            .expect("caller retry must attempt refresh again");
503        assert_eq!(calls.load(Ordering::SeqCst), 2);
504        assert_eq!(
505            parts.headers.get("x-credential-generation"),
506            Some(&HeaderValue::from_static("2"))
507        );
508    }
509
510    #[test]
511    fn debug_is_opaque() {
512        let signer = Signer::new(
513            Context::new(),
514            StaticProvider,
515            MutatingSigner { fail: false },
516        );
517        *signer.credential.lock().expect("lock poisoned") = Some(TestCredential);
518
519        let debug = format!("{signer:?}");
520        assert!(debug.starts_with("Signer"));
521        assert!(!debug.contains("StaticProvider"));
522        assert!(!debug.contains("MutatingSigner"));
523        assert!(!debug.contains("credential:"));
524    }
525}