Skip to main content

rpc/
auth.rs

1use rust_zero_core::{
2    decode_jwt_hs256, sign_request, AuthFailure, JwtClaimProjection, RequestSignature,
3    RequestSignatureVerifier, AUTH_KEY_ID_HEADER, AUTH_SIGNATURE_HEADER, AUTH_TIMESTAMP_HEADER,
4};
5use serde::de::DeserializeOwned;
6use std::{
7    collections::BTreeMap,
8    sync::Arc,
9    time::{SystemTime, UNIX_EPOCH},
10};
11
12use tonic::{
13    metadata::{Ascii, MetadataValue},
14    service::Interceptor,
15    Request, Status,
16};
17
18/// Adds a bearer credential to every outgoing RPC request.
19#[derive(Clone)]
20pub struct BearerToken {
21    authorization: MetadataValue<Ascii>,
22}
23
24impl BearerToken {
25    pub fn new(token: &str) -> Result<Self, tonic::metadata::errors::InvalidMetadataValue> {
26        let authorization = format!("Bearer {token}").parse()?;
27        Ok(Self { authorization })
28    }
29
30    pub(crate) fn authorization(&self) -> MetadataValue<Ascii> {
31        self.authorization.clone()
32    }
33}
34
35impl Interceptor for BearerToken {
36    fn call(&mut self, mut request: Request<()>) -> Result<Request<()>, Status> {
37        request
38            .metadata_mut()
39            .insert("authorization", self.authorization.clone());
40        Ok(request)
41    }
42}
43
44type Validator<T> = dyn Fn(&str) -> Option<T> + Send + Sync;
45
46/// Validates bearer credentials on incoming RPC requests.
47pub struct RpcBearerAuth<T> {
48    validator: Arc<Validator<T>>,
49}
50
51impl<T> Clone for RpcBearerAuth<T> {
52    fn clone(&self) -> Self {
53        Self {
54            validator: Arc::clone(&self.validator),
55        }
56    }
57}
58
59impl<T> RpcBearerAuth<T>
60where
61    T: Clone + Send + Sync + 'static,
62{
63    pub fn new(validator: impl Fn(&str) -> Option<T> + Send + Sync + 'static) -> Self {
64        Self {
65            validator: Arc::new(validator),
66        }
67    }
68
69    /// Returns the identity installed in a validated request.
70    pub fn authenticated<U>(request: &Request<U>) -> Option<T> {
71        request.extensions().get::<T>().cloned()
72    }
73}
74
75impl<T> Interceptor for RpcBearerAuth<T>
76where
77    T: Clone + Send + Sync + 'static,
78{
79    fn call(&mut self, mut request: Request<()>) -> Result<Request<()>, Status> {
80        let identity = request
81            .metadata()
82            .get("authorization")
83            .and_then(|value| value.to_str().ok())
84            .and_then(bearer_token)
85            .and_then(|token| (self.validator)(token))
86            .ok_or_else(|| auth_status(AuthFailure::InvalidCredentials))?;
87
88        request.extensions_mut().insert(identity);
89        Ok(request)
90    }
91}
92
93fn bearer_token(value: &str) -> Option<&str> {
94    let (scheme, token) = value.split_once(char::is_whitespace)?;
95    let token = token.trim();
96    (scheme.eq_ignore_ascii_case("bearer")
97        && !token.is_empty()
98        && !token.contains(char::is_whitespace))
99    .then_some(token)
100}
101
102fn auth_status(failure: AuthFailure) -> Status {
103    Status::unauthenticated(format!("{}: {}", failure.code(), failure.message()))
104}
105
106/// Validates HS256 bearer tokens and exposes typed and projected claims to gRPC handlers.
107#[derive(Clone)]
108pub struct RpcJwtAuth<T> {
109    secrets: Vec<Arc<[u8]>>,
110    leeway_seconds: u64,
111    projection: JwtClaimProjection,
112    marker: std::marker::PhantomData<fn() -> T>,
113}
114
115impl<T> RpcJwtAuth<T>
116where
117    T: Clone + Send + Sync + DeserializeOwned + 'static,
118{
119    pub fn new(secret: impl AsRef<[u8]>) -> Self {
120        let secret = secret.as_ref();
121        assert!(!secret.is_empty(), "JWT secret cannot be empty");
122        Self {
123            secrets: vec![Arc::from(secret)],
124            leeway_seconds: 0,
125            projection: JwtClaimProjection::default(),
126            marker: std::marker::PhantomData,
127        }
128    }
129
130    pub fn with_previous_secret(mut self, secret: impl AsRef<[u8]>) -> Self {
131        let secret = secret.as_ref();
132        assert!(!secret.is_empty(), "previous JWT secret cannot be empty");
133        self.secrets.push(Arc::from(secret));
134        self
135    }
136
137    pub fn with_leeway(mut self, seconds: u64) -> Self {
138        self.leeway_seconds = seconds;
139        self
140    }
141
142    pub fn with_claim_projection(mut self, projection: JwtClaimProjection) -> Self {
143        self.projection = projection;
144        self
145    }
146
147    pub fn claims<U>(request: &Request<U>) -> Option<T> {
148        request
149            .extensions()
150            .get::<RpcJwtClaims<T>>()
151            .map(|v| v.0.clone())
152    }
153
154    pub fn projected_claims<U>(
155        request: &Request<U>,
156    ) -> Option<BTreeMap<String, serde_json::Value>> {
157        request
158            .extensions()
159            .get::<RpcProjectedClaims>()
160            .map(|v| v.0.clone())
161    }
162}
163
164#[derive(Clone)]
165struct RpcJwtClaims<T>(T);
166
167#[derive(Clone)]
168struct RpcProjectedClaims(BTreeMap<String, serde_json::Value>);
169
170impl<T> Interceptor for RpcJwtAuth<T>
171where
172    T: Clone + Send + Sync + DeserializeOwned + 'static,
173{
174    fn call(&mut self, mut request: Request<()>) -> Result<Request<()>, Status> {
175        let token = request
176            .metadata()
177            .get("authorization")
178            .and_then(|value| value.to_str().ok())
179            .and_then(bearer_token)
180            .ok_or_else(|| auth_status(AuthFailure::MissingCredentials))?;
181        let claims: T = decode_jwt_hs256(
182            token,
183            &self.secrets,
184            self.leeway_seconds,
185            unix_seconds() as u64,
186        )
187        .map_err(AuthFailure::from)
188        .map_err(auth_status)?;
189        let projected = decode_jwt_hs256::<serde_json::Value>(
190            token,
191            &self.secrets,
192            self.leeway_seconds,
193            unix_seconds() as u64,
194        )
195        .ok()
196        .map(|value| self.projection.project(&value))
197        .unwrap_or_default();
198        request
199            .extensions_mut()
200            .insert(RpcProjectedClaims(projected));
201        request.extensions_mut().insert(RpcJwtClaims(claims));
202        Ok(request)
203    }
204}
205
206/// Adds an HMAC request signature to outgoing gRPC metadata.
207#[derive(Clone)]
208pub struct RpcRequestSigner {
209    key_id: String,
210    secret: Arc<[u8]>,
211    target: String,
212}
213
214impl RpcRequestSigner {
215    pub fn new(
216        key_id: impl Into<String>,
217        secret: impl AsRef<[u8]>,
218        target: impl Into<String>,
219    ) -> Self {
220        Self {
221            key_id: key_id.into(),
222            secret: Arc::from(secret.as_ref()),
223            target: target.into(),
224        }
225    }
226}
227
228impl Interceptor for RpcRequestSigner {
229    fn call(&mut self, mut request: Request<()>) -> Result<Request<()>, Status> {
230        let signature = sign_request(
231            self.key_id.clone(),
232            &self.secret,
233            unix_seconds(),
234            "POST",
235            &self.target,
236        )
237        .map_err(auth_status)?;
238        insert_signature(&mut request, &signature)?;
239        Ok(request)
240    }
241}
242
243/// Validates gRPC request signatures for one canonical service method target.
244#[derive(Clone)]
245pub struct RpcRequestSignatureAuth {
246    verifier: RequestSignatureVerifier,
247    target: String,
248}
249
250impl RpcRequestSignatureAuth {
251    pub fn new(verifier: RequestSignatureVerifier, target: impl Into<String>) -> Self {
252        Self {
253            verifier,
254            target: target.into(),
255        }
256    }
257
258    pub fn key_id<U>(request: &Request<U>) -> Option<String> {
259        request
260            .extensions()
261            .get::<RpcSignatureKeyId>()
262            .map(|id| id.0.clone())
263    }
264}
265
266#[derive(Clone)]
267struct RpcSignatureKeyId(String);
268
269impl Interceptor for RpcRequestSignatureAuth {
270    fn call(&mut self, mut request: Request<()>) -> Result<Request<()>, Status> {
271        let signature = parse_signature(&request)?;
272        self.verifier
273            .verify(&signature, "POST", &self.target, unix_seconds())
274            .map_err(auth_status)?;
275        request
276            .extensions_mut()
277            .insert(RpcSignatureKeyId(signature.key_id));
278        Ok(request)
279    }
280}
281
282#[allow(clippy::result_large_err)] // Tonic interceptors conventionally return `Status` directly.
283fn insert_signature(request: &mut Request<()>, signature: &RequestSignature) -> Result<(), Status> {
284    for (name, value) in [
285        (AUTH_KEY_ID_HEADER, signature.key_id.clone()),
286        (AUTH_TIMESTAMP_HEADER, signature.timestamp.to_string()),
287        (AUTH_SIGNATURE_HEADER, signature.signature.clone()),
288    ] {
289        request.metadata_mut().insert(
290            name,
291            value
292                .parse()
293                .map_err(|_| auth_status(AuthFailure::InvalidSignature))?,
294        );
295    }
296    Ok(())
297}
298
299#[allow(clippy::result_large_err)] // Tonic interceptors conventionally return `Status` directly.
300fn parse_signature(request: &Request<()>) -> Result<RequestSignature, Status> {
301    let value = |name| {
302        request
303            .metadata()
304            .get(name)
305            .and_then(|value| value.to_str().ok())
306    };
307    let key_id =
308        value(AUTH_KEY_ID_HEADER).ok_or_else(|| auth_status(AuthFailure::MissingSignature))?;
309    let timestamp = value(AUTH_TIMESTAMP_HEADER)
310        .ok_or_else(|| auth_status(AuthFailure::MissingSignature))?
311        .parse()
312        .map_err(|_| auth_status(AuthFailure::InvalidSignature))?;
313    let signature =
314        value(AUTH_SIGNATURE_HEADER).ok_or_else(|| auth_status(AuthFailure::MissingSignature))?;
315    Ok(RequestSignature {
316        key_id: key_id.to_owned(),
317        timestamp,
318        signature: signature.to_owned(),
319    })
320}
321
322fn unix_seconds() -> i64 {
323    SystemTime::now()
324        .duration_since(UNIX_EPOCH)
325        .unwrap_or_default()
326        .as_secs() as i64
327}
328
329#[cfg(test)]
330mod tests {
331    use super::*;
332    use serde::{Deserialize, Serialize};
333    use std::time::Duration;
334    use tonic::Code;
335
336    #[test]
337    fn client_and_server_interceptors_exchange_identity() {
338        let mut client = BearerToken::new("valid").unwrap();
339        let request = client.call(Request::new(())).unwrap();
340        let mut server =
341            RpcBearerAuth::new(|token| (token == "valid").then(|| "service-account".to_owned()));
342        let request = server.call(request).unwrap();
343
344        assert_eq!(
345            RpcBearerAuth::<String>::authenticated(&request),
346            Some("service-account".to_owned())
347        );
348    }
349
350    #[test]
351    fn server_rejects_missing_credentials() {
352        let mut server = RpcBearerAuth::new(|_| Some(()));
353        let error = server.call(Request::new(())).unwrap_err();
354
355        assert_eq!(error.code(), Code::Unauthenticated);
356    }
357
358    #[derive(Clone, Deserialize, Serialize)]
359    struct Claims {
360        sub: String,
361        exp: u64,
362    }
363
364    #[test]
365    fn jwt_auth_projects_selected_claims() {
366        let claims = Claims {
367            sub: "service-42".to_owned(),
368            exp: unix_seconds() as u64 + 60,
369        };
370        let token = rust_zero_core::encode_jwt_hs256(&claims, b"secret").unwrap();
371        let mut request = Request::new(());
372        request
373            .metadata_mut()
374            .insert("authorization", format!("Bearer {token}").parse().unwrap());
375        let mut auth = RpcJwtAuth::<Claims>::new("secret").with_claim_projection(
376            JwtClaimProjection::new([("caller".to_owned(), "sub".to_owned())]),
377        );
378        let request = auth.call(request).unwrap();
379
380        assert_eq!(
381            RpcJwtAuth::<Claims>::claims(&request).unwrap().sub,
382            "service-42"
383        );
384        assert_eq!(
385            RpcJwtAuth::<Claims>::projected_claims(&request).unwrap()["caller"],
386            "service-42"
387        );
388    }
389
390    #[test]
391    fn rpc_request_signatures_round_trip_and_bind_the_target() {
392        let verifier = RequestSignatureVerifier::new(
393            [("client".to_owned(), b"secret".to_vec())],
394            Duration::from_secs(30),
395        )
396        .unwrap();
397        let target = "/rust_zero.echo.Echo/Ping";
398        let request = RpcRequestSigner::new("client", "secret", target)
399            .call(Request::new(()))
400            .unwrap();
401        let request = RpcRequestSignatureAuth::new(verifier.clone(), target)
402            .call(request)
403            .unwrap();
404        assert_eq!(
405            RpcRequestSignatureAuth::key_id(&request).as_deref(),
406            Some("client")
407        );
408
409        let request = RpcRequestSigner::new("client", "secret", target)
410            .call(Request::new(()))
411            .unwrap();
412        let error = RpcRequestSignatureAuth::new(verifier, "/other.Service/Call")
413            .call(request)
414            .unwrap_err();
415        assert_eq!(error.code(), Code::Unauthenticated);
416        assert!(error.message().starts_with("auth_invalid_signature:"));
417    }
418}