1use alloy::primitives::keccak256;
4use async_trait::async_trait;
5use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
6use k256::ecdsa::Signature;
7use once_cell::sync::Lazy;
8use reqwest::Client;
9use serde_json::Value;
10use std::{
11 collections::HashMap,
12 env,
13 time::{Duration, Instant},
14};
15use tokio::fs;
16use tokio::sync::RwLock;
17
18#[cfg(test)]
19use mockall::automock;
20
21use crate::{
22 models::{Address, AzureKeyVaultAuthType, AzureKeyVaultSignerConfig},
23 utils::{recover_public_key, recover_public_key_from_hash, Secp256k1Error},
24};
25
26const AZURE_API_VERSION: &str = "7.4";
27const AZURE_SCOPE: &str = "https://vault.azure.net/.default";
28const AZURE_MANAGED_IDENTITY_RESOURCE: &str = "https://vault.azure.net";
29const AZURE_SIGN_ALGORITHM: &str = "ES256K";
30const AZURE_IMDS_API_VERSION: &str = "2018-02-01";
31const AZURE_IMDS_TOKEN_URL: &str = "http://169.254.169.254/metadata/identity/oauth2/token";
32const AZURE_HTTP_CONNECT_TIMEOUT: Duration = Duration::from_secs(5);
33const AZURE_HTTP_REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
34const AZURE_HTTP_POOL_IDLE_TIMEOUT: Duration = Duration::from_secs(90);
35const AZURE_TOKEN_CACHE_TTL_FALLBACK: Duration = Duration::from_secs(300);
36const AZURE_TOKEN_CACHE_REFRESH_BUFFER: Duration = Duration::from_secs(60);
37
38#[derive(Debug, thiserror::Error, serde::Serialize)]
39pub enum AzureKeyVaultError {
40 #[error("Azure Key Vault HTTP error: {0}")]
41 HttpError(String),
42 #[error("Azure Key Vault API error: {0}")]
43 ApiError(String),
44 #[error("Azure Key Vault response parse error: {0}")]
45 ParseError(String),
46 #[error("Azure Key Vault missing field: {0}")]
47 MissingField(String),
48 #[error("Azure Key Vault recovery error: {0}")]
49 RecoveryError(#[from] Secp256k1Error),
50}
51
52pub type AzureKeyVaultResult<T> = Result<T, AzureKeyVaultError>;
53
54#[async_trait]
55#[cfg_attr(test, automock)]
56pub trait AzureKeyVaultEvmService: Send + Sync {
57 async fn get_evm_address(&self) -> AzureKeyVaultResult<Address>;
59 async fn sign_payload_evm(&self, payload: &[u8]) -> AzureKeyVaultResult<Vec<u8>>;
61 async fn sign_hash_evm(&self, hash: &[u8; 32]) -> AzureKeyVaultResult<Vec<u8>>;
63}
64
65#[derive(Clone, Debug)]
66pub struct AzureKeyVaultService {
67 pub config: AzureKeyVaultSignerConfig,
68 client: Client,
69}
70
71#[derive(Clone, Debug, Eq, Hash, PartialEq)]
72struct AzureAccessTokenCacheKey {
73 auth_type: String,
74 tenant_id: String,
75 client_id: String,
76 vault_url: String,
77 federated_token_file: Option<String>,
78}
79
80#[derive(Clone, Debug, Eq, Hash, PartialEq)]
81struct AzurePublicKeyCacheKey {
82 vault_url: String,
83 key_path: String,
84}
85
86#[derive(Clone, Debug)]
87struct AzureAccessTokenCacheEntry {
88 token: String,
89 expires_at: Instant,
90}
91
92static AZURE_ACCESS_TOKEN_CACHE: Lazy<
94 RwLock<HashMap<AzureAccessTokenCacheKey, AzureAccessTokenCacheEntry>>,
95> = Lazy::new(|| RwLock::new(HashMap::new()));
96
97static AZURE_PUBLIC_KEY_CACHE: Lazy<RwLock<HashMap<AzurePublicKeyCacheKey, [u8; 64]>>> =
99 Lazy::new(|| RwLock::new(HashMap::new()));
100
101impl AzureKeyVaultService {
102 pub fn new(config: &AzureKeyVaultSignerConfig) -> AzureKeyVaultResult<Self> {
104 Ok(Self {
105 config: config.clone(),
106 client: Self::build_http_client()?,
107 })
108 }
109
110 fn build_http_client() -> AzureKeyVaultResult<Client> {
112 Client::builder()
113 .connect_timeout(AZURE_HTTP_CONNECT_TIMEOUT)
114 .timeout(AZURE_HTTP_REQUEST_TIMEOUT)
115 .pool_idle_timeout(AZURE_HTTP_POOL_IDLE_TIMEOUT)
116 .build()
117 .map_err(|e| AzureKeyVaultError::HttpError(e.to_string()))
118 }
119
120 fn tenant_id(&self) -> String {
122 self.config
123 .tenant_id
124 .as_ref()
125 .map(|value| value.to_str().to_string())
126 .unwrap_or_default()
127 }
128
129 fn client_id(&self) -> String {
131 self.config
132 .client_id
133 .as_ref()
134 .map(|value| value.to_str().to_string())
135 .unwrap_or_default()
136 }
137
138 fn client_secret(&self) -> String {
140 self.config
141 .client_secret
142 .as_ref()
143 .map(|value| value.to_str().to_string())
144 .unwrap_or_default()
145 }
146
147 fn federated_token_file(&self) -> Option<String> {
149 self.config
150 .federated_token_file
151 .as_ref()
152 .map(|value| value.to_str().to_string())
153 .or_else(|| env::var("AZURE_FEDERATED_TOKEN_FILE").ok())
154 }
155
156 fn vault_url(&self) -> String {
158 self.config
159 .vault_url
160 .to_str()
161 .trim_end_matches('/')
162 .to_string()
163 }
164
165 fn key_name(&self) -> String {
167 self.config.key_name.to_str().to_string()
168 }
169
170 fn oauth_token_url(&self) -> String {
172 let tenant_id = self.tenant_id();
173 if tenant_id.starts_with("http://") || tenant_id.starts_with("https://") {
174 format!("{}/oauth2/v2.0/token", tenant_id.trim_end_matches('/'))
175 } else {
176 format!(
177 "https://login.microsoftonline.com/{}/oauth2/v2.0/token",
178 tenant_id
179 )
180 }
181 }
182
183 fn key_path(&self) -> String {
185 match self.config.key_version.as_deref() {
186 Some(version) if !version.is_empty() => format!("keys/{}/{}", self.key_name(), version),
187 _ => format!("keys/{}", self.key_name()),
188 }
189 }
190
191 fn key_url(&self) -> String {
193 format!(
194 "{}/{}?api-version={}",
195 self.vault_url(),
196 self.key_path(),
197 AZURE_API_VERSION
198 )
199 }
200
201 fn sign_url(&self) -> String {
203 format!(
204 "{}/{}/sign?api-version={}",
205 self.vault_url(),
206 self.key_path(),
207 AZURE_API_VERSION
208 )
209 }
210
211 fn auth_type(&self) -> AzureKeyVaultAuthType {
213 self.config.auth_type()
214 }
215
216 fn auth_type_cache_key(&self) -> String {
218 match self.auth_type() {
219 AzureKeyVaultAuthType::ClientSecret => "client_secret",
220 AzureKeyVaultAuthType::ManagedIdentity => "managed_identity",
221 AzureKeyVaultAuthType::WorkloadIdentity => "workload_identity",
222 }
223 .to_string()
224 }
225
226 fn managed_identity_token_url(&self) -> String {
228 env::var("AZURE_IMDS_TOKEN_URL").unwrap_or_else(|_| AZURE_IMDS_TOKEN_URL.to_string())
229 }
230
231 fn access_token_cache_key(&self) -> AzureAccessTokenCacheKey {
233 AzureAccessTokenCacheKey {
234 auth_type: self.auth_type_cache_key(),
235 tenant_id: self.tenant_id(),
236 client_id: self.client_id(),
237 vault_url: self.vault_url(),
238 federated_token_file: self.federated_token_file(),
239 }
240 }
241
242 fn public_key_cache_key(&self) -> AzurePublicKeyCacheKey {
244 AzurePublicKeyCacheKey {
245 vault_url: self.vault_url(),
246 key_path: self.key_path(),
247 }
248 }
249
250 async fn get_client_secret_access_token(
252 &self,
253 ) -> AzureKeyVaultResult<AzureAccessTokenCacheEntry> {
254 let client_id = self.client_id();
255 let client_secret = self.client_secret();
256 let url = self.oauth_token_url();
257
258 let response = self
259 .client
260 .post(url)
261 .form(&[
262 ("grant_type", "client_credentials"),
263 ("client_id", client_id.as_str()),
264 ("client_secret", client_secret.as_str()),
265 ("scope", AZURE_SCOPE),
266 ])
267 .send()
268 .await
269 .map_err(|e| AzureKeyVaultError::HttpError(e.to_string()))?;
270
271 Self::parse_access_token_response(response).await
272 }
273
274 async fn get_managed_identity_access_token(
276 &self,
277 ) -> AzureKeyVaultResult<AzureAccessTokenCacheEntry> {
278 let request = self
279 .client
280 .get(self.managed_identity_token_url())
281 .header("Metadata", "true");
282 let client_id = self.client_id();
283
284 let mut query = vec![
285 ("api-version", AZURE_IMDS_API_VERSION),
286 ("resource", AZURE_MANAGED_IDENTITY_RESOURCE),
287 ];
288 if !client_id.is_empty() {
289 query.push(("client_id", client_id.as_str()));
290 }
291
292 let response = request
293 .query(&query)
294 .send()
295 .await
296 .map_err(|e| AzureKeyVaultError::HttpError(e.to_string()))?;
297
298 Self::parse_access_token_response(response).await
299 }
300
301 async fn get_workload_identity_access_token(
303 &self,
304 ) -> AzureKeyVaultResult<AzureAccessTokenCacheEntry> {
305 let client_id = self.client_id();
306 let url = self.oauth_token_url();
307 let token_file = self
308 .federated_token_file()
309 .ok_or_else(|| AzureKeyVaultError::MissingField("federated_token_file".to_string()))?;
310 let federated_token = fs::read_to_string(&token_file).await.map_err(|e| {
311 AzureKeyVaultError::HttpError(format!(
312 "failed to read federated token file {token_file}: {e}"
313 ))
314 })?;
315
316 let response = self
317 .client
318 .post(url)
319 .form(&[
320 ("grant_type", "client_credentials"),
321 ("client_id", client_id.as_str()),
322 (
323 "client_assertion_type",
324 "urn:ietf:params:oauth:client-assertion-type:jwt-bearer",
325 ),
326 ("client_assertion", federated_token.trim()),
327 ("scope", AZURE_SCOPE),
328 ])
329 .send()
330 .await
331 .map_err(|e| AzureKeyVaultError::HttpError(e.to_string()))?;
332
333 Self::parse_access_token_response(response).await
334 }
335
336 async fn parse_access_token_response(
338 response: reqwest::Response,
339 ) -> AzureKeyVaultResult<AzureAccessTokenCacheEntry> {
340 let status = response.status();
341 let text = response.text().await.unwrap_or_default();
342
343 if !status.is_success() {
344 return Err(AzureKeyVaultError::ApiError(format!(
345 "token request failed ({status}): {text}"
346 )));
347 }
348
349 let body: Value = serde_json::from_str(&text)
350 .map_err(|e| AzureKeyVaultError::ParseError(format!("{e}: {text}")))?;
351
352 let token = body
353 .get("access_token")
354 .and_then(Value::as_str)
355 .map(ToOwned::to_owned)
356 .ok_or_else(|| AzureKeyVaultError::MissingField("access_token".to_string()))?;
357
358 Ok(AzureAccessTokenCacheEntry {
359 token,
360 expires_at: Instant::now() + Self::parse_access_token_ttl(&body),
361 })
362 }
363
364 fn parse_access_token_ttl(body: &Value) -> Duration {
366 let expires_in = body.get("expires_in").and_then(|value| match value {
367 Value::Number(number) => number.as_u64(),
368 Value::String(text) => text.parse::<u64>().ok(),
369 _ => None,
370 });
371
372 let expires_on = body.get("expires_on").and_then(|value| match value {
373 Value::Number(number) => number.as_u64(),
374 Value::String(text) => text.parse::<u64>().ok(),
375 _ => None,
376 });
377
378 let ttl = expires_in.or_else(|| {
379 expires_on.and_then(|unix_seconds| {
380 let now = std::time::SystemTime::now()
381 .duration_since(std::time::UNIX_EPOCH)
382 .ok()?
383 .as_secs();
384 unix_seconds.checked_sub(now)
385 })
386 });
387
388 let ttl = ttl
389 .map(Duration::from_secs)
390 .unwrap_or(AZURE_TOKEN_CACHE_TTL_FALLBACK);
391 ttl.saturating_sub(AZURE_TOKEN_CACHE_REFRESH_BUFFER)
392 }
393
394 async fn get_access_token(&self) -> AzureKeyVaultResult<String> {
396 let cache_key = self.access_token_cache_key();
397
398 let cached = {
400 let cache_read = AZURE_ACCESS_TOKEN_CACHE.read().await;
401 cache_read.get(&cache_key).cloned()
402 };
403 if let Some(cached) = cached {
404 if Instant::now() < cached.expires_at {
405 return Ok(cached.token);
406 }
407 }
408
409 let entry = match self.auth_type() {
411 AzureKeyVaultAuthType::ClientSecret => self.get_client_secret_access_token().await,
412 AzureKeyVaultAuthType::ManagedIdentity => {
413 self.get_managed_identity_access_token().await
414 }
415 AzureKeyVaultAuthType::WorkloadIdentity => {
416 self.get_workload_identity_access_token().await
417 }
418 }?;
419
420 let token = entry.token.clone();
422 let mut cache_write = AZURE_ACCESS_TOKEN_CACHE.write().await;
423 cache_write.insert(cache_key, entry);
424 Ok(token)
425 }
426
427 async fn key_vault_get(&self, url: &str) -> AzureKeyVaultResult<Value> {
429 let token = self.get_access_token().await?;
430 let response = self
431 .client
432 .get(url)
433 .bearer_auth(token)
434 .send()
435 .await
436 .map_err(|e| AzureKeyVaultError::HttpError(e.to_string()))?;
437
438 let status = response.status();
439 let text = response.text().await.unwrap_or_default();
440
441 if !status.is_success() {
442 return Err(AzureKeyVaultError::ApiError(format!(
443 "key vault request failed ({status}): {text}"
444 )));
445 }
446
447 serde_json::from_str(&text)
448 .map_err(|e| AzureKeyVaultError::ParseError(format!("{e}: {text}")))
449 }
450
451 async fn key_vault_post(&self, url: &str, body: &Value) -> AzureKeyVaultResult<Value> {
453 let token = self.get_access_token().await?;
454 let response = self
455 .client
456 .post(url)
457 .bearer_auth(token)
458 .json(body)
459 .send()
460 .await
461 .map_err(|e| AzureKeyVaultError::HttpError(e.to_string()))?;
462
463 let status = response.status();
464 let text = response.text().await.unwrap_or_default();
465
466 if !status.is_success() {
467 return Err(AzureKeyVaultError::ApiError(format!(
468 "key vault request failed ({status}): {text}"
469 )));
470 }
471
472 serde_json::from_str(&text)
473 .map_err(|e| AzureKeyVaultError::ParseError(format!("{e}: {text}")))
474 }
475
476 async fn get_public_key(&self) -> AzureKeyVaultResult<[u8; 64]> {
478 let cache_key = self.public_key_cache_key();
479 let cached = {
481 let cache_read = AZURE_PUBLIC_KEY_CACHE.read().await;
482 cache_read.get(&cache_key).copied()
483 };
484 if let Some(cached) = cached {
485 return Ok(cached);
486 }
487
488 let body = self.key_vault_get(&self.key_url()).await?;
490 let key = body
491 .get("key")
492 .ok_or_else(|| AzureKeyVaultError::MissingField("key".to_string()))?;
493
494 let x = key
495 .get("x")
496 .and_then(Value::as_str)
497 .ok_or_else(|| AzureKeyVaultError::MissingField("key.x".to_string()))?;
498 let y = key
499 .get("y")
500 .and_then(Value::as_str)
501 .ok_or_else(|| AzureKeyVaultError::MissingField("key.y".to_string()))?;
502
503 let x_bytes = URL_SAFE_NO_PAD
504 .decode(x)
505 .map_err(|e| AzureKeyVaultError::ParseError(e.to_string()))?;
506 let y_bytes = URL_SAFE_NO_PAD
507 .decode(y)
508 .map_err(|e| AzureKeyVaultError::ParseError(e.to_string()))?;
509
510 if x_bytes.len() != 32 || y_bytes.len() != 32 {
511 return Err(AzureKeyVaultError::ParseError(format!(
512 "expected 32-byte secp256k1 coordinates, got x={}, y={}",
513 x_bytes.len(),
514 y_bytes.len()
515 )));
516 }
517
518 let mut public_key = [0u8; 64];
519 public_key[..32].copy_from_slice(&x_bytes);
520 public_key[32..].copy_from_slice(&y_bytes);
521
522 let mut cache_write = AZURE_PUBLIC_KEY_CACHE.write().await;
523 cache_write.insert(cache_key, public_key);
524
525 Ok(public_key)
526 }
527
528 async fn sign_digest(&self, digest: [u8; 32]) -> AzureKeyVaultResult<Vec<u8>> {
530 let body = serde_json::json!({
531 "alg": AZURE_SIGN_ALGORITHM,
532 "value": URL_SAFE_NO_PAD.encode(digest),
533 });
534
535 let response = self.key_vault_post(&self.sign_url(), &body).await?;
536 let signature = response
537 .get("value")
538 .and_then(Value::as_str)
539 .ok_or_else(|| AzureKeyVaultError::MissingField("value".to_string()))?;
540
541 URL_SAFE_NO_PAD
542 .decode(signature)
543 .map_err(|e| AzureKeyVaultError::ParseError(e.to_string()))
544 }
545
546 async fn sign_and_recover_evm(
548 &self,
549 digest: [u8; 32],
550 original_bytes: &[u8],
551 use_prehash_recovery: bool,
552 ) -> AzureKeyVaultResult<Vec<u8>> {
553 let raw_signature = self.sign_digest(digest).await?;
554 if raw_signature.len() != 64 {
555 return Err(AzureKeyVaultError::ParseError(format!(
556 "expected 64-byte ES256K signature, got {} bytes",
557 raw_signature.len()
558 )));
559 }
560
561 let mut rs = Signature::from_slice(&raw_signature)
562 .map_err(|e| AzureKeyVaultError::ParseError(e.to_string()))?;
563
564 if let Some(normalized) = rs.normalize_s() {
565 rs = normalized;
566 }
567
568 let public_key = self.get_public_key().await?;
569 let recovery_id = if use_prehash_recovery {
570 recover_public_key_from_hash(&public_key, &rs, &digest)?
571 } else {
572 recover_public_key(&public_key, &rs, original_bytes)?
573 };
574
575 let mut signature = rs.to_vec();
576 signature.push(27 + recovery_id);
577 Ok(signature)
578 }
579}
580
581#[async_trait]
582impl AzureKeyVaultEvmService for AzureKeyVaultService {
583 async fn get_evm_address(&self) -> AzureKeyVaultResult<Address> {
585 let public_key = self.get_public_key().await?;
586 let hash = keccak256(public_key);
587 let mut address = [0u8; 20];
588 address.copy_from_slice(&hash[12..]);
589 Ok(Address::Evm(address))
590 }
591
592 async fn sign_payload_evm(&self, payload: &[u8]) -> AzureKeyVaultResult<Vec<u8>> {
594 let digest = keccak256(payload).0;
595 self.sign_and_recover_evm(digest, payload, false).await
596 }
597
598 async fn sign_hash_evm(&self, hash: &[u8; 32]) -> AzureKeyVaultResult<Vec<u8>> {
600 self.sign_and_recover_evm(*hash, hash, true).await
601 }
602}
603
604#[cfg(test)]
605mod tests {
606 use super::*;
607 use crate::models::SecretString;
608 use alloy::primitives::utils::eip191_message;
609 use k256::{
610 ecdsa::{signature::hazmat::PrehashSigner, SigningKey},
611 elliptic_curve::rand_core::OsRng,
612 };
613 use mockito::Server;
614 use std::io::Write;
615 use tempfile::NamedTempFile;
616
617 async fn clear_test_caches() {
618 AZURE_ACCESS_TOKEN_CACHE.write().await.clear();
619 AZURE_PUBLIC_KEY_CACHE.write().await.clear();
620 }
621
622 fn test_config(base_url: &str) -> AzureKeyVaultSignerConfig {
623 AzureKeyVaultSignerConfig {
624 auth_type: Some(AzureKeyVaultAuthType::ClientSecret),
625 tenant_id: Some(SecretString::new(base_url)),
626 client_id: Some(SecretString::new("test-client")),
627 client_secret: Some(SecretString::new("test-secret")),
628 federated_token_file: None,
629 vault_url: SecretString::new(base_url),
630 key_name: SecretString::new("test-key"),
631 key_version: Some("test-version".to_string()),
632 }
633 }
634
635 fn managed_identity_config(base_url: &str) -> AzureKeyVaultSignerConfig {
636 AzureKeyVaultSignerConfig {
637 auth_type: Some(AzureKeyVaultAuthType::ManagedIdentity),
638 tenant_id: None,
639 client_id: Some(SecretString::new("managed-client-id")),
640 client_secret: None,
641 federated_token_file: None,
642 vault_url: SecretString::new(base_url),
643 key_name: SecretString::new("test-key"),
644 key_version: Some("test-version".to_string()),
645 }
646 }
647
648 #[tokio::test]
649 async fn test_get_evm_address() {
650 clear_test_caches().await;
651 let mut server = Server::new_async().await;
652 let signing_key = SigningKey::random(&mut OsRng);
653 let point = signing_key.verifying_key().to_encoded_point(false);
654 let x = URL_SAFE_NO_PAD.encode(point.x().unwrap());
655 let y = URL_SAFE_NO_PAD.encode(point.y().unwrap());
656
657 let _token = server
658 .mock("POST", "/oauth2/v2.0/token")
659 .match_body(mockito::Matcher::Any)
660 .with_status(200)
661 .with_header("content-type", "application/json")
662 .with_body(r#"{"access_token":"test-token","expires_in":3600}"#)
663 .expect(1)
664 .create_async()
665 .await;
666
667 let _key = server
668 .mock("GET", "/keys/test-key/test-version")
669 .match_query(mockito::Matcher::UrlEncoded(
670 "api-version".into(),
671 AZURE_API_VERSION.into(),
672 ))
673 .match_header("authorization", "Bearer test-token")
674 .with_status(200)
675 .with_header("content-type", "application/json")
676 .with_body(
677 serde_json::json!({
678 "key": {
679 "x": x,
680 "y": y,
681 }
682 })
683 .to_string(),
684 )
685 .expect(1)
686 .create_async()
687 .await;
688
689 let service = AzureKeyVaultService::new(&test_config(&server.url())).unwrap();
690 let address = service.get_evm_address().await.unwrap();
691 assert!(matches!(address, Address::Evm(_)));
692 }
693
694 #[tokio::test]
695 async fn test_sign_payload_evm() {
696 clear_test_caches().await;
697 let mut server = Server::new_async().await;
698 let signing_key = SigningKey::random(&mut OsRng);
699 let point = signing_key.verifying_key().to_encoded_point(false);
700 let x = URL_SAFE_NO_PAD.encode(point.x().unwrap());
701 let y = URL_SAFE_NO_PAD.encode(point.y().unwrap());
702
703 let message = eip191_message(b"hello azure");
704 let digest = keccak256(&message).0;
705 let signature: Signature = signing_key.sign_prehash(&digest).unwrap();
706 let raw_signature = signature.to_bytes();
707
708 let _token_1 = server
709 .mock("POST", "/oauth2/v2.0/token")
710 .match_body(mockito::Matcher::Any)
711 .with_status(200)
712 .with_header("content-type", "application/json")
713 .with_body(r#"{"access_token":"test-token","expires_in":3600}"#)
714 .expect(1)
715 .create_async()
716 .await;
717
718 let _sign = server
719 .mock("POST", "/keys/test-key/test-version/sign")
720 .match_query(mockito::Matcher::UrlEncoded(
721 "api-version".into(),
722 AZURE_API_VERSION.into(),
723 ))
724 .match_header("authorization", "Bearer test-token")
725 .with_status(200)
726 .with_header("content-type", "application/json")
727 .with_body(
728 serde_json::json!({
729 "value": URL_SAFE_NO_PAD.encode(raw_signature),
730 })
731 .to_string(),
732 )
733 .expect(1)
734 .create_async()
735 .await;
736
737 let _key = server
738 .mock("GET", "/keys/test-key/test-version")
739 .match_query(mockito::Matcher::UrlEncoded(
740 "api-version".into(),
741 AZURE_API_VERSION.into(),
742 ))
743 .match_header("authorization", "Bearer test-token")
744 .with_status(200)
745 .with_header("content-type", "application/json")
746 .with_body(
747 serde_json::json!({
748 "key": {
749 "x": x,
750 "y": y,
751 }
752 })
753 .to_string(),
754 )
755 .expect(1)
756 .create_async()
757 .await;
758
759 let service = AzureKeyVaultService::new(&test_config(&server.url())).unwrap();
760 let signature = service.sign_payload_evm(&message).await.unwrap();
761
762 assert_eq!(signature.len(), 65);
763 assert!(signature[64] == 27 || signature[64] == 28);
764 }
765
766 #[tokio::test]
767 async fn test_managed_identity_access_token() {
768 clear_test_caches().await;
769 let mut server = Server::new_async().await;
770 unsafe {
771 env::set_var(
772 "AZURE_IMDS_TOKEN_URL",
773 format!("{}/metadata/identity/oauth2/token", server.url()),
774 );
775 }
776 let _token = server
777 .mock("GET", "/metadata/identity/oauth2/token")
778 .match_header("metadata", "true")
779 .match_query(mockito::Matcher::AllOf(vec![
780 mockito::Matcher::UrlEncoded("api-version".into(), AZURE_IMDS_API_VERSION.into()),
781 mockito::Matcher::UrlEncoded(
782 "resource".into(),
783 AZURE_MANAGED_IDENTITY_RESOURCE.into(),
784 ),
785 mockito::Matcher::UrlEncoded("client_id".into(), "managed-client-id".into()),
786 ]))
787 .with_status(200)
788 .with_header("content-type", "application/json")
789 .with_body(r#"{"access_token":"managed-token","expires_in":3600}"#)
790 .expect(1)
791 .create_async()
792 .await;
793
794 let service = AzureKeyVaultService {
795 config: managed_identity_config(&server.url()),
796 client: AzureKeyVaultService::build_http_client().unwrap(),
797 };
798
799 let token = service.get_access_token().await.unwrap();
800 assert_eq!(token, "managed-token");
801
802 unsafe {
803 env::remove_var("AZURE_IMDS_TOKEN_URL");
804 }
805 }
806
807 #[tokio::test]
808 async fn test_workload_identity_access_token() {
809 clear_test_caches().await;
810 let mut server = Server::new_async().await;
811 let mut token_file = NamedTempFile::new().unwrap();
812 writeln!(token_file, "federated-jwt").unwrap();
813
814 let _token = server
815 .mock("POST", "/oauth2/v2.0/token")
816 .match_body(mockito::Matcher::AllOf(vec![
817 mockito::Matcher::UrlEncoded("grant_type".into(), "client_credentials".into()),
818 mockito::Matcher::UrlEncoded("client_id".into(), "workload-client-id".into()),
819 mockito::Matcher::UrlEncoded(
820 "client_assertion_type".into(),
821 "urn:ietf:params:oauth:client-assertion-type:jwt-bearer".into(),
822 ),
823 mockito::Matcher::UrlEncoded("client_assertion".into(), "federated-jwt".into()),
824 mockito::Matcher::UrlEncoded("scope".into(), AZURE_SCOPE.into()),
825 ]))
826 .with_status(200)
827 .with_header("content-type", "application/json")
828 .with_body(r#"{"access_token":"workload-token","expires_in":3600}"#)
829 .expect(1)
830 .create_async()
831 .await;
832
833 let config = AzureKeyVaultSignerConfig {
834 auth_type: Some(AzureKeyVaultAuthType::WorkloadIdentity),
835 tenant_id: Some(SecretString::new(&server.url())),
836 client_id: Some(SecretString::new("workload-client-id")),
837 client_secret: None,
838 federated_token_file: Some(SecretString::new(
839 token_file.path().to_string_lossy().as_ref(),
840 )),
841 vault_url: SecretString::new(&server.url()),
842 key_name: SecretString::new("test-key"),
843 key_version: Some("test-version".to_string()),
844 };
845
846 let service = AzureKeyVaultService::new(&config).unwrap();
847 let token = service.get_access_token().await.unwrap();
848 assert_eq!(token, "workload-token");
849 }
850
851 #[tokio::test]
852 async fn test_access_token_is_cached() {
853 clear_test_caches().await;
854 let mut server = Server::new_async().await;
855
856 let _token = server
857 .mock("POST", "/oauth2/v2.0/token")
858 .match_body(mockito::Matcher::Any)
859 .with_status(200)
860 .with_header("content-type", "application/json")
861 .with_body(r#"{"access_token":"cached-token","expires_in":3600}"#)
862 .expect(1)
863 .create_async()
864 .await;
865
866 let service = AzureKeyVaultService::new(&test_config(&server.url())).unwrap();
867
868 let first = service.get_access_token().await.unwrap();
869 let second = service.get_access_token().await.unwrap();
870
871 assert_eq!(first, "cached-token");
872 assert_eq!(second, "cached-token");
873 }
874
875 #[tokio::test]
876 async fn test_public_key_is_cached() {
877 clear_test_caches().await;
878 let mut server = Server::new_async().await;
879 let signing_key = SigningKey::random(&mut OsRng);
880 let point = signing_key.verifying_key().to_encoded_point(false);
881 let x = URL_SAFE_NO_PAD.encode(point.x().unwrap());
882 let y = URL_SAFE_NO_PAD.encode(point.y().unwrap());
883
884 let _token = server
885 .mock("POST", "/oauth2/v2.0/token")
886 .match_body(mockito::Matcher::Any)
887 .with_status(200)
888 .with_header("content-type", "application/json")
889 .with_body(r#"{"access_token":"test-token","expires_in":3600}"#)
890 .expect(1)
891 .create_async()
892 .await;
893
894 let _key = server
895 .mock("GET", "/keys/test-key/test-version")
896 .match_query(mockito::Matcher::UrlEncoded(
897 "api-version".into(),
898 AZURE_API_VERSION.into(),
899 ))
900 .match_header("authorization", "Bearer test-token")
901 .with_status(200)
902 .with_header("content-type", "application/json")
903 .with_body(
904 serde_json::json!({
905 "key": {
906 "x": x,
907 "y": y,
908 }
909 })
910 .to_string(),
911 )
912 .expect(1)
913 .create_async()
914 .await;
915
916 let service = AzureKeyVaultService::new(&test_config(&server.url())).unwrap();
917
918 let first = service.get_public_key().await.unwrap();
919 let second = service.get_public_key().await.unwrap();
920
921 assert_eq!(first, second);
922 }
923}