Skip to main content

solana_primitives/types/
pda.rs

1use crate::error::{Result, SolanaError};
2use crate::types::Pubkey;
3use ed25519_dalek::VerifyingKey;
4use sha2::{Digest, Sha256};
5
6/// Maximum number of seeds allowed in a PDA
7pub const MAX_SEEDS: usize = 16;
8/// Maximum length of a seed in bytes
9pub const MAX_SEED_LEN: usize = 32;
10
11/// Find a program address and bump seed for the given seeds
12pub fn find_program_address(program_id: &Pubkey, seeds: &[&[u8]]) -> Result<(Pubkey, u8)> {
13    // The bump seed occupies one of the MAX_SEEDS slots.
14    if seeds.len() >= MAX_SEEDS {
15        return Err(SolanaError::InvalidPubkey(format!(
16            "too many seeds: {}, max: {}",
17            seeds.len(),
18            MAX_SEEDS
19        )));
20    }
21    for seed in seeds {
22        if seed.len() > MAX_SEED_LEN {
23            return Err(SolanaError::InvalidPubkey(format!(
24                "seed too long: {}, max: {}",
25                seed.len(),
26                MAX_SEED_LEN
27            )));
28        }
29    }
30
31    // Try each bump seed until we find a valid PDA
32    let mut bump = 255;
33    loop {
34        let mut hasher = Sha256::new();
35
36        // Hash all seeds
37        for seed in seeds {
38            hasher.update(seed);
39        }
40
41        // Add bump seed
42        hasher.update([bump]);
43
44        // Add program ID
45        hasher.update(program_id.as_bytes());
46
47        // Add "ProgramDerivedAddress" as a domain separator
48        hasher.update(b"ProgramDerivedAddress");
49
50        // Get the hash result
51        let hash = hasher.finalize();
52
53        // Convert hash to pubkey
54        let mut pubkey_bytes = [0u8; 32];
55        pubkey_bytes.copy_from_slice(&hash[..32]);
56
57        // Check if it's on curve
58        if !is_on_curve(&pubkey_bytes) {
59            // Found a valid PDA
60            return Ok((Pubkey::new(pubkey_bytes), bump));
61        }
62
63        if bump == 0 {
64            return Err(SolanaError::InvalidPubkey(
65                "unable to find valid PDA, all bump seeds exhausted".to_string(),
66            ));
67        }
68        bump -= 1;
69    }
70}
71
72/// Create a program address from seeds and a bump seed
73pub fn create_program_address(
74    program_id: &Pubkey,
75    seeds: &[&[u8]],
76    bump_seed: u8,
77) -> Result<Pubkey> {
78    // The bump seed occupies one of the MAX_SEEDS slots.
79    if seeds.len() >= MAX_SEEDS {
80        return Err(SolanaError::InvalidPubkey(format!(
81            "too many seeds: {}, max: {}",
82            seeds.len(),
83            MAX_SEEDS
84        )));
85    }
86    for seed in seeds {
87        if seed.len() > MAX_SEED_LEN {
88            return Err(SolanaError::InvalidPubkey(format!(
89                "seed too long: {}, max: {}",
90                seed.len(),
91                MAX_SEED_LEN
92            )));
93        }
94    }
95
96    let mut hasher = Sha256::new();
97
98    // Hash all seeds
99    for seed in seeds {
100        hasher.update(seed);
101    }
102
103    // Add bump seed
104    hasher.update([bump_seed]);
105
106    // Add program ID
107    hasher.update(program_id.as_bytes());
108
109    // Add "ProgramDerivedAddress" as a domain separator
110    hasher.update(b"ProgramDerivedAddress");
111
112    // Get the hash result
113    let hash = hasher.finalize();
114
115    // Convert hash to pubkey
116    let mut pubkey_bytes = [0u8; 32];
117    pubkey_bytes.copy_from_slice(&hash[..32]);
118
119    // Check if it's on curve
120    if is_on_curve(&pubkey_bytes) {
121        return Err(SolanaError::InvalidPubkey(
122            "resulting address is on curve (invalid PDA)".to_string(),
123        ));
124    }
125
126    Ok(Pubkey::new(pubkey_bytes))
127}
128
129/// Check if a public key is on the ed25519 curve
130pub fn is_on_curve(bytes: &[u8; 32]) -> bool {
131    // Check if the point is all zeros
132    if bytes.iter().all(|&b| b == 0) {
133        return false;
134    }
135    // Try to decompress the point
136    VerifyingKey::from_bytes(bytes).is_ok()
137}
138
139#[cfg(test)]
140mod tests {
141    use super::*;
142    use std::str::FromStr;
143
144    fn create_test_program_id() -> Pubkey {
145        crate::instructions::program_ids::system_program()
146    }
147
148    #[test]
149    fn test_find_program_address() {
150        let program_id = create_test_program_id();
151        let seed = b"test_seed";
152        let seeds = [seed.as_ref()];
153
154        let (pda, bump) = find_program_address(&program_id, &seeds).unwrap();
155
156        // Verify the PDA is off curve
157        assert!(!is_on_curve(pda.as_bytes()));
158
159        // Verify we can recreate the PDA with the bump
160        let recreated_pda = create_program_address(&program_id, &seeds, bump).unwrap();
161        assert_eq!(pda, recreated_pda);
162    }
163
164    #[test]
165    fn test_find_program_address_multiple_seeds() {
166        let program_id = create_test_program_id();
167        let seed1 = b"seed1";
168        let seed2 = b"seed2";
169        let seed3 = b"seed3";
170        let seeds = [seed1.as_ref(), seed2.as_ref(), seed3.as_ref()];
171
172        let (pda, bump) = find_program_address(&program_id, &seeds).unwrap();
173
174        // Verify the PDA is off curve
175        assert!(!is_on_curve(pda.as_bytes()));
176
177        // Verify we can recreate the PDA with the bump
178        let recreated_pda = create_program_address(&program_id, &seeds, bump).unwrap();
179        assert_eq!(pda, recreated_pda);
180    }
181
182    #[test]
183    fn test_find_program_address_too_many_seeds() {
184        let program_id = create_test_program_id();
185        let seed_strings: Vec<String> = (0..MAX_SEEDS + 1).map(|i| format!("seed{i}")).collect();
186        let seed_refs: Vec<&[u8]> = seed_strings.iter().map(|s| s.as_bytes()).collect();
187
188        let result = find_program_address(&program_id, &seed_refs);
189        assert!(matches!(result, Err(SolanaError::InvalidPubkey(_))));
190    }
191
192    #[test]
193    fn test_find_program_address_max_seeds_rejected() {
194        let program_id = create_test_program_id();
195        let seed_strings: Vec<String> = (0..MAX_SEEDS).map(|i| format!("seed{i}")).collect();
196        let seed_refs: Vec<&[u8]> = seed_strings.iter().map(|s| s.as_bytes()).collect();
197
198        let result = find_program_address(&program_id, &seed_refs);
199        assert!(matches!(result, Err(SolanaError::InvalidPubkey(_))));
200    }
201
202    #[test]
203    fn test_find_program_address_max_seeds_minus_one_succeeds() {
204        let program_id = create_test_program_id();
205        let seed_strings: Vec<String> = (0..MAX_SEEDS - 1).map(|i| format!("seed{i}")).collect();
206        let seed_refs: Vec<&[u8]> = seed_strings.iter().map(|s| s.as_bytes()).collect();
207
208        let (pda, bump) = find_program_address(&program_id, &seed_refs).unwrap();
209        assert!(!is_on_curve(pda.as_bytes()));
210
211        let recreated_pda = create_program_address(&program_id, &seed_refs, bump).unwrap();
212        assert_eq!(pda, recreated_pda);
213    }
214
215    #[test]
216    fn test_find_program_address_seed_too_long() {
217        let program_id = create_test_program_id();
218        let seed = [0u8; MAX_SEED_LEN + 1];
219        let seeds = [&seed[..]];
220
221        let result = find_program_address(&program_id, &seeds);
222        assert!(matches!(result, Err(SolanaError::InvalidPubkey(_))));
223    }
224
225    #[test]
226    fn test_create_program_address() {
227        let program_id = create_test_program_id();
228        let seed = b"test_seed";
229        let seeds = [seed.as_ref()];
230        let bump = 255;
231
232        let pda = create_program_address(&program_id, &seeds, bump).unwrap();
233
234        // Verify the PDA is off curve
235        assert!(!is_on_curve(pda.as_bytes()));
236    }
237
238    #[test]
239    fn test_create_program_address_max_seeds_rejected() {
240        let program_id = create_test_program_id();
241        let seed_strings: Vec<String> = (0..MAX_SEEDS).map(|i| format!("seed{i}")).collect();
242        let seed_refs: Vec<&[u8]> = seed_strings.iter().map(|s| s.as_bytes()).collect();
243
244        let result = create_program_address(&program_id, &seed_refs, 255);
245        assert!(matches!(result, Err(SolanaError::InvalidPubkey(_))));
246    }
247
248    #[test]
249    fn test_create_program_address_max_seeds_minus_one_succeeds() {
250        let program_id = create_test_program_id();
251        let seed_strings: Vec<String> = (0..MAX_SEEDS - 1).map(|i| format!("seed{i}")).collect();
252        let seed_refs: Vec<&[u8]> = seed_strings.iter().map(|s| s.as_bytes()).collect();
253
254        let (_, bump) = find_program_address(&program_id, &seed_refs).unwrap();
255        let pda = create_program_address(&program_id, &seed_refs, bump).unwrap();
256        assert!(!is_on_curve(pda.as_bytes()));
257    }
258
259    #[test]
260    fn test_create_program_address_on_curve() {
261        let program_id = create_test_program_id();
262        // Try different seeds and bumps to verify we never get a point on the curve
263        for i in 0..10 {
264            let seed = format!("test_seed_{i}");
265            let seeds = [seed.as_bytes()];
266            for bump in 0..10 {
267                if let Ok(pubkey) = create_program_address(&program_id, &seeds, bump) {
268                    assert!(
269                        !is_on_curve(pubkey.as_bytes()),
270                        "Found point on curve with seed {i} and bump {bump}"
271                    );
272                }
273            }
274        }
275    }
276
277    #[test]
278    fn test_is_on_curve() {
279        // Test a valid ed25519 public key (base point)
280        let valid_key = [
281            0x58, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66,
282            0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66, 0x66,
283            0x66, 0x66, 0x66, 0x66,
284        ];
285        assert!(is_on_curve(&valid_key));
286
287        // Test an invalid public key (all zeros)
288        let invalid_key = [0u8; 32];
289        assert!(!is_on_curve(&invalid_key));
290    }
291
292    #[test]
293    fn test_pda_deterministic() {
294        let program_id = create_test_program_id();
295        let seed = b"test_seed";
296        let seeds = [seed.as_ref()];
297
298        // Generate PDA twice with same inputs
299        let (pda1, bump1) = find_program_address(&program_id, &seeds).unwrap();
300        let (pda2, bump2) = find_program_address(&program_id, &seeds).unwrap();
301
302        // Verify results are identical
303        assert_eq!(pda1, pda2);
304        assert_eq!(bump1, bump2);
305    }
306
307    #[test]
308    fn test_pda_different_program_ids() {
309        let program_id1 = create_test_program_id();
310        let program_id2 = Pubkey::new([
311            1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3, 4, 4, 4, 4, 4,
312            4, 4, 4,
313        ]);
314        let seed = b"test_seed";
315        let seeds = [seed.as_ref()];
316
317        let (pda1, _bump1) = find_program_address(&program_id1, &seeds).unwrap();
318        let (pda2, _bump2) = find_program_address(&program_id2, &seeds).unwrap();
319
320        // Verify PDAs are different
321        assert_ne!(pda1, pda2);
322        // Both PDAs should be off curve
323        assert!(!is_on_curve(pda1.as_bytes()));
324        assert!(!is_on_curve(pda2.as_bytes()));
325    }
326
327    #[test]
328    fn test_pda_matches_js_example() {
329        let program_id = crate::instructions::program_ids::system_program();
330        let string = b"helloWorld";
331        let seeds = [string.as_ref()];
332
333        let (pda, bump) = find_program_address(&program_id, &seeds).unwrap();
334
335        // Expected values from JS example:
336        // PDA: 46GZzzetjCURsdFPb7rcnspbEMnCBXe9kpjrsZAkKb6X
337        // Bump: 254
338        let expected_pda =
339            Pubkey::from_str("46GZzzetjCURsdFPb7rcnspbEMnCBXe9kpjrsZAkKb6X").unwrap();
340        let expected_bump = 254;
341
342        assert_eq!(pda, expected_pda, "PDA does not match expected value");
343        assert_eq!(
344            bump, expected_bump,
345            "Bump seed does not match expected value"
346        );
347
348        // Verify we can recreate the PDA with the bump
349        let recreated_pda = create_program_address(&program_id, &seeds, bump).unwrap();
350        assert_eq!(recreated_pda, expected_pda);
351    }
352}