Skip to main content

ethers_gcp_kms_signer/
lib.rs

1use async_trait::async_trait;
2use ethers::prelude::k256::pkcs8::DecodePublicKey;
3use ethers::types::transaction::eip2718::TypedTransaction;
4use ethers::types::transaction::eip712::Eip712;
5use ethers::{
6    prelude::k256::{
7        ecdsa::{RecoveryId, Signature as KSig, VerifyingKey},
8        FieldBytes,
9    },
10    signers::Signer,
11    types::{Address, Signature, H256, U256},
12    utils::{hash_message, keccak256},
13};
14use gcloud_sdk::{
15    google::cloud::kms::{
16        self,
17        v1::{
18            key_management_service_client::KeyManagementServiceClient, AsymmetricSignRequest,
19            GetPublicKeyRequest,
20        },
21    },
22    GoogleApi, GoogleAuthMiddleware,
23};
24use std::fmt::Debug;
25use tonic::Request;
26use tracing::{debug, instrument};
27
28mod error;
29pub use error::CKMSError;
30
31/// Convert a verifying key to an ethereum address
32fn verifying_key_to_address(key: &VerifyingKey) -> Address {
33    // false for uncompressed
34    let uncompressed_pub_key = key.to_encoded_point(false);
35    let public_key = uncompressed_pub_key.to_bytes();
36    debug_assert_eq!(public_key[0], 0x04);
37    let hash = keccak256(&public_key[1..]);
38    Address::from_slice(&hash[12..])
39}
40
41pub fn apply_eip155(sig: &mut Signature, chain_id: u64) {
42    let v = (chain_id * 2 + 35) + sig.v;
43    sig.v = v;
44}
45
46/// Makes a trial recovery to check whether an RSig corresponds to a known
47/// `VerifyingKey`
48fn check_candidate(
49    sig: &KSig,
50    recovery_id: RecoveryId,
51    digest: [u8; 32],
52    vk: &VerifyingKey,
53) -> bool {
54    VerifyingKey::recover_from_prehash(digest.as_slice(), sig, recovery_id)
55        .map(|key| key == *vk)
56        .unwrap_or(false)
57}
58
59pub fn sig_from_digest_bytes_trial_recovery(
60    sig: &KSig,
61    digest: [u8; 32],
62    vk: &VerifyingKey,
63) -> Signature {
64    let r_bytes: FieldBytes = sig.r().into();
65    let s_bytes: FieldBytes = sig.s().into();
66    let r = U256::from_big_endian(r_bytes.as_slice());
67    let s = U256::from_big_endian(s_bytes.as_slice());
68
69    if check_candidate(sig, RecoveryId::from_byte(0).unwrap(), digest, vk) {
70        Signature { r, s, v: 0 }
71    } else if check_candidate(sig, RecoveryId::from_byte(1).unwrap(), digest, vk) {
72        Signature { r, s, v: 1 }
73    } else {
74        panic!("bad sig");
75    }
76}
77
78#[derive(Clone, Debug)]
79pub struct GcpKeyRingRef {
80    pub google_project_id: String,
81    pub location: String,
82    pub key_ring: String,
83}
84
85impl GcpKeyRingRef {
86    pub fn new(google_project_id: &str, location: &str, key_ring: &str) -> Self {
87        Self {
88            google_project_id: google_project_id.to_string(),
89            location: location.to_string(),
90            key_ring: key_ring.to_string(),
91        }
92    }
93
94    fn to_google_ref(&self) -> String {
95        format!(
96            "projects/{}/locations/{}/keyRings/{}",
97            self.google_project_id, self.location, self.key_ring
98        )
99    }
100
101    fn to_key_version_ref(&self, key_id: &str, key_version: u64) -> String {
102        format!(
103            "{}/cryptoKeys/{}/cryptoKeyVersions/{}",
104            self.to_google_ref(),
105            key_id,
106            key_version,
107        )
108    }
109}
110
111#[derive(Clone)]
112pub struct GcpKmsProvider {
113    client: GoogleApi<KeyManagementServiceClient<GoogleAuthMiddleware>>,
114    kms_key_ref: GcpKeyRingRef,
115}
116
117impl Debug for GcpKmsProvider {
118    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
119        f.debug_struct("GcpKmsProvider")
120            .field("kms_key_ref", &self.kms_key_ref)
121            .finish()
122    }
123}
124
125impl GcpKmsProvider {
126    pub async fn new(kms_key_ref: GcpKeyRingRef) -> Result<Self, CKMSError> {
127        debug!(
128            "Initialising Google KMS envelope encryption for {}",
129            kms_key_ref.to_google_ref()
130        );
131
132        let client = GoogleApi::from_function(
133            KeyManagementServiceClient::new,
134            "https://cloudkms.googleapis.com",
135            None,
136        )
137        .await?;
138
139        Ok(Self {
140            kms_key_ref,
141            client,
142        })
143    }
144
145    pub async fn get_verifying_key(
146        &self,
147        key_id: &str,
148        key_version: u64,
149    ) -> Result<VerifyingKey, CKMSError> {
150        let kms_key_name = self.kms_key_ref.to_key_version_ref(key_id, key_version);
151
152        let mut request = tonic::Request::new(GetPublicKeyRequest {
153            name: kms_key_name.clone(),
154        });
155
156        // Add metadata for request routing: https://cloud.google.com/kms/docs/grpc
157        request.metadata_mut().insert(
158            "x-goog-request-params",
159            format!("name={}", kms_key_name.clone()).parse().unwrap(),
160        );
161
162        let response = self.client.get().get_public_key(request).await?;
163        let pem = response.into_inner().pem;
164        let public_key = VerifyingKey::from_public_key_pem(&pem)?;
165        Ok(public_key)
166    }
167
168    pub async fn sign_digest(
169        &self,
170        key_id: &str,
171        key_version: u64,
172        digest: &[u8],
173    ) -> Result<Vec<u8>, CKMSError> {
174        let kms_key_name = self.kms_key_ref.to_key_version_ref(key_id, key_version);
175
176        let mut request = Request::new(AsymmetricSignRequest {
177            name: kms_key_name.clone(),
178            digest: Some(kms::v1::Digest {
179                digest: Some(kms::v1::digest::Digest::Sha256(digest.to_vec())),
180            }),
181            ..Default::default()
182        });
183
184        // Add metadata for request routing: https://cloud.google.com/kms/docs/grpc
185        request.metadata_mut().insert(
186            "x-goog-request-params",
187            format!("name={}", kms_key_name.clone()).parse().unwrap(),
188        );
189
190        let response = self.client.get().asymmetric_sign(request).await?;
191        let signature = response.into_inner().signature;
192        Ok(signature)
193    }
194}
195
196#[derive(Clone, Debug)]
197pub struct GcpKmsSigner {
198    provider: GcpKmsProvider,
199    key_id: String,
200    key_version: u64,
201    chain_id: u64,
202    verifying_key: VerifyingKey,
203}
204
205impl GcpKmsSigner {
206    pub async fn new(
207        provider: GcpKmsProvider,
208        key_id: String,
209        key_version: u64,
210        chain_id: u64,
211    ) -> Result<Self, CKMSError> {
212        let verifying_key = provider.get_verifying_key(&key_id, key_version).await?;
213        Ok(Self {
214            provider,
215            key_id,
216            key_version,
217            chain_id,
218            verifying_key,
219        })
220    }
221
222    /// Sign a digest with this signer's key
223    pub async fn sign_digest(&self, digest: [u8; 32]) -> Result<KSig, CKMSError> {
224        let signature = self
225            .provider
226            .sign_digest(self.key_id.as_ref(), self.key_version, digest.as_ref())
227            .await?;
228        let sig = KSig::from_der(&signature)?;
229        let sig = sig.normalize_s().unwrap_or(sig);
230        Ok(sig)
231    }
232
233    /// Sign a digest with this signer's key and add the eip155 `v` value
234    /// corresponding to the input chain_id
235    #[instrument(err, skip(digest))]
236    async fn sign_digest_with_eip155(
237        &self,
238        digest: H256,
239        chain_id: u64,
240    ) -> Result<Signature, CKMSError> {
241        let sig = self.sign_digest(digest.into()).await?;
242        let mut sig =
243            sig_from_digest_bytes_trial_recovery(&sig, digest.into(), &self.verifying_key);
244        apply_eip155(&mut sig, chain_id);
245        Ok(sig)
246    }
247}
248
249#[async_trait]
250impl Signer for GcpKmsSigner {
251    type Error = CKMSError;
252
253    /// Signs the message
254    #[instrument(err, skip(message))]
255    async fn sign_message<S: Send + Sync + AsRef<[u8]>>(
256        &self,
257        message: S,
258    ) -> Result<Signature, Self::Error> {
259        let message = message.as_ref();
260        let message_hash = hash_message(message);
261        self.sign_digest_with_eip155(message_hash, self.chain_id)
262            .await
263    }
264
265    /// Signs the transaction
266    #[instrument(err)]
267    async fn sign_transaction(&self, tx: &TypedTransaction) -> Result<Signature, Self::Error> {
268        let mut tx_with_chain = tx.clone();
269        let chain_id = tx_with_chain
270            .chain_id()
271            .map(|id| id.as_u64())
272            .unwrap_or(self.chain_id);
273        tx_with_chain.set_chain_id(chain_id);
274
275        let sighash = tx_with_chain.sighash();
276        self.sign_digest_with_eip155(sighash, chain_id).await
277    }
278
279    /// Encodes and signs the typed data according EIP-712.
280    /// Payload must implement Eip712 trait.
281    async fn sign_typed_data<T: Eip712 + Send + Sync>(
282        &self,
283        payload: &T,
284    ) -> Result<Signature, Self::Error> {
285        let digest = payload
286            .encode_eip712()
287            .map_err(|e| CKMSError::Eip712Error(e.to_string()))?;
288
289        let sig = self.sign_digest(digest).await?;
290        let sig = sig_from_digest_bytes_trial_recovery(&sig, digest, &self.verifying_key);
291
292        Ok(sig)
293    }
294
295    /// Returns the signer's Ethereum Address
296    fn address(&self) -> Address {
297        verifying_key_to_address(&self.verifying_key)
298    }
299
300    /// Returns the signer's chain id
301    fn chain_id(&self) -> u64 {
302        self.chain_id
303    }
304
305    /// Sets the signer's chain id
306    #[must_use]
307    fn with_chain_id<T: Into<u64>>(self, chain_id: T) -> Self {
308        let mut this = self;
309        this.chain_id = chain_id.into();
310        this
311    }
312}
313
314#[cfg(test)]
315mod tests {
316    use super::*;
317
318    #[test_log::test(tokio::test)]
319    async fn it_works() {
320        // skip test if no credentials are provided
321        if std::env::var("GOOGLE_APPLICATION_CREDENTIALS").is_err() {
322            return;
323        }
324
325        let project_id = std::env::var("GOOGLE_PROJECT_ID").expect("GOOGLE_PROJECT_ID");
326        let location = std::env::var("GOOGLE_LOCATION").expect("GOOGLE_LOCATION");
327        let keyring = std::env::var("GOOGLE_KEYRING").expect("GOOGLE_KEYRING");
328        let key_name = std::env::var("GOOGLE_KEY_NAME").expect("GOOGLE_KEY_NAME");
329
330        let keyring = GcpKeyRingRef::new(&project_id, &location, &keyring);
331        let provider = GcpKmsProvider::new(keyring)
332            .await
333            .expect("Failed to create GCP KMS provider");
334        let signer = GcpKmsSigner::new(provider, key_name, 1, 1)
335            .await
336            .expect("get key");
337
338        let message = vec![0, 1, 2, 3];
339        let sig = signer.sign_message(&message).await.unwrap();
340        sig.verify(message, signer.address()).expect("valid sig");
341    }
342}