// Copyright (c) Facebook, Inc. and its affiliates. // // This source code is licensed under both the MIT license found in the // LICENSE-MIT file in the root directory of this source tree and the Apache // License, Version 2.0 found in the LICENSE-APACHE file in the root directory // of this source tree. use alloc::string::String; use alloc::vec; use alloc::vec::Vec; use core::ops::Add; use digest::core_api::BlockSizeUser; use digest::OutputSizeUser; use generic_array::typenum::{IsLess, IsLessOrEqual, Sum, U256}; use generic_array::ArrayLength; use serde_json::Value; use crate::tests::mock_rng::CycleRng; use crate::tests::parser::*; use crate::{ BlindedElement, CipherSuite, EvaluationElement, Group, OprfClient, OprfServer, PoprfClient, PoprfServer, PoprfServerBatchEvaluateFinishResult, PoprfServerBatchEvaluatePrepareResult, Proof, Result, VoprfClient, VoprfServer, VoprfServerBatchEvaluateFinishResult, }; #[derive(Debug)] struct VOPRFTestVectorParameters { seed: Vec, sksm: Vec, pksm: Vec, input: Vec>, info: Vec, key_info: Vec, blind: Vec>, blinded_element: Vec>, evaluation_element: Vec>, proof: Vec, proof_random_scalar: Vec, output: Vec>, } fn populate_test_vectors(values: &Value) -> VOPRFTestVectorParameters { VOPRFTestVectorParameters { seed: decode(values, "Seed"), sksm: decode(values, "skSm"), pksm: decode(values, "pkSm"), input: decode_vec(values, "Input"), info: decode(values, "Info"), key_info: decode(values, "KeyInfo"), blind: decode_vec(values, "Blind"), blinded_element: decode_vec(values, "BlindedElement"), evaluation_element: decode_vec(values, "EvaluationElement"), proof: decode(values, "Proof"), proof_random_scalar: decode(values, "ProofRandomScalar"), output: decode_vec(values, "Output"), } } fn decode(values: &Value, key: &str) -> Vec { values[key] .as_str() .and_then(|s| hex::decode(s).ok()) .unwrap_or_default() } fn decode_vec(values: &Value, key: &str) -> Vec> { let s = values[key].as_str().unwrap(); let res = match s.contains(',') { true => Some(s.split(',').map(|x| hex::decode(x).unwrap()).collect()), false => Some(vec![hex::decode(s).unwrap()]), }; res.unwrap() } macro_rules! json_to_test_vectors { ( $v:ident, $cs:expr, $mode:expr ) => { $v[$cs][$mode] .as_array() .into_iter() .flatten() .map(populate_test_vectors) .collect::>() }; } #[test] fn test_vectors() -> Result<()> { use p256::NistP256; let rfc: Value = serde_json::from_str(rfc_to_json(super::cfrg_vectors::VECTORS).as_str()) .expect("Could not parse json"); #[cfg(feature = "ristretto255")] { use crate::Ristretto255; let ristretto_oprf_tvs = json_to_test_vectors!( rfc, String::from("ristretto255, SHA-512"), String::from("OPRF") ); assert_ne!(ristretto_oprf_tvs.len(), 0); test_oprf_seed_to_key::(&ristretto_oprf_tvs)?; test_oprf_blind::(&ristretto_oprf_tvs)?; test_oprf_blind_evaluate::(&ristretto_oprf_tvs)?; test_oprf_finalize::(&ristretto_oprf_tvs)?; test_oprf_evaluate::(&ristretto_oprf_tvs)?; let ristretto_voprf_tvs = json_to_test_vectors!( rfc, String::from("ristretto255, SHA-512"), String::from("VOPRF") ); assert_ne!(ristretto_voprf_tvs.len(), 0); test_voprf_seed_to_key::(&ristretto_voprf_tvs)?; test_voprf_blind::(&ristretto_voprf_tvs)?; test_voprf_blind_evaluate::(&ristretto_voprf_tvs)?; test_voprf_finalize::(&ristretto_voprf_tvs)?; test_voprf_evaluate::(&ristretto_voprf_tvs)?; let ristretto_poprf_tvs = json_to_test_vectors!( rfc, String::from("ristretto255, SHA-512"), String::from("POPRF") ); assert_ne!(ristretto_poprf_tvs.len(), 0); test_poprf_seed_to_key::(&ristretto_poprf_tvs)?; test_poprf_blind::(&ristretto_poprf_tvs)?; test_poprf_blind_evaluate::(&ristretto_poprf_tvs)?; test_poprf_finalize::(&ristretto_poprf_tvs)?; test_poprf_evaluate::(&ristretto_poprf_tvs)?; } let p256_oprf_tvs = json_to_test_vectors!(rfc, String::from("P-256, SHA-256"), String::from("OPRF")); assert_ne!(p256_oprf_tvs.len(), 0); test_oprf_seed_to_key::(&p256_oprf_tvs)?; test_oprf_blind::(&p256_oprf_tvs)?; test_oprf_blind_evaluate::(&p256_oprf_tvs)?; test_oprf_finalize::(&p256_oprf_tvs)?; test_oprf_evaluate::(&p256_oprf_tvs)?; let p256_voprf_tvs = json_to_test_vectors!(rfc, String::from("P-256, SHA-256"), String::from("VOPRF")); assert_ne!(p256_voprf_tvs.len(), 0); test_voprf_seed_to_key::(&p256_voprf_tvs)?; test_voprf_blind::(&p256_voprf_tvs)?; test_voprf_blind_evaluate::(&p256_voprf_tvs)?; test_voprf_finalize::(&p256_voprf_tvs)?; test_voprf_evaluate::(&p256_voprf_tvs)?; let p256_poprf_tvs = json_to_test_vectors!(rfc, String::from("P-256, SHA-256"), String::from("POPRF")); assert_ne!(p256_poprf_tvs.len(), 0); test_poprf_seed_to_key::(&p256_poprf_tvs)?; test_poprf_blind::(&p256_poprf_tvs)?; test_poprf_blind_evaluate::(&p256_poprf_tvs)?; test_poprf_finalize::(&p256_poprf_tvs)?; test_poprf_evaluate::(&p256_poprf_tvs)?; Ok(()) } fn test_oprf_seed_to_key(tvs: &[VOPRFTestVectorParameters]) -> Result<()> where ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, { for parameters in tvs { let server = OprfServer::::new_from_seed(¶meters.seed, ¶meters.key_info)?; assert_eq!( ¶meters.sksm, &CS::Group::serialize_scalar(server.get_private_key()).to_vec() ); } Ok(()) } fn test_voprf_seed_to_key(tvs: &[VOPRFTestVectorParameters]) -> Result<()> where ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, { for parameters in tvs { let server = VoprfServer::::new_from_seed(¶meters.seed, ¶meters.key_info)?; assert_eq!( ¶meters.sksm, &CS::Group::serialize_scalar(server.get_private_key()).to_vec() ); assert_eq!( ¶meters.pksm, CS::Group::serialize_elem(server.get_public_key()).as_slice() ); } Ok(()) } fn test_poprf_seed_to_key(tvs: &[VOPRFTestVectorParameters]) -> Result<()> where ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, { for parameters in tvs { let server = PoprfServer::::new_from_seed(¶meters.seed, ¶meters.key_info)?; assert_eq!( ¶meters.sksm, &CS::Group::serialize_scalar(server.get_private_key()).to_vec() ); assert_eq!( ¶meters.pksm, CS::Group::serialize_elem(server.get_public_key()).as_slice() ); } Ok(()) } // Tests input -> blind, blinded_element fn test_oprf_blind(tvs: &[VOPRFTestVectorParameters]) -> Result<()> where ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, { for parameters in tvs { for i in 0..parameters.input.len() { let blind = CS::Group::deserialize_scalar(¶meters.blind[i])?; let client_result = OprfClient::::deterministic_blind_unchecked(¶meters.input[i], blind)?; assert_eq!( ¶meters.blind[i], &CS::Group::serialize_scalar(client_result.state.blind).to_vec() ); assert_eq!( parameters.blinded_element[i].as_slice(), client_result.message.serialize().as_slice(), ); } } Ok(()) } // Tests input -> blind, blinded_element fn test_voprf_blind(tvs: &[VOPRFTestVectorParameters]) -> Result<()> where ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, { for parameters in tvs { for i in 0..parameters.input.len() { let blind = CS::Group::deserialize_scalar(¶meters.blind[i])?; let client_blind_result = VoprfClient::::deterministic_blind_unchecked(¶meters.input[i], blind)?; assert_eq!( ¶meters.blind[i], &CS::Group::serialize_scalar(client_blind_result.state.get_blind()).to_vec() ); assert_eq!( parameters.blinded_element[i].as_slice(), client_blind_result.message.serialize().as_slice(), ); } } Ok(()) } // Tests input -> blind, blinded_element fn test_poprf_blind(tvs: &[VOPRFTestVectorParameters]) -> Result<()> where ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, { for parameters in tvs { for i in 0..parameters.input.len() { let blind = CS::Group::deserialize_scalar(¶meters.blind[i])?; let client_blind_result = PoprfClient::::deterministic_blind_unchecked(¶meters.input[i], blind)?; assert_eq!( ¶meters.blind[i], &CS::Group::serialize_scalar(client_blind_result.state.get_blind()).to_vec() ); assert_eq!( parameters.blinded_element[i].as_slice(), client_blind_result.message.serialize().as_slice(), ); } } Ok(()) } // Tests sksm, blinded_element -> evaluation_element fn test_oprf_blind_evaluate(tvs: &[VOPRFTestVectorParameters]) -> Result<()> where ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, { for parameters in tvs { for i in 0..parameters.input.len() { let server = OprfServer::::new_with_key(¶meters.sksm)?; let message = server.blind_evaluate(&BlindedElement::deserialize( ¶meters.blinded_element[i], )?); assert_eq!( ¶meters.evaluation_element[i], &message.serialize().as_slice() ); } } Ok(()) } fn test_voprf_blind_evaluate(tvs: &[VOPRFTestVectorParameters]) -> Result<()> where ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, ::ScalarLen: Add<::ScalarLen>, Sum<::ScalarLen, ::ScalarLen>: ArrayLength, { for parameters in tvs { let mut rng = CycleRng::new(parameters.proof_random_scalar.clone()); let server = VoprfServer::::new_with_key(¶meters.sksm)?; let mut blinded_elements = vec![]; for blinded_element_bytes in ¶meters.blinded_element { blinded_elements.push(BlindedElement::deserialize(blinded_element_bytes)?); } let prepared_evaluation_elements = server.batch_blind_evaluate_prepare(blinded_elements.iter()); let prepared_elements: Vec<_> = prepared_evaluation_elements.collect(); let VoprfServerBatchEvaluateFinishResult { messages, proof } = server .batch_blind_evaluate_finish(&mut rng, blinded_elements.iter(), &prepared_elements)?; let messages: Vec<_> = messages.collect(); for (parameter, message) in parameters.evaluation_element.iter().zip(messages) { assert_eq!(¶meter, &message.serialize().as_slice()); } assert_eq!(¶meters.proof, &proof.serialize().as_slice()); } Ok(()) } fn test_poprf_blind_evaluate(tvs: &[VOPRFTestVectorParameters]) -> Result<()> where ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, ::ScalarLen: Add<::ScalarLen>, Sum<::ScalarLen, ::ScalarLen>: ArrayLength, { for parameters in tvs { let mut rng = CycleRng::new(parameters.proof_random_scalar.clone()); let server = PoprfServer::::new_with_key(¶meters.sksm)?; let mut blinded_elements = vec![]; for blinded_element_bytes in ¶meters.blinded_element { blinded_elements.push(BlindedElement::deserialize(blinded_element_bytes)?); } let PoprfServerBatchEvaluatePrepareResult { prepared_evaluation_elements, prepared_tweak, } = server.batch_blind_evaluate_prepare(blinded_elements.iter(), Some(¶meters.info))?; let prepared_evaluation_elements: Vec<_> = prepared_evaluation_elements.collect(); let PoprfServerBatchEvaluateFinishResult { messages, proof } = PoprfServer::batch_blind_evaluate_finish::<_, _, Vec<_>>( &mut rng, blinded_elements.iter(), &prepared_evaluation_elements, &prepared_tweak, ) .unwrap(); let messages: Vec<_> = messages.collect(); for (parameter, message) in parameters.evaluation_element.iter().zip(messages) { assert_eq!(¶meter, &message.serialize().as_slice()); } assert_eq!(¶meters.proof, &proof.serialize().as_slice()); } Ok(()) } // Tests input, blind, evaluation_element -> output fn test_oprf_finalize(tvs: &[VOPRFTestVectorParameters]) -> Result<()> where ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, { for parameters in tvs { for i in 0..parameters.input.len() { let client = OprfClient::::from_blind(CS::Group::deserialize_scalar(¶meters.blind[i])?); let client_finalize_result = client.finalize( ¶meters.input[i], &EvaluationElement::deserialize(¶meters.evaluation_element[i])?, )?; assert_eq!(¶meters.output[i], &client_finalize_result.to_vec()); } } Ok(()) } fn test_voprf_finalize(tvs: &[VOPRFTestVectorParameters]) -> Result<()> where ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, { for parameters in tvs { let mut clients = vec![]; for i in 0..parameters.input.len() { let client = VoprfClient::::from_blind_and_element( CS::Group::deserialize_scalar(¶meters.blind[i])?, CS::Group::deserialize_elem(¶meters.blinded_element[i])?, ); clients.push(client.clone()); } let messages: Vec<_> = parameters .evaluation_element .iter() .map(|x| EvaluationElement::deserialize(x).unwrap()) .collect(); let batch_result = VoprfClient::batch_finalize( ¶meters.input, &clients, &messages, &Proof::deserialize(¶meters.proof)?, CS::Group::deserialize_elem(¶meters.pksm)?, )?; assert_eq!( parameters.output, batch_result .map(|arr| arr.map(|message| message.to_vec())) .collect::>>()? ); } Ok(()) } fn test_poprf_finalize(tvs: &[VOPRFTestVectorParameters]) -> Result<()> where ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, { for parameters in tvs { let mut clients = vec![]; for i in 0..parameters.input.len() { let blind = CS::Group::deserialize_scalar(¶meters.blind[i])?; let client_blind_result = PoprfClient::::deterministic_blind_unchecked(¶meters.input[i], blind)?; let client = client_blind_result.state; clients.push(client.clone()); } let messages: Vec<_> = parameters .evaluation_element .iter() .map(|x| EvaluationElement::deserialize(x).unwrap()) .collect(); let batch_result = PoprfClient::batch_finalize( parameters.input.iter().map(|input| input.as_slice()), &clients, &messages, &Proof::deserialize(¶meters.proof)?, CS::Group::deserialize_elem(¶meters.pksm)?, Some(¶meters.info), )?; let result: Vec> = batch_result.map(|arr| arr.unwrap().to_vec()).collect(); assert_eq!(parameters.output, result); } Ok(()) } // Tests input, sksm -> output fn test_oprf_evaluate(tvs: &[VOPRFTestVectorParameters]) -> Result<()> where ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, { for parameters in tvs { for i in 0..parameters.input.len() { let server = OprfServer::::new_with_key(¶meters.sksm)?; let server_evaluate_result = server.evaluate(¶meters.input[i])?; assert_eq!(¶meters.output[i], &server_evaluate_result.to_vec()); } } Ok(()) } fn test_voprf_evaluate(tvs: &[VOPRFTestVectorParameters]) -> Result<()> where ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, { for parameters in tvs { for i in 0..parameters.input.len() { let server = VoprfServer::::new_with_key(¶meters.sksm)?; let server_evaluate_result = server.evaluate(¶meters.input[i])?; assert_eq!(¶meters.output[i], &server_evaluate_result.to_vec()); } } Ok(()) } fn test_poprf_evaluate(tvs: &[VOPRFTestVectorParameters]) -> Result<()> where ::OutputSize: IsLess + IsLessOrEqual<::BlockSize>, { for parameters in tvs { for i in 0..parameters.input.len() { let server = PoprfServer::::new_with_key(¶meters.sksm)?; let server_evaluate_result = server.evaluate(¶meters.input[i], Some(¶meters.info))?; assert_eq!(¶meters.output[i], &server_evaluate_result.to_vec()); } } Ok(()) }