Skip to main content

keymaster_multisig/
multisig.rs

1use crate::error::{MultisigError, Result};
2use crate::types::{encode_varint, PrivateKey, PublicKey, Transaction};
3use k256::{
4    ecdsa::{
5        signature::{hazmat::PrehashSigner, SignatureEncoding},
6        Signature as EcdsaSignature, SigningKey,
7    },
8    SecretKey,
9};
10use serde::{Deserialize, Serialize};
11use sha2::{Digest, Sha256};
12
13const OP_0: u8 = 0x00;
14const OP_CHECKMULTISIG: u8 = 0xae;
15const SIGHASH_ALL_FORKID: u8 = 0x41;
16
17#[derive(Serialize, Deserialize, Debug)]
18pub struct Multisig {
19    private_keys: Option<Vec<PrivateKey>>,
20    public_keys: Vec<PublicKey>,
21    m: usize,
22    n: usize,
23    sig_hash_type: u8,
24}
25
26impl Multisig {
27    pub fn new(
28        private_keys: Option<Vec<PrivateKey>>,
29        public_keys: Vec<PublicKey>,
30        m: usize,
31    ) -> Result<Self> {
32        if public_keys.is_empty() || public_keys.len() > 20 {
33            return Err(MultisigError::InvalidPublicKeys);
34        }
35
36        if m == 0 || m > public_keys.len() {
37            return Err(MultisigError::InvalidM(format!(
38                "m={} must be between 1 and n={}",
39                m,
40                public_keys.len()
41            )));
42        }
43
44        if let Some(ref keys) = private_keys {
45            if keys.len() < m {
46                return Err(MultisigError::NoPrivateKeys);
47            }
48        }
49
50        let n = public_keys.len();
51        Ok(Multisig {
52            private_keys,
53            public_keys,
54            m,
55            n,
56            sig_hash_type: SIGHASH_ALL_FORKID,
57        })
58    }
59
60    pub fn lock(&self) -> Result<Vec<u8>> {
61        if self.m == 0 || self.m > self.n {
62            return Err(MultisigError::InvalidM(format!(
63                "m={} must be between 1 and n={}",
64                self.m, self.n
65            )));
66        }
67        if self.n == 0 || self.n > 20 {
68            return Err(MultisigError::InvalidPublicKeys);
69        }
70
71        let mut script = Vec::new();
72
73        script.push(0x01 + (self.m as u8) - 1);
74
75        for pub_key in &self.public_keys {
76            script.push(pub_key.key.len() as u8);
77            script.extend(&pub_key.key);
78        }
79
80        script.push(0x01 + (self.n as u8) - 1);
81        script.push(OP_CHECKMULTISIG);
82
83        Ok(script)
84    }
85
86    pub fn sign(&self, tx: &Transaction, input_index: usize) -> Result<Vec<Vec<u8>>> {
87        if let Some(ref priv_keys) = self.private_keys {
88            if priv_keys.len() < self.m {
89                return Err(MultisigError::NoPrivateKeys);
90            }
91
92            let mut signatures = Vec::new();
93
94            for private_key in priv_keys.iter().take(self.m) {
95                let sig = self.sign_one(tx, input_index, private_key)?;
96                signatures.push(sig);
97            }
98
99            Ok(signatures)
100        } else {
101            Err(MultisigError::NoPrivateKeys)
102        }
103    }
104
105    pub fn sign_one(
106        &self,
107        tx: &Transaction,
108        input_index: usize,
109        private_key: &PrivateKey,
110    ) -> Result<Vec<u8>> {
111        if input_index >= tx.inputs.len() {
112            return Err(MultisigError::TransactionError(
113                "Input index out of bounds".to_string(),
114            ));
115        }
116
117        let sighash = self.calculate_signature_hash(tx, input_index)?;
118
119        let signature = self.generate_signature(&sighash, private_key)?;
120
121        Ok(signature)
122    }
123
124    fn calculate_signature_hash(&self, tx: &Transaction, input_index: usize) -> Result<Vec<u8>> {
125        if input_index >= tx.inputs.len() {
126            return Err(MultisigError::TransactionError(
127                "Input index out of bounds".to_string(),
128            ));
129        }
130        let source = tx.inputs[input_index]
131            .source_output
132            .as_ref()
133            .ok_or_else(|| {
134                MultisigError::TransactionError("Source output is required".to_string())
135            })?;
136        let hash256 = |value: &[u8]| -> [u8; 32] {
137            let first = Sha256::digest(value);
138            Sha256::digest(first).into()
139        };
140        let mut prevouts = Vec::new();
141        let mut sequences = Vec::new();
142        for input in &tx.inputs {
143            let mut txid = hex::decode(&input.source_txid)
144                .map_err(|_| MultisigError::TransactionError("Invalid source txid".to_string()))?;
145            if txid.len() != 32 {
146                return Err(MultisigError::TransactionError(
147                    "Invalid source txid length".to_string(),
148                ));
149            }
150            txid.reverse();
151            prevouts.extend(txid);
152            prevouts.extend_from_slice(&input.source_output_index.to_le_bytes());
153            sequences.extend_from_slice(&input.sequence.to_le_bytes());
154        }
155        let mut outputs = Vec::new();
156        for output in &tx.outputs {
157            outputs.extend_from_slice(&output.satoshis.to_le_bytes());
158            outputs.extend(encode_varint(output.locking_script.len() as u64));
159            outputs.extend(&output.locking_script);
160        }
161        let input = &tx.inputs[input_index];
162        let mut outpoint_txid = hex::decode(&input.source_txid)
163            .map_err(|_| MultisigError::TransactionError("Invalid source txid".to_string()))?;
164        outpoint_txid.reverse();
165        let mut preimage = Vec::new();
166        preimage.extend_from_slice(&tx.version.to_le_bytes());
167        preimage.extend(hash256(&prevouts));
168        preimage.extend(hash256(&sequences));
169        preimage.extend(outpoint_txid);
170        preimage.extend_from_slice(&input.source_output_index.to_le_bytes());
171        preimage.extend(encode_varint(source.locking_script.len() as u64));
172        preimage.extend(&source.locking_script);
173        preimage.extend_from_slice(&source.satoshis.to_le_bytes());
174        preimage.extend_from_slice(&input.sequence.to_le_bytes());
175        preimage.extend(hash256(&outputs));
176        preimage.extend_from_slice(&tx.lock_time.to_le_bytes());
177        preimage.extend_from_slice(&(self.sig_hash_type as u32).to_le_bytes());
178        Ok(hash256(&preimage).to_vec())
179    }
180
181    fn generate_signature(&self, sighash: &[u8], private_key: &PrivateKey) -> Result<Vec<u8>> {
182        // Convert private key bytes to SecretKey
183        let secret_key = SecretKey::from_slice(&private_key.key)
184            .map_err(|_| MultisigError::InvalidPrivateKey)?;
185
186        let signing_key = SigningKey::from(secret_key);
187        let mut signature: EcdsaSignature = signing_key
188            .sign_prehash(sighash)
189            .map_err(|_| MultisigError::SignatureError("Failed to create signature".to_string()))?;
190        if let Some(normalized) = signature.normalize_s() {
191            signature = normalized;
192        }
193
194        // Convert to DER format and add SIGHASH type
195        let der_sig = signature.to_der();
196        let mut sig_with_hash = der_sig.to_vec();
197        sig_with_hash.push(self.sig_hash_type);
198
199        Ok(sig_with_hash)
200    }
201
202    pub fn estimate_length(&self) -> usize {
203        1 + self.m * (71 + 1)
204    }
205
206    pub fn create_fake_sign(&self) -> Result<Vec<u8>> {
207        let mut script = vec![OP_0];
208
209        for _ in 0..self.m {
210            script.extend(vec![0u8; 72]);
211            script.push(self.sig_hash_type);
212        }
213
214        Ok(script)
215    }
216
217    pub fn build_sign_script(&self, signatures: &[Vec<u8>]) -> Result<Vec<u8>> {
218        let mut script = vec![OP_0];
219
220        for sig in signatures {
221            script.push(sig.len() as u8);
222            script.extend(sig);
223        }
224
225        Ok(script)
226    }
227
228    pub fn get_m(&self) -> usize {
229        self.m
230    }
231
232    pub fn get_n(&self) -> usize {
233        self.n
234    }
235
236    pub fn get_sig_hash_type(&self) -> u8 {
237        self.sig_hash_type
238    }
239
240    pub fn get_public_keys(&self) -> &[PublicKey] {
241        &self.public_keys
242    }
243}
244
245#[cfg(test)]
246mod tests {
247    use super::Multisig;
248    use crate::types::{PrivateKey, PublicKey, Transaction, TransactionInput, TransactionOutput};
249
250    #[test]
251    fn supports_all_two_of_three_signature_pairs() {
252        let public_keys = vec![
253            PublicKey::new(vec![0x02; 33]),
254            PublicKey::new(vec![0x03; 33]),
255            PublicKey::new(vec![0x04; 33]),
256        ];
257        let transaction = Transaction::new(
258            1,
259            vec![TransactionInput {
260                source_txid: "aa".repeat(32),
261                source_output_index: 0,
262                unlocking_script: Vec::new(),
263                sequence: 1,
264                source_output: Some(TransactionOutput::new(1000, vec![0x51])),
265            }],
266            vec![TransactionOutput::new(1000, vec![0x51])],
267            0,
268        );
269        let signer = Multisig::new(None, public_keys, 2).unwrap();
270        let buyer = signer
271            .sign_one(&transaction, 0, &PrivateKey::new(vec![1; 32]))
272            .unwrap();
273        let seller = signer
274            .sign_one(&transaction, 0, &PrivateKey::new(vec![2; 32]))
275            .unwrap();
276        let arbiter = signer
277            .sign_one(&transaction, 0, &PrivateKey::new(vec![3; 32]))
278            .unwrap();
279
280        let buyer_seller = signer
281            .build_sign_script(&[buyer.clone(), seller.clone()])
282            .unwrap();
283        let buyer_arbiter = signer
284            .build_sign_script(&[buyer.clone(), arbiter.clone()])
285            .unwrap();
286        let seller_arbiter = signer.build_sign_script(&[seller, arbiter]).unwrap();
287
288        assert_eq!(buyer_seller[0], 0);
289        assert_eq!(buyer_arbiter[0], 0);
290        assert_eq!(seller_arbiter[0], 0);
291        assert_ne!(buyer_seller, buyer_arbiter);
292        assert_ne!(buyer_arbiter, seller_arbiter);
293    }
294}