Skip to main content

binius_examples/circuits/
zklogin.rs

1// Copyright 2025 Irreducible Inc.
2use anyhow::{Result, ensure};
3use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD as BASE64_URL_SAFE_NO_PAD};
4use binius_circuits::{
5	base64::Base64UrlSafe,
6	bytes::swap_bytes,
7	concat::concat,
8	fixed_byte_vec::ByteVec,
9	jwt_claims::{Attribute, JwtClaims},
10	rs256::Rs256Verify,
11	sha256::sha256_varlen,
12	slice::assert_slice_eq,
13};
14use binius_core::Word;
15use binius_frontend::{CircuitBuilder, Wire, WitnessFiller};
16use clap::Args;
17use jwt_simple::prelude::*;
18use rand::prelude::*;
19use sha2::{Digest, Sha256};
20
21use crate::ExampleCircuit;
22
23/// The configuration of the ZKLogin circuit.
24///
25/// Picking the numbers are a tradeoff. Picking a large number will require a larger circuit and
26/// thus more proving time. Picking a small number may make some statements unprovable.
27#[derive(Debug, Clone)]
28pub struct Config {
29	/// Maximum length in wires of the base64 decoded JWT header. Must be a multiple of 8.
30	pub max_len_json_jwt_header: usize,
31	/// Maximum length in wires of the base64 decoded JWT payload. Must be a multiple of 8.
32	pub max_len_json_jwt_payload: usize,
33	/// Maximum length in wires of the base64 decoded JWT signature. Must be a multiple of 8.
34	pub max_len_jwt_signature: usize,
35	pub max_len_jwt_sub: usize,
36	pub max_len_jwt_aud: usize,
37	pub max_len_jwt_iss: usize,
38	pub max_len_salt: usize,
39	pub max_len_nonce_r: usize,
40	pub max_len_t_max: usize,
41}
42
43impl Default for Config {
44	fn default() -> Self {
45		Self {
46			max_len_json_jwt_header: 33,
47			max_len_json_jwt_payload: 63,
48			max_len_jwt_signature: 33,
49			max_len_jwt_sub: 9,
50			max_len_jwt_aud: 9,
51			max_len_jwt_iss: 9,
52			max_len_salt: 9,
53			max_len_nonce_r: 6,
54			max_len_t_max: 6,
55		}
56	}
57}
58
59impl Config {
60	pub const fn max_len_base64_jwt_header(&self) -> usize {
61		self.max_len_json_jwt_header.div_ceil(3) * 4
62	}
63
64	pub const fn max_len_base64_jwt_payload(&self) -> usize {
65		self.max_len_json_jwt_payload.div_ceil(3) * 4
66	}
67
68	pub const fn max_len_base64_jwt_signature(&self) -> usize {
69		self.max_len_jwt_signature.div_ceil(3) * 4
70	}
71}
72
73/// A circuit that implements zk login.
74pub struct ZkLogin {
75	/// The sub claim value
76	pub sub: ByteVec,
77	/// The aud claim value
78	pub aud: ByteVec,
79	/// The iss claim value
80	pub iss: ByteVec,
81	/// The salt value
82	pub salt: ByteVec,
83	/// The zkaddr (SHA256 hash of concat(sub, aud, iss, salt))
84	pub zkaddr: [Wire; 4],
85	/// The message ByteVec whose SHA-256 is asserted to equal `zkaddr`
86	pub zkaddr_sha256: ByteVec,
87	/// The subcircuit that verifies the JWT header.
88	pub jwt_claims_header: JwtClaims,
89	/// The subcircuit that verifies the JWT in the payload.
90	pub jwt_claims_payload: JwtClaims,
91	/// The subcircuit that verifies the RS256 signature in the JWT.
92	pub jwt_signature_verify: Rs256Verify,
93	/// The JWT header
94	pub base64_jwt_header: ByteVec,
95	/// The JWT payload
96	pub base64_jwt_payload: ByteVec,
97	/// The JWT signature
98	pub base64_jwt_signature: ByteVec,
99	/// The decoded JWT header
100	pub jwt_header: ByteVec,
101	/// The decoded jwt_payload
102	pub jwt_payload: ByteVec,
103	/// The decoded jwt_signature (264 bytes for Base64, little-endian packing)
104	pub jwt_signature: ByteVec,
105	/// The base64 encoded nonce
106	pub base64_jwt_payload_nonce: [Wire; 6],
107	/// The message ByteVec whose SHA-256 is asserted to equal `nonce`
108	pub nonce_sha256: ByteVec,
109	/// The nonce value (32 bytes SHA256 hash)
110	pub nonce: [Wire; 4],
111	/// The vk_u public key (32 bytes)
112	pub vk_u: [Wire; 4],
113	/// The t_max value
114	pub t_max: ByteVec,
115	/// The nonce_r value
116	pub nonce_r: ByteVec,
117}
118
119impl ZkLogin {
120	pub fn new(b: &mut CircuitBuilder, config: &Config) -> Self {
121		let sub = ByteVec::new_inout(b, config.max_len_jwt_sub);
122		let aud = ByteVec::new_inout(b, config.max_len_jwt_aud);
123		let iss = ByteVec::new_inout(b, config.max_len_jwt_iss);
124		let salt = ByteVec::new_inout(b, config.max_len_salt);
125
126		let base64_jwt_header = ByteVec::new_inout(b, config.max_len_base64_jwt_header());
127		let base64_jwt_payload = ByteVec::new_inout(b, config.max_len_base64_jwt_payload());
128		let base64_jwt_signature = ByteVec::new_inout(b, config.max_len_base64_jwt_signature());
129
130		let jwt_header = ByteVec::new_inout(b, config.max_len_json_jwt_header);
131		let jwt_payload = ByteVec::new_witness(b, config.max_len_json_jwt_payload);
132		let jwt_signature = ByteVec::new_witness(b, config.max_len_jwt_signature);
133
134		let t_max = ByteVec::new_inout(b, config.max_len_t_max);
135		let nonce_r = ByteVec::new_witness(b, config.max_len_nonce_r);
136
137		let zkaddr: [Wire; 4] = std::array::from_fn(|_| b.add_inout());
138		let vk_u: [Wire; 4] = std::array::from_fn(|_| b.add_inout());
139		let nonce: [Wire; 4] = std::array::from_fn(|_| b.add_witness());
140
141		// The base64 encoded nonce in the JWT payload. This must have
142		// 6 wires = 48 bytes to accommodate the 43-byte base64 nonce with padding.
143		let base64_jwt_payload_nonce: [Wire; 6] = std::array::from_fn(|_| b.add_witness());
144
145		// RSA modulus as public input (256 bytes for 2048-bit RSA)
146		let rsa_modulus = ByteVec::new_inout(b, 32);
147
148		// Decode JWT.
149		// 1. header
150		// 2. payload
151		// 3. signature
152
153		let _base64decode_check_header = Base64UrlSafe::new(
154			&b.subcircuit("base64_check_header"),
155			jwt_header.data.clone(),
156			base64_jwt_header.data.clone(),
157			jwt_header.len_bytes,
158		);
159		let _base64decode_check_payload = Base64UrlSafe::new(
160			&b.subcircuit("base64_check_payload"),
161			jwt_payload.data.clone(),
162			base64_jwt_payload.data.clone(),
163			jwt_payload.len_bytes,
164		);
165		let _base64decode_check_signature = Base64UrlSafe::new(
166			&b.subcircuit("base64_check_signature"),
167			jwt_signature.data.clone(),
168			base64_jwt_signature.data.clone(),
169			jwt_signature.len_bytes,
170		);
171
172		// We need to check
173		//
174		// X = concat(JWT.sub, JWT.aud, JWT.iss, salt)
175		// assert zkaddr == SHA256(X)
176		let max_len_zkaddr_preimage = config.max_len_jwt_sub
177			+ config.max_len_jwt_aud
178			+ config.max_len_jwt_iss
179			+ config.max_len_salt;
180
181		let zkaddr_preimage = concat(
182			&b.subcircuit("zkaddr_preimage_concat"),
183			&[sub.clone(), aud.clone(), iss.clone(), salt.clone()],
184		);
185
186		let zkaddr_sha256_message: Vec<Wire> = (0..max_len_zkaddr_preimage)
187			.map(|_| b.add_witness())
188			.collect();
189		let zkaddr_sha256 = ByteVec::new(zkaddr_sha256_message, zkaddr_preimage.len_bytes);
190		let zkaddr_computed = sha256_varlen(&b.subcircuit("zkaddr_sha256"), &zkaddr_sha256);
191		for i in 0..4 {
192			b.assert_eq(format!("zkaddr_digest[{i}]"), zkaddr_computed[i], zkaddr[i]);
193		}
194
195		assert_slice_eq(
196			b,
197			"zkaddr_preimage_eq",
198			zkaddr_preimage.len_bytes,
199			&zkaddr_sha256.data,
200			&zkaddr_preimage.data,
201		);
202
203		// We need to check:
204		//
205		// nonce_preimage = concat(vk_u, T_max, r) where vk_u is a public key
206		// assert nonce = SHA256(nonce_preimage)
207		// assert nonce = base64_decode(base64_jwt_payload_nonce)
208		let max_len_nonce_preimage = 4 + config.max_len_t_max + config.max_len_nonce_r;
209
210		let nonce_preimage = concat(
211			&b.subcircuit("nonce_preimage_concat"),
212			&[
213				ByteVec::new(vk_u.to_vec(), b.add_constant_64(32)),
214				t_max.clone(),
215				nonce_r.clone(),
216			],
217		);
218
219		let nonce_sha256_message: Vec<Wire> = (0..max_len_nonce_preimage)
220			.map(|_| b.add_witness())
221			.collect();
222		let nonce_sha256 = ByteVec::new(nonce_sha256_message, nonce_preimage.len_bytes);
223		let nonce_computed = sha256_varlen(&b.subcircuit("nonce_sha256"), &nonce_sha256);
224		for i in 0..4 {
225			b.assert_eq(format!("nonce_digest[{i}]"), nonce_computed[i], nonce[i]);
226		}
227
228		assert_slice_eq(
229			b,
230			"nonce_preimage_eq",
231			nonce_preimage.len_bytes,
232			&nonce_sha256.data,
233			&nonce_preimage.data,
234		);
235
236		let nonce_le = nonce_computed.map(|x| swap_bytes(b, x));
237
238		// Base64 requires 48 bytes (6 wires) for alignment, so add zero padding
239		let zero = b.add_constant(Word::ZERO);
240		let nonce_le_for_base64: Vec<Wire> = nonce_le.into_iter().chain([zero, zero]).collect();
241
242		// The zklogin nonce claim is Base64 URL encoded without padding (i.e.
243		// in the same way as JWS components)
244		// <https://github.com/MystenLabs/ts-sdks/blob/eb23fc1c122a1495e52d0bd613bf5e8e6eb816cc/packages/typescript/src/zklogin/nonce.ts#L33>
245		//
246		// The nonce is 32 bytes which encodes to 43 base64 characters.
247		// minimal wires those will fit into: 6 wires.
248		let base64_check_nonce_builder = b.subcircuit("base64_check_nonce");
249		let _base64decode_check_nonce = Base64UrlSafe::new(
250			&base64_check_nonce_builder,
251			nonce_le_for_base64,
252			base64_jwt_payload_nonce.to_vec(),
253			base64_check_nonce_builder.add_constant_64(32),
254		);
255
256		// Check signing payload. The JWT signed payload L is a concatenation of:
257		//
258		// L = concat(jwt.header | "." | jwt.payload)
259		//
260		let max_len_jwt_signing_payload =
261			config.max_len_base64_jwt_header() + 1 + config.max_len_base64_jwt_payload();
262
263		let signing_payload = concat(
264			&b.subcircuit("jwt_signing_payload_concat"),
265			&[
266				base64_jwt_header.clone(),
267				ByteVec::new(vec![b.add_constant_zx_8(b'.')], b.add_constant_64(1)),
268				base64_jwt_payload.clone(),
269			],
270		);
271
272		let jwt_signing_payload_sha256_message: Vec<Wire> = (0..max_len_jwt_signing_payload)
273			.map(|_| b.add_witness())
274			.collect();
275
276		let jwt_signing_payload =
277			ByteVec::new(jwt_signing_payload_sha256_message, signing_payload.len_bytes);
278
279		let jwt_signature_verify =
280			Rs256Verify::new(b, jwt_signing_payload, jwt_signature.clone(), rsa_modulus);
281
282		let jwt_signing_payload_le_wires = jwt_signature_verify.message.data.clone();
283		assert_slice_eq(
284			b,
285			"jwt_signing_payload_eq",
286			signing_payload.len_bytes,
287			&jwt_signing_payload_le_wires,
288			&signing_payload.data,
289		);
290
291		let jwt_claims_header = jwt_header_check(b, &jwt_header);
292		let jwt_claims_payload =
293			jwt_payload_check(b, &jwt_payload, &sub, &aud, &iss, &base64_jwt_payload_nonce);
294
295		Self {
296			sub,
297			aud,
298			iss,
299			salt,
300			zkaddr,
301			zkaddr_sha256,
302			jwt_claims_header,
303			jwt_claims_payload,
304			jwt_signature_verify,
305			base64_jwt_header,
306			base64_jwt_payload,
307			base64_jwt_signature,
308			jwt_header,
309			jwt_payload,
310			jwt_signature,
311			base64_jwt_payload_nonce,
312			nonce_sha256,
313			nonce,
314			vk_u,
315			t_max,
316			nonce_r,
317		}
318	}
319
320	pub fn populate_sub(&self, w: &mut WitnessFiller<'_>, sub_bytes: &[u8]) {
321		self.sub.populate_bytes_le(w, sub_bytes);
322	}
323
324	pub fn populate_aud(&self, w: &mut WitnessFiller<'_>, aud_bytes: &[u8]) {
325		self.aud.populate_bytes_le(w, aud_bytes);
326	}
327
328	pub fn populate_iss(&self, w: &mut WitnessFiller<'_>, iss_bytes: &[u8]) {
329		self.iss.populate_bytes_le(w, iss_bytes);
330	}
331
332	pub fn populate_salt(&self, w: &mut WitnessFiller<'_>, salt_bytes: &[u8]) {
333		self.salt.populate_bytes_le(w, salt_bytes);
334	}
335
336	pub fn populate_zkaddr(&self, w: &mut WitnessFiller<'_>, zkaddr_hash: &[u8; 32]) {
337		for (i, chunk) in zkaddr_hash.chunks(8).enumerate() {
338			w[self.zkaddr[i]] = Word(u64::from_be_bytes(chunk.try_into().unwrap()));
339		}
340	}
341
342	pub fn populate_zkaddr_preimage(&self, w: &mut WitnessFiller<'_>, zkaddr_preimage: &[u8]) {
343		self.zkaddr_sha256
344			.populate_len_bytes(w, zkaddr_preimage.len());
345		self.zkaddr_sha256.populate_data(w, zkaddr_preimage);
346	}
347
348	pub fn populate_jwt_header(&self, w: &mut WitnessFiller<'_>, header_bytes: &[u8]) {
349		self.jwt_header.populate_bytes_le(w, header_bytes);
350	}
351
352	pub fn populate_jwt_payload(&self, w: &mut WitnessFiller<'_>, payload_bytes: &[u8]) {
353		self.jwt_payload.populate_bytes_le(w, payload_bytes);
354	}
355
356	pub fn populate_jwt_signature(&self, w: &mut WitnessFiller<'_>, signature_bytes: &[u8]) {
357		assert_eq!(signature_bytes.len(), 256, "RSA signature must be 256 bytes");
358		self.jwt_signature.populate_bytes_le(w, signature_bytes);
359	}
360
361	pub fn populate_base64_jwt_header(&self, w: &mut WitnessFiller<'_>, bytes: &[u8]) {
362		self.base64_jwt_header.populate_bytes_le(w, bytes);
363	}
364
365	pub fn populate_base64_jwt_payload(&self, w: &mut WitnessFiller<'_>, bytes: &[u8]) {
366		self.base64_jwt_payload.populate_bytes_le(w, bytes);
367	}
368
369	pub fn populate_base64_jwt_signature(&self, w: &mut WitnessFiller<'_>, bytes: &[u8]) {
370		self.base64_jwt_signature.populate_bytes_le(w, bytes);
371	}
372
373	pub fn populate_rsa_modulus(&self, w: &mut WitnessFiller<'_>, modulus_bytes: &[u8]) {
374		self.jwt_signature_verify
375			.modulus
376			.populate_bytes_le(w, modulus_bytes);
377	}
378
379	pub fn populate_jwt_header_attributes(&self, w: &mut WitnessFiller<'_>) {
380		// Populate the expected lengths for "alg" and "typ" attributes
381		self.jwt_claims_header.attributes[0].populate_len_bytes(w, 5); // "RS256" is 5 bytes
382		self.jwt_claims_header.attributes[1].populate_len_bytes(w, 3); // "JWT" is 3 bytes
383	}
384
385	pub fn populate_nonce(&self, w: &mut WitnessFiller<'_>, nonce_hash: &[u8; 32]) {
386		for (i, chunk) in nonce_hash.chunks(8).enumerate() {
387			w[self.nonce[i]] = Word(u64::from_be_bytes(chunk.try_into().unwrap()));
388		}
389	}
390
391	pub fn populate_nonce_preimage(&self, w: &mut WitnessFiller<'_>, nonce_preimage: &[u8]) {
392		self.nonce_sha256
393			.populate_len_bytes(w, nonce_preimage.len());
394		self.nonce_sha256.populate_data(w, nonce_preimage);
395	}
396
397	pub fn populate_vk_u(&self, w: &mut WitnessFiller<'_>, vk_u_bytes: &[u8; 32]) {
398		w.pack_bytes_le(&self.vk_u, vk_u_bytes);
399	}
400
401	pub fn populate_t_max(&self, w: &mut WitnessFiller<'_>, t_max_bytes: &[u8]) {
402		self.t_max.populate_bytes_le(w, t_max_bytes);
403	}
404
405	pub fn populate_nonce_r(&self, w: &mut WitnessFiller<'_>, nonce_r_bytes: &[u8]) {
406		self.nonce_r.populate_bytes_le(w, nonce_r_bytes);
407	}
408
409	pub fn populate_base64_jwt_payload_nonce(
410		&self,
411		w: &mut WitnessFiller<'_>,
412		base64_nonce: &[u8],
413	) {
414		// The base64 nonce is 43 characters, but we need to pad to 48 bytes (6 wires)
415		let mut padded = vec![0u8; 48];
416		padded[..base64_nonce.len()].copy_from_slice(&base64_nonce[..base64_nonce.len()]);
417		w.pack_bytes_le(&self.base64_jwt_payload_nonce, &padded);
418	}
419}
420
421/// A check that verifies that JWT header has the expected constant values in the `alg` and `typ`
422/// fields.
423fn jwt_header_check(b: &CircuitBuilder, jwt_header: &ByteVec) -> JwtClaims {
424	JwtClaims::new(
425		&b.subcircuit("jwt_claims_header"),
426		jwt_header.len_bytes,
427		jwt_header.data.clone(),
428		vec![
429			Attribute {
430				name: "alg",
431				len_bytes: b.add_inout(),
432				value: vec![b.add_constant_64(u64::from_le_bytes(*b"RS256\0\0\0"))],
433			},
434			Attribute {
435				name: "typ",
436				len_bytes: b.add_inout(),
437				value: vec![b.add_constant_64(u64::from_le_bytes(*b"JWT\0\0\0\0\0"))],
438			},
439		],
440	)
441}
442
443/// A check that verifies that the payload has all the claimed values of `sub`, `aud`, `iss`
444/// and `nonce`.
445fn jwt_payload_check(
446	b: &CircuitBuilder,
447	jwt_payload: &ByteVec,
448	sub_byte_vec: &ByteVec,
449	aud_byte_vec: &ByteVec,
450	iss_byte_vec: &ByteVec,
451	base64_nonce: &[Wire; 6],
452) -> JwtClaims {
453	JwtClaims::new(
454		&b.subcircuit("jwt_claims_payload"),
455		jwt_payload.len_bytes,
456		jwt_payload.data.clone(),
457		vec![
458			Attribute {
459				name: "sub",
460				len_bytes: sub_byte_vec.len_bytes,
461				value: sub_byte_vec.data.clone(),
462			},
463			Attribute {
464				name: "aud",
465				len_bytes: aud_byte_vec.len_bytes,
466				value: aud_byte_vec.data.clone(),
467			},
468			Attribute {
469				name: "iss",
470				len_bytes: iss_byte_vec.len_bytes,
471				value: iss_byte_vec.data.clone(),
472			},
473			Attribute {
474				name: "nonce",
475				len_bytes: b.add_constant_64(43), /* Base64 encoded 32 bytes without padding = 43
476				                                   * chars */
477				value: base64_nonce.to_vec(),
478			},
479		],
480	)
481}
482
483pub struct ZkLoginExample {
484	zklogin: ZkLogin,
485}
486
487#[derive(Args, Debug, Clone)]
488pub struct Params {
489	/// Optional config for testing - if None, uses default
490	#[clap(skip)]
491	pub config: Option<Config>,
492}
493
494#[derive(Args, Debug, Clone)]
495pub struct Instance {
496	/// Subject claim value
497	#[arg(long, default_value = "1234567890")]
498	pub sub: String,
499
500	/// Audience claim value
501	#[arg(long, default_value = "4074087")]
502	pub aud: String,
503
504	/// Issuer claim value
505	#[arg(long, default_value = "google.com")]
506	pub iss: String,
507
508	/// Salt value for zkaddr computation
509	#[arg(long, default_value = "test_salt_value")]
510	pub salt: String,
511}
512
513struct JwtGenerationResult {
514	jwt: String,
515	zkaddr_hash: [u8; 32],
516	vk_u: [u8; 32],
517	zkaddr_preimage: Vec<u8>,
518	nonce_preimage: Vec<u8>,
519	jwt_key_pair: RS256KeyPair,
520}
521
522impl JwtGenerationResult {
523	fn generate(sub: &str, aud: &str, iss: &str, salt: &str, rng: &mut impl Rng) -> Result<Self> {
524		// Generate VK_u (verifier public key)
525		let mut vk_u = [0u8; 32];
526		rng.fill_bytes(&mut vk_u);
527
528		// Fixed values for nonce computation
529		let t_max = b"t_max";
530		let nonce_r = b"nonce_r";
531
532		// Calculate zkaddr = SHA256(concat(sub, aud, iss, salt))
533		let mut zkaddr_preimage = Vec::new();
534		zkaddr_preimage.extend_from_slice(sub.as_bytes());
535		zkaddr_preimage.extend_from_slice(aud.as_bytes());
536		zkaddr_preimage.extend_from_slice(iss.as_bytes());
537		zkaddr_preimage.extend_from_slice(salt.as_bytes());
538		let zkaddr_hash: [u8; 32] = Sha256::digest(&zkaddr_preimage).into();
539
540		// Calculate nonce = SHA256(concat(vk_u, t_max, nonce_r))
541		let mut nonce_preimage = Vec::new();
542		nonce_preimage.extend_from_slice(&vk_u);
543		nonce_preimage.extend_from_slice(t_max);
544		nonce_preimage.extend_from_slice(nonce_r);
545		let nonce_hash: [u8; 32] = Sha256::digest(&nonce_preimage).into();
546		let nonce_hash_base64 = BASE64_URL_SAFE_NO_PAD.encode(nonce_hash);
547
548		// Generate JWT key pair
549		let jwt_key_pair = RS256KeyPair::generate(2048).unwrap();
550
551		// Create and sign JWT
552		let claims = Claims::create(Duration::from_hours(2))
553			.with_issuer(iss)
554			.with_audience(aud)
555			.with_subject(sub)
556			.with_nonce(nonce_hash_base64);
557
558		let jwt = jwt_key_pair.sign(claims).unwrap();
559
560		Ok(Self {
561			jwt,
562			zkaddr_hash,
563			vk_u,
564			zkaddr_preimage,
565			nonce_preimage,
566			jwt_key_pair,
567		})
568	}
569}
570
571impl ExampleCircuit for ZkLoginExample {
572	type Params = Params;
573	type Instance = Instance;
574
575	fn build(params: Params, builder: &mut CircuitBuilder) -> Result<Self> {
576		let config = params.config.unwrap_or_default();
577		let zklogin = ZkLogin::new(builder, &config);
578
579		Ok(Self { zklogin })
580	}
581
582	fn populate_witness(&self, instance: Instance, w: &mut WitnessFiller<'_>) -> Result<()> {
583		let mut rng = StdRng::seed_from_u64(42);
584
585		// Generate JWT and related data
586		let JwtGenerationResult {
587			jwt,
588			zkaddr_hash,
589			vk_u,
590			zkaddr_preimage,
591			nonce_preimage,
592			jwt_key_pair,
593		} = JwtGenerationResult::generate(
594			&instance.sub,
595			&instance.aud,
596			&instance.iss,
597			&instance.salt,
598			&mut rng,
599		)?;
600
601		// Parse JWT components
602		let jwt_components = jwt.split(".").collect::<Vec<_>>();
603		let [header_base64, payload_base64, signature_base64] = jwt_components.as_slice() else {
604			anyhow::bail!("JWT should have format: header.payload.signature");
605		};
606
607		// Decode JWT components
608		let signature_bytes = BASE64_URL_SAFE_NO_PAD.decode(signature_base64)?;
609		let modulus_bytes = jwt_key_pair.public_key().to_components().n;
610		let header = BASE64_URL_SAFE_NO_PAD.decode(header_base64)?;
611		let payload = BASE64_URL_SAFE_NO_PAD.decode(payload_base64)?;
612
613		ensure!(
614			signature_bytes.len() == 256,
615			"RSA signature must be 256 bytes, got {}",
616			signature_bytes.len()
617		);
618
619		// Populate JWT components
620		self.zklogin
621			.populate_base64_jwt_header(w, header_base64.as_bytes());
622		self.zklogin
623			.populate_base64_jwt_payload(w, payload_base64.as_bytes());
624		self.zklogin
625			.populate_base64_jwt_signature(w, signature_base64.as_bytes());
626		self.zklogin.populate_jwt_header(w, &header);
627		self.zklogin.populate_jwt_header_attributes(w);
628		self.zklogin.populate_jwt_payload(w, &payload);
629		self.zklogin.populate_jwt_signature(w, &signature_bytes);
630
631		// Populate claim values
632		self.zklogin.populate_sub(w, instance.sub.as_bytes());
633		self.zklogin.populate_aud(w, instance.aud.as_bytes());
634		self.zklogin.populate_iss(w, instance.iss.as_bytes());
635		self.zklogin.populate_salt(w, instance.salt.as_bytes());
636
637		// Populate zkaddr
638		self.zklogin.populate_zkaddr(w, &zkaddr_hash);
639		self.zklogin.populate_zkaddr_preimage(w, &zkaddr_preimage);
640		self.zklogin.populate_vk_u(w, &vk_u);
641		self.zklogin.populate_t_max(w, b"t_max");
642		self.zklogin.populate_nonce_r(w, b"nonce_r");
643
644		// Populate nonce
645		let nonce_hash: [u8; 32] = Sha256::digest(&nonce_preimage).into();
646		let nonce_hash_base64 = BASE64_URL_SAFE_NO_PAD.encode(nonce_hash);
647		self.zklogin.populate_nonce(w, &nonce_hash);
648		self.zklogin.populate_nonce_preimage(w, &nonce_preimage);
649		self.zklogin
650			.populate_base64_jwt_payload_nonce(w, nonce_hash_base64.as_bytes());
651
652		// Populate JWS signature verification data
653		let message_str = format!("{header_base64}.{payload_base64}");
654		let message = message_str.as_bytes();
655		self.zklogin.populate_rsa_modulus(w, &modulus_bytes);
656		self.zklogin
657			.jwt_signature_verify
658			.populate_len_bytes(w, message.len());
659		self.zklogin
660			.jwt_signature_verify
661			.populate_message(w, message);
662		self.zklogin.jwt_signature_verify.populate_intermediates(
663			w,
664			&signature_bytes,
665			&modulus_bytes,
666		);
667
668		Ok(())
669	}
670}
671
672#[cfg(test)]
673mod tests {
674
675	use super::*;
676
677	fn run_zk_login_with_jwt_population(config: Config) {
678		let params = Params {
679			config: Some(config),
680		};
681		let instance = Instance {
682			sub: "1234567890".to_string(),
683			aud: "4074087".to_string(),
684			iss: "google.com".to_string(),
685			salt: "test_salt_value".to_string(),
686		};
687
688		let mut builder = CircuitBuilder::new();
689		let zklogin_example = ZkLoginExample::build(params, &mut builder).unwrap();
690		let circuit = builder.build();
691
692		let mut w = circuit.new_witness_filler();
693		zklogin_example.populate_witness(instance, &mut w).unwrap();
694
695		circuit.populate_wire_witness(&mut w).unwrap();
696		let cs = circuit.constraint_system();
697		cs.verify(&w.into_value_vec()).unwrap();
698	}
699
700	#[test]
701	fn test_zk_login_with_jwt_population() {
702		run_zk_login_with_jwt_population(Config::default());
703	}
704
705	#[test]
706	fn test_zk_login_with_jwt_population_weird_lengths() {
707		run_zk_login_with_jwt_population(Config {
708			max_len_json_jwt_header: 35,
709			max_len_json_jwt_payload: 61,
710			max_len_jwt_signature: 32,
711			max_len_jwt_sub: 8,
712			max_len_jwt_aud: 9,
713			max_len_jwt_iss: 10,
714			max_len_salt: 9,
715			max_len_nonce_r: 6,
716			max_len_t_max: 7,
717		});
718	}
719}