use alloy_primitives::{Address, keccak256};
use om_primitives_types::transaction::{B264, MultiSigSignatureEntry, Signature};
use thiserror::Error;
pub const DST_MULTISIG_ADDR_V1: &[u8] = b"MULTISIG_V1";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SignerConfig {
pub public_key: Vec<u8>,
pub weight: u8,
}
impl SignerConfig {
pub fn new(public_key: Vec<u8>, weight: u8) -> Result<Self, MultiSigError> {
if public_key.len() != 33 {
return Err(MultiSigError::InvalidPublicKeyLength {
expected: 33,
actual: public_key.len(),
});
}
if weight == 0 {
return Err(MultiSigError::InvalidWeight);
}
Ok(Self { public_key, weight })
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ThresholdConfig {
pub threshold: u16,
}
impl ThresholdConfig {
pub fn new(threshold: u16) -> Result<Self, MultiSigError> {
if threshold == 0 {
return Err(MultiSigError::InvalidThreshold);
}
Ok(Self { threshold })
}
}
#[derive(Debug, Error)]
pub enum MultiSigError {
#[error("invalid public key length: expected {expected} bytes, got {actual} bytes")]
InvalidPublicKeyLength { expected: usize, actual: usize },
#[error("invalid weight: must be greater than 0")]
InvalidWeight,
#[error("invalid threshold: must be greater than 0")]
InvalidThreshold,
#[error("no signers provided")]
NoSigners,
#[error("threshold {threshold} exceeds total weight {total_weight}")]
ThresholdExceedsTotalWeight { threshold: u16, total_weight: u16 },
#[error("duplicate public key found")]
DuplicatePublicKey,
#[error("signature count mismatch: expected at least 1, got {count}")]
InsufficientSignatures { count: usize },
}
pub fn derive_multisig_address(
signers: &[SignerConfig],
threshold: &ThresholdConfig,
) -> Result<Address, MultiSigError> {
if signers.is_empty() {
return Err(MultiSigError::NoSigners);
}
if threshold.threshold == 0 {
return Err(MultiSigError::InvalidThreshold);
}
for signer in signers {
if signer.weight == 0 {
return Err(MultiSigError::InvalidWeight);
}
if signer.public_key.len() != 33 {
return Err(MultiSigError::InvalidPublicKeyLength {
expected: 33,
actual: signer.public_key.len(),
});
}
}
let total_weight: u16 = signers.iter().map(|s| s.weight as u16).sum();
if threshold.threshold > total_weight {
return Err(MultiSigError::ThresholdExceedsTotalWeight {
threshold: threshold.threshold,
total_weight,
});
}
let mut sorted_signers: Vec<&SignerConfig> = signers.iter().collect();
sorted_signers.sort_by(|a, b| a.public_key.cmp(&b.public_key));
for i in 1..sorted_signers.len() {
if sorted_signers[i].public_key == sorted_signers[i - 1].public_key {
return Err(MultiSigError::DuplicatePublicKey);
}
}
let mut data = Vec::with_capacity(DST_MULTISIG_ADDR_V1.len() + signers.len() * 34 + 2);
data.extend_from_slice(DST_MULTISIG_ADDR_V1);
for signer in &sorted_signers {
data.extend_from_slice(&signer.public_key);
data.push(signer.weight);
}
data.extend_from_slice(&threshold.threshold.to_be_bytes());
let hash = keccak256(&data);
let mut address_bytes = [0u8; 20];
address_bytes.copy_from_slice(&hash[12..32]);
Ok(Address::from(address_bytes))
}
pub struct MultiSigSignatureCollector {
signatures: Vec<MultiSigSignatureEntry>,
}
impl MultiSigSignatureCollector {
pub fn new() -> Self {
Self { signatures: Vec::new() }
}
pub fn add_signature(&mut self, signer_pubkey: B264, signature: Signature) -> &mut Self {
self.signatures.push(MultiSigSignatureEntry {
signer_pubkey,
signature,
});
self
}
pub fn add_signatures(&mut self, mut signatures: Vec<MultiSigSignatureEntry>) -> &mut Self {
self.signatures.append(&mut signatures);
self
}
pub fn signature_count(&self) -> usize {
self.signatures.len()
}
pub fn signatures(self) -> Vec<MultiSigSignatureEntry> {
self.signatures
}
pub fn is_empty(&self) -> bool {
self.signatures.is_empty()
}
}
impl Default for MultiSigSignatureCollector {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_signer_config_new_valid() {
let pubkey = vec![2; 33];
let signer = SignerConfig::new(pubkey.clone(), 1).unwrap();
assert_eq!(signer.public_key, pubkey);
assert_eq!(signer.weight, 1);
}
#[test]
fn test_signer_config_invalid_pubkey_length() {
let pubkey = vec![2; 32]; let result = SignerConfig::new(pubkey, 1);
assert!(matches!(result, Err(MultiSigError::InvalidPublicKeyLength { .. })));
}
#[test]
fn test_signer_config_zero_weight() {
let pubkey = vec![2; 33];
let result = SignerConfig::new(pubkey, 0);
assert!(matches!(result, Err(MultiSigError::InvalidWeight)));
}
#[test]
fn test_threshold_config_new_valid() {
let threshold = ThresholdConfig::new(2).unwrap();
assert_eq!(threshold.threshold, 2);
}
#[test]
fn test_threshold_config_zero() {
let result = ThresholdConfig::new(0);
assert!(matches!(result, Err(MultiSigError::InvalidThreshold)));
}
#[test]
fn test_derive_multisig_address_2_of_3() {
let signer1 = SignerConfig::new(vec![2; 33], 1).unwrap();
let signer2 = SignerConfig::new(vec![3; 33], 1).unwrap();
let signer3 = SignerConfig::new(vec![4; 33], 1).unwrap();
let threshold = ThresholdConfig::new(2).unwrap();
let address =
derive_multisig_address(&[signer1.clone(), signer2.clone(), signer3.clone()], &threshold).unwrap();
let address2 = derive_multisig_address(&[signer3, signer1, signer2], &threshold).unwrap();
assert_eq!(
address, address2,
"Address should be deterministic regardless of input order"
);
}
#[test]
fn test_derive_multisig_address_no_signers() {
let threshold = ThresholdConfig::new(1).unwrap();
let result = derive_multisig_address(&[], &threshold);
assert!(matches!(result, Err(MultiSigError::NoSigners)));
}
#[test]
fn test_derive_multisig_address_threshold_exceeds_weight() {
let signer1 = SignerConfig::new(vec![2; 33], 1).unwrap();
let signer2 = SignerConfig::new(vec![3; 33], 1).unwrap();
let threshold = ThresholdConfig::new(3).unwrap();
let result = derive_multisig_address(&[signer1, signer2], &threshold);
assert!(matches!(result, Err(MultiSigError::ThresholdExceedsTotalWeight { .. })));
}
#[test]
fn test_derive_multisig_address_duplicate_pubkey() {
let signer1 = SignerConfig::new(vec![2; 33], 1).unwrap();
let signer2 = SignerConfig::new(vec![2; 33], 2).unwrap(); let threshold = ThresholdConfig::new(2).unwrap();
let result = derive_multisig_address(&[signer1, signer2], &threshold);
assert!(matches!(result, Err(MultiSigError::DuplicatePublicKey)));
}
#[test]
fn test_multisig_collector() {
let mut collector = MultiSigSignatureCollector::new();
assert_eq!(collector.signature_count(), 0);
assert!(collector.is_empty());
let sig1 = Signature::test_signature();
let sig2 = Signature::test_signature();
collector.add_signature(B264::repeat_byte(2), sig1);
collector.add_signature(B264::repeat_byte(3), sig2);
assert_eq!(collector.signature_count(), 2);
assert!(!collector.is_empty());
let signatures = collector.signatures();
assert_eq!(signatures.len(), 2);
}
#[test]
fn test_multisig_collector_default() {
let collector = MultiSigSignatureCollector::default();
assert!(collector.is_empty());
assert_eq!(collector.signature_count(), 0);
}
}