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#[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
46pub 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 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#[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#[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#[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)] fn 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)] fn 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}