1use crate::error::{CryptoError, Result};
9use crate::internal::secp256k1::{AffinePoint, EncodedPoint, ProjectivePoint, Scalar};
10use crate::internal::subtle::{ct_option_to_option, ConstantTimeEq};
11use crate::primitives::sha3::sha3_256;
12use rand::RngCore;
13
14#[derive(Debug, Clone)]
16pub struct EcSchnorrProof {
17 pub commitment: Vec<u8>,
19 pub response: Vec<u8>,
21}
22
23#[derive(Debug)]
25pub enum SchnorrError {
26 InvalidCommitment(String),
28 InvalidResponse(String),
30 InvalidPublicKey(String),
32 VerificationFailed,
34 RngFailed,
36}
37
38impl std::fmt::Display for SchnorrError {
39 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
40 match self {
41 SchnorrError::InvalidCommitment(msg) => write!(f, "Invalid commitment point: {msg}"),
42 SchnorrError::InvalidResponse(msg) => write!(f, "Invalid response scalar: {msg}"),
43 SchnorrError::InvalidPublicKey(msg) => write!(f, "Invalid public key: {msg}"),
44 SchnorrError::VerificationFailed => write!(f, "Proof verification failed"),
45 SchnorrError::RngFailed => write!(f, "Random number generation failed"),
46 }
47 }
48}
49
50impl std::error::Error for SchnorrError {}
51
52impl From<SchnorrError> for CryptoError {
53 fn from(e: SchnorrError) -> Self {
54 CryptoError::InvalidParameter(e.to_string())
55 }
56}
57
58pub fn prove(secret_key: &[u8; 32], public_key: &[u8], message: &[u8]) -> Result<EcSchnorrProof> {
60 let x = ct_option_to_option(Scalar::from_repr(secret_key))
61 .ok_or_else(|| CryptoError::InvalidParameter("Invalid secret key scalar".into()))?;
62
63 let r = {
65 let mut rng = rand::thread_rng();
66 let mut buf = [0u8; 32];
67 rng.fill_bytes(&mut buf);
68 ct_option_to_option(Scalar::from_repr(&buf)).ok_or(SchnorrError::RngFailed)?
69 };
70
71 let r_point = ProjectivePoint::GENERATOR.mul(&r);
73 let r_compressed = r_point.to_encoded_point(true);
74
75 let challenge = compute_challenge(r_compressed.as_bytes(), public_key, message);
77
78 let ex = challenge.mul(&x);
80 let s = r.add(&ex);
81
82 Ok(EcSchnorrProof {
83 commitment: r_compressed.as_bytes().to_vec(),
84 response: s.to_bytes().to_vec(),
85 })
86}
87
88pub fn verify(proof: &EcSchnorrProof, public_key: &[u8], message: &[u8]) -> Result<bool> {
90 let r_ctoption = EncodedPoint::from_bytes(&proof.commitment);
92 if !bool::from(r_ctoption.is_some()) {
93 return Err(SchnorrError::InvalidCommitment("Invalid encoding".into()).into());
94 }
95 let r_encoded = r_ctoption.unwrap();
96 let r_affine = ct_option_to_option(AffinePoint::from_encoded_point(&r_encoded))
97 .ok_or_else(|| SchnorrError::InvalidCommitment("Identity point".into()))?;
98
99 let s_bytes: [u8; 32] = proof
101 .response
102 .as_slice()
103 .try_into()
104 .map_err(|_| SchnorrError::InvalidResponse("Wrong length".into()))?;
105 let s = ct_option_to_option(Scalar::from_repr(&s_bytes))
106 .ok_or_else(|| SchnorrError::InvalidResponse("Not a valid scalar".into()))?;
107
108 let p_ctoption = EncodedPoint::from_bytes(public_key);
110 if !bool::from(p_ctoption.is_some()) {
111 return Err(SchnorrError::InvalidPublicKey("Invalid encoding".into()).into());
112 }
113 let p_encoded = p_ctoption.unwrap();
114 let p_affine = ct_option_to_option(AffinePoint::from_encoded_point(&p_encoded))
115 .ok_or_else(|| SchnorrError::InvalidPublicKey("Identity point".into()))?;
116
117 let challenge = compute_challenge(&proof.commitment, public_key, message);
119
120 let s_g = ProjectivePoint::GENERATOR.mul(&s);
122 let p_projective = ProjectivePoint::from(p_affine);
123 let e_p = p_projective.mul(&challenge);
124 let r_projective = ProjectivePoint::from(r_affine);
125 let r_plus_ep = r_projective.add(&e_p);
126
127 let equal = s_g.ct_eq(&r_plus_ep);
128 Ok(bool::from(equal))
129}
130
131pub fn batch_verify(
133 proofs: &[EcSchnorrProof],
134 public_keys: &[Vec<u8>],
135 messages: &[Vec<u8>],
136) -> Result<bool> {
137 if proofs.len() != public_keys.len() || proofs.len() != messages.len() {
138 return Err(CryptoError::InvalidParameter(
139 "Mismatched input lengths for batch verify".into(),
140 ));
141 }
142 for i in 0..proofs.len() {
143 if !verify(&proofs[i], &public_keys[i], &messages[i])? {
144 return Ok(false);
145 }
146 }
147 Ok(true)
148}
149
150fn compute_challenge(commitment: &[u8], public_key: &[u8], message: &[u8]) -> Scalar {
152 let mut input = Vec::with_capacity(23 + commitment.len() + public_key.len() + message.len());
154 input.extend_from_slice(b"ec-schnorr-challenge-v1");
155 input.extend_from_slice(commitment);
156 input.extend_from_slice(public_key);
157 input.extend_from_slice(message);
158 let hash = sha3_256(&input);
159
160 Scalar::from_repr_reduced(&hash)
165}
166
167pub fn generate_keypair(seed: &[u8; 32]) -> ([u8; 32], Vec<u8>) {
169 let mut input = Vec::with_capacity(23 + 32);
170 input.extend_from_slice(b"ec-schnorr-keygen-v2");
171 input.extend_from_slice(seed);
172 let mut counter = 0u32;
173
174 let secret = loop {
175 let hash = sha3_256(&input);
176 if let Some(s) = ct_option_to_option(Scalar::from_repr(&hash)) {
177 break s;
178 }
179 counter += 1;
180 input.truncate(23 + 32);
181 input.extend_from_slice(&counter.to_le_bytes());
182 };
183
184 let public = ProjectivePoint::GENERATOR.mul(&secret);
185 let public_compressed = public.to_encoded_point(true);
186
187 (
188 secret.to_bytes().into(),
189 public_compressed.as_bytes().to_vec(),
190 )
191}
192
193#[cfg(test)]
194mod tests {
195 use super::*;
196
197 #[test]
198 fn test_prove_verify_roundtrip() {
199 let seed = [42u8; 32];
200 let (sk_bytes, pk_bytes) = generate_keypair(&seed);
201 let message = b"test message for EC-Schnorr";
202 let proof = prove(&sk_bytes, &pk_bytes, message).unwrap();
203 assert!(verify(&proof, &pk_bytes, message).unwrap());
204 }
205
206 #[test]
207 fn test_verify_wrong_message_fails() {
208 let seed = [42u8; 32];
209 let (sk_bytes, pk_bytes) = generate_keypair(&seed);
210 let proof = prove(&sk_bytes, &pk_bytes, b"original message").unwrap();
211 assert!(!verify(&proof, &pk_bytes, b"different message").unwrap());
212 }
213
214 #[test]
215 fn test_verify_wrong_public_key_fails() {
216 let (sk1, pk1) = generate_keypair(&[42u8; 32]);
217 let (_, pk2) = generate_keypair(&[99u8; 32]);
218 let proof = prove(&sk1, &pk1, b"test message").unwrap();
219 assert!(!verify(&proof, &pk2, b"test message").unwrap());
220 }
221
222 #[test]
223 fn test_deterministic_keypair() {
224 let (sk1, pk1) = generate_keypair(&[7u8; 32]);
225 let (sk2, pk2) = generate_keypair(&[7u8; 32]);
226 assert_eq!(sk1, sk2);
227 assert_eq!(pk1, pk2);
228 }
229
230 #[test]
231 fn test_different_seeds_different_keys() {
232 let (sk1, pk1) = generate_keypair(&[1u8; 32]);
233 let (sk2, pk2) = generate_keypair(&[2u8; 32]);
234 assert_ne!(sk1, sk2);
235 assert_ne!(pk1, pk2);
236 }
237
238 #[test]
239 fn test_batch_verify_all_valid() {
240 let mut proofs = Vec::new();
241 let mut pks = Vec::new();
242 let mut msgs = Vec::new();
243 for i in 0..10u8 {
244 let (sk, pk) = generate_keypair(&[i; 32]);
245 let msg = format!("message {}", i).into_bytes();
246 proofs.push(prove(&sk, &pk, &msg).unwrap());
247 pks.push(pk);
248 msgs.push(msg);
249 }
250 assert!(batch_verify(&proofs, &pks, &msgs).unwrap());
251 }
252
253 #[test]
254 fn test_batch_verify_one_invalid() {
255 let mut proofs = Vec::new();
256 let mut pks = Vec::new();
257 let mut msgs = Vec::new();
258 for i in 0..3u8 {
259 let (sk, pk) = generate_keypair(&[i; 32]);
260 let msg = format!("message {}", i).into_bytes();
261 proofs.push(prove(&sk, &pk, &msg).unwrap());
262 pks.push(pk);
263 msgs.push(msg);
264 }
265 let (sk, pk) = generate_keypair(&[99u8; 32]);
266 proofs.push(prove(&sk, &pk, b"correct message").unwrap());
267 pks.push(pk);
268 msgs.push(b"wrong message".to_vec());
269 assert!(!batch_verify(&proofs, &pks, &msgs).unwrap());
270 }
271
272 #[test]
273 fn test_empty_message() {
274 let seed = [42u8; 32];
275 let (sk_bytes, pk_bytes) = generate_keypair(&seed);
276 let proof = prove(&sk_bytes, &pk_bytes, b"").unwrap();
277 assert!(verify(&proof, &pk_bytes, b"").unwrap());
278 }
279
280 #[test]
281 fn test_large_message() {
282 let seed = [42u8; 32];
283 let (sk_bytes, pk_bytes) = generate_keypair(&seed);
284 let message = vec![0xABu8; 1_000_000];
285 let proof = prove(&sk_bytes, &pk_bytes, &message).unwrap();
286 assert!(verify(&proof, &pk_bytes, &message).unwrap());
287 }
288}