openzeppelin_relayer/services/azure_key_vault/
mod.rs

1//! Azure Key Vault service for EVM secp256k1 signing.
2
3use 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    /// Returns the EVM address derived from the configured Azure Key Vault key.
58    async fn get_evm_address(&self) -> AzureKeyVaultResult<Address>;
59    /// Signs a payload using the EVM signing scheme (hashes before signing).
60    async fn sign_payload_evm(&self, payload: &[u8]) -> AzureKeyVaultResult<Vec<u8>>;
61    /// Signs a pre-computed hash using the EVM signing scheme (no hashing).
62    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
92// Global cache for Azure access tokens - HashMap keyed by auth configuration
93static AZURE_ACCESS_TOKEN_CACHE: Lazy<
94    RwLock<HashMap<AzureAccessTokenCacheKey, AzureAccessTokenCacheEntry>>,
95> = Lazy::new(|| RwLock::new(HashMap::new()));
96
97// Global cache for secp256k1 public keys - HashMap keyed by vault/key path
98static AZURE_PUBLIC_KEY_CACHE: Lazy<RwLock<HashMap<AzurePublicKeyCacheKey, [u8; 64]>>> =
99    Lazy::new(|| RwLock::new(HashMap::new()));
100
101impl AzureKeyVaultService {
102    /// Creates a new Azure Key Vault service with a shared HTTP client configuration.
103    pub fn new(config: &AzureKeyVaultSignerConfig) -> AzureKeyVaultResult<Self> {
104        Ok(Self {
105            config: config.clone(),
106            client: Self::build_http_client()?,
107        })
108    }
109
110    /// Builds the reqwest client used for Azure AD and Key Vault requests.
111    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    /// Returns the configured tenant identifier, or an empty string if unset.
121    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    /// Returns the configured client identifier, or an empty string if unset.
130    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    /// Returns the configured client secret, or an empty string if unset.
139    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    /// Resolves the workload identity federated token file from config or environment.
148    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    /// Returns the configured vault base URL without a trailing slash.
157    fn vault_url(&self) -> String {
158        self.config
159            .vault_url
160            .to_str()
161            .trim_end_matches('/')
162            .to_string()
163    }
164
165    /// Returns the configured key name.
166    fn key_name(&self) -> String {
167        self.config.key_name.to_str().to_string()
168    }
169
170    /// Builds the Azure AD OAuth token endpoint for the configured tenant.
171    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    /// Builds the Key Vault key path, including the key version when configured.
184    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    /// Builds the Key Vault URL used to fetch the public key material.
192    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    /// Builds the Key Vault URL used to request signatures.
202    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    /// Returns the configured Azure authentication type.
212    fn auth_type(&self) -> AzureKeyVaultAuthType {
213        self.config.auth_type()
214    }
215
216    /// Returns a stable string representation of the configured authentication type.
217    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    /// Returns the IMDS token endpoint, allowing tests to override it via environment.
227    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    /// Builds the cache key for Azure access token reuse.
232    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    /// Builds the cache key for the Azure secp256k1 public key.
243    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    /// Fetches an Azure AD access token using the client secret flow.
251    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    /// Fetches an Azure AD access token using the managed identity flow.
275    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    /// Fetches an Azure AD access token using the workload identity flow.
302    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    /// Parses a token response and converts it into a cache entry with an expiry timestamp.
337    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    /// Derives a token TTL from Azure response fields and applies a refresh buffer.
365    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    /// Returns a cached Azure access token or fetches and caches a fresh one.
395    async fn get_access_token(&self) -> AzureKeyVaultResult<String> {
396        let cache_key = self.access_token_cache_key();
397
398        // Try cache first with minimal lock time
399        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        // Fetch a fresh token from Azure AD or IMDS
410        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        // Update the cache
421        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    /// Sends an authenticated GET request to Azure Key Vault and parses the JSON response.
428    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    /// Sends an authenticated POST request to Azure Key Vault and parses the JSON response.
452    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    /// Returns the uncompressed secp256k1 public key, using the cache when available.
477    async fn get_public_key(&self) -> AzureKeyVaultResult<[u8; 64]> {
478        let cache_key = self.public_key_cache_key();
479        // Try cache first with minimal lock time
480        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        // Fetch from Azure Key Vault
489        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    /// Requests a raw ES256K signature for the provided 32-byte digest.
529    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    /// Signs a digest and converts the Azure response into a recoverable EVM signature.
547    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    /// Returns the EVM address derived from the configured Azure Key Vault key.
584    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    /// Signs a payload using the EVM signing scheme (hashes before signing).
593    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    /// Signs a pre-computed hash using the EVM signing scheme (no hashing).
599    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}