1use crate::error::{Result, SolanaError};
2use crate::types::Pubkey;
3use ed25519_dalek::VerifyingKey;
4use sha2::{Digest, Sha256};
5
6pub const MAX_SEEDS: usize = 16;
8pub const MAX_SEED_LEN: usize = 32;
10
11pub fn find_program_address(program_id: &Pubkey, seeds: &[&[u8]]) -> Result<(Pubkey, u8)> {
13 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 let mut bump = 255;
33 loop {
34 let mut hasher = Sha256::new();
35
36 for seed in seeds {
38 hasher.update(seed);
39 }
40
41 hasher.update([bump]);
43
44 hasher.update(program_id.as_bytes());
46
47 hasher.update(b"ProgramDerivedAddress");
49
50 let hash = hasher.finalize();
52
53 let mut pubkey_bytes = [0u8; 32];
55 pubkey_bytes.copy_from_slice(&hash[..32]);
56
57 if !is_on_curve(&pubkey_bytes) {
59 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
72pub fn create_program_address(
74 program_id: &Pubkey,
75 seeds: &[&[u8]],
76 bump_seed: u8,
77) -> Result<Pubkey> {
78 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 for seed in seeds {
100 hasher.update(seed);
101 }
102
103 hasher.update([bump_seed]);
105
106 hasher.update(program_id.as_bytes());
108
109 hasher.update(b"ProgramDerivedAddress");
111
112 let hash = hasher.finalize();
114
115 let mut pubkey_bytes = [0u8; 32];
117 pubkey_bytes.copy_from_slice(&hash[..32]);
118
119 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
129pub fn is_on_curve(bytes: &[u8; 32]) -> bool {
131 if bytes.iter().all(|&b| b == 0) {
133 return false;
134 }
135 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 assert!(!is_on_curve(pda.as_bytes()));
158
159 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 assert!(!is_on_curve(pda.as_bytes()));
176
177 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 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 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 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 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 let (pda1, bump1) = find_program_address(&program_id, &seeds).unwrap();
300 let (pda2, bump2) = find_program_address(&program_id, &seeds).unwrap();
301
302 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 assert_ne!(pda1, pda2);
322 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 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 let recreated_pda = create_program_address(&program_id, &seeds, bump).unwrap();
350 assert_eq!(recreated_pda, expected_pda);
351 }
352}