use k256::ecdsa::{Signature as EcdsaSig, SigningKey};
use crate::{recover, signer_address, Address, Sig, Word, U256};
pub const SECP256K1_N: Word = [
0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFE,
0xBA, 0xAE, 0xDC, 0xE6, 0xAF, 0x48, 0xA0, 0x3B, 0xBF, 0xD2, 0x5E, 0x8C, 0xD0, 0x36, 0x41, 0x41,
];
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum SignerError {
#[error("signing backend failed to produce a signature")]
SigningFailed,
#[error("no recovery id recovers the expected signer address")]
RecoveryMismatch,
#[error("malformed DER-encoded ECDSA signature")]
BadDer,
#[error("invalid secp256k1 scalar in signature")]
BadScalar,
}
pub trait Signer {
fn address(&self) -> Address;
fn sign_digest(&self, digest: &Word) -> Result<Sig, SignerError>;
}
#[derive(Clone)]
pub struct LocalSigner {
sk: SigningKey,
address: Address,
}
impl LocalSigner {
pub fn from_bytes(bytes: &[u8; 32]) -> Result<Self, SignerError> {
let fb: k256::FieldBytes = (*bytes).into();
let sk = SigningKey::from_bytes(&fb).map_err(|_| SignerError::BadScalar)?;
let address = signer_address(&sk);
Ok(Self { sk, address })
}
}
impl core::fmt::Debug for LocalSigner {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("LocalSigner")
.field("address", &self.address)
.finish_non_exhaustive()
}
}
impl Signer for LocalSigner {
fn address(&self) -> Address {
self.address
}
fn sign_digest(&self, digest: &Word) -> Result<Sig, SignerError> {
Ok(crate::sign_digest(&self.sk, digest))
}
}
pub fn sig_from_rs(
digest: &Word,
r: &Word,
s: &Word,
expected_signer: &Address,
) -> Result<Sig, SignerError> {
let mut rs = [0u8; 64];
rs[..32].copy_from_slice(r);
rs[32..].copy_from_slice(s);
let parsed = EcdsaSig::from_slice(&rs).map_err(|_| SignerError::BadScalar)?;
let low = parsed.normalize_s().unwrap_or(parsed);
let bytes = low.to_bytes();
let mut lr = [0u8; 32];
let mut ls = [0u8; 32];
lr.copy_from_slice(&bytes[..32]);
ls.copy_from_slice(&bytes[32..]);
for v in [27u8, 28u8] {
let candidate = Sig { r: lr, s: ls, v };
if recover(digest, &candidate).as_ref() == Some(expected_signer) {
return Ok(candidate);
}
}
Err(SignerError::RecoveryMismatch)
}
pub fn sig_from_der(
digest: &Word,
der: &[u8],
expected_signer: &Address,
) -> Result<Sig, SignerError> {
let parsed = EcdsaSig::from_der(der).map_err(|_| SignerError::BadDer)?;
let bytes = parsed.to_bytes();
let mut r = [0u8; 32];
let mut s = [0u8; 32];
r.copy_from_slice(&bytes[..32]);
s.copy_from_slice(&bytes[32..]);
sig_from_rs(digest, &r, &s, expected_signer)
}
pub struct ExternalSigner<F> {
address: Address,
sign_der: F,
}
impl<F> ExternalSigner<F>
where
F: Fn(&Word) -> Result<Vec<u8>, SignerError>,
{
pub fn new(address: Address, sign_der: F) -> Self {
Self { address, sign_der }
}
}
impl<F> core::fmt::Debug for ExternalSigner<F> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("ExternalSigner")
.field("address", &self.address)
.finish_non_exhaustive()
}
}
impl<F> Signer for ExternalSigner<F>
where
F: Fn(&Word) -> Result<Vec<u8>, SignerError>,
{
fn address(&self) -> Address {
self.address
}
fn sign_digest(&self, digest: &Word) -> Result<Sig, SignerError> {
let der = (self.sign_der)(digest)?;
sig_from_der(digest, &der, &self.address)
}
}
pub fn is_high_s(s: &Word) -> bool {
let half_n = U256::from_be_bytes(SECP256K1_N) >> 1;
U256::from_be_bytes(*s) > half_n
}
pub fn high_s_counterpart(s: &Word) -> Word {
let n = U256::from_be_bytes(SECP256K1_N);
let s_val = U256::from_be_bytes(*s);
(n - s_val).to_be_bytes::<32>()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn local_signer_round_trips() {
let signer = LocalSigner::from_bytes(&[0x42; 32]).unwrap();
let digest = crate::keccak(b"local signer unit test digest...");
let sig = signer.sign_digest(&digest).unwrap();
assert_eq!(recover(&digest, &sig).as_ref(), Some(&signer.address()));
}
#[test]
fn is_high_s_boundary() {
let half_n: U256 = U256::from_be_bytes(SECP256K1_N) >> 1;
assert!(!is_high_s(&half_n.to_be_bytes::<32>()));
assert!(is_high_s(&(half_n + U256::from(1u64)).to_be_bytes::<32>()));
}
#[test]
fn high_s_counterpart_is_high() {
let one = {
let mut w = [0u8; 32];
w[31] = 1;
w
};
let back = high_s_counterpart(&high_s_counterpart(&one));
assert_eq!(back, one);
}
}