Files
opaque-vx/src/oprf.rs
T

163 lines
5.5 KiB
Rust
Raw Normal View History

2020-06-05 09:35:14 -07:00
// Copyright (c) Facebook, Inc. and its affiliates.
//
// This source code is licensed under the MIT license found in the
// LICENSE file in the root directory of this source tree.
2020-07-27 15:25:04 -07:00
use crate::{
errors::InternalPakeError, group::Group, hash::Hash, map_to_curve::GroupWithMapToCurve,
};
use digest::Digest;
use generic_array::GenericArray;
2020-06-05 09:35:14 -07:00
use hkdf::Hkdf;
use rand_core::{CryptoRng, RngCore};
2020-11-02 13:51:43 -08:00
/// Used to store the OPRF input and blinding factor
pub struct Token<Grp: Group> {
pub(crate) data: Vec<u8>,
pub(crate) blind: Grp::Scalar,
2020-06-05 09:35:14 -07:00
}
2020-11-02 13:51:43 -08:00
static STR_VOPRF: &[u8] = b"VOPRF05";
2020-06-05 09:35:14 -07:00
/// Computes the first step for the multiplicative blinding version of DH-OPRF. This
/// message is sent from the client (who holds the input) to the server (who holds the OPRF key).
/// The client can also pass in an optional "pepper" string to be mixed in with the input through
/// an HKDF computation.
2020-11-02 13:51:43 -08:00
pub(crate) fn blind_with_postprocessing<R: RngCore + CryptoRng, G: GroupWithMapToCurve>(
2020-06-05 09:35:14 -07:00
input: &[u8],
blinding_factor_rng: &mut R,
2020-11-02 13:51:43 -08:00
postprocess: fn(G::Scalar) -> G::Scalar,
) -> Result<(Token<G>, G), InternalPakeError> {
let mapped_point = G::map_to_curve(input, Some(STR_VOPRF)); // TODO: add contextString from RFC
2020-06-05 09:35:14 -07:00
let blinding_factor = G::random_scalar(blinding_factor_rng);
2020-11-02 13:51:43 -08:00
let blind = postprocess(blinding_factor);
let blind_token = mapped_point * &blind;
Ok((
Token {
data: input.to_vec(),
blind,
},
blind_token,
))
2020-06-05 09:35:14 -07:00
}
/// Computes the second step for the multiplicative blinding version of DH-OPRF. This
/// message is sent from the server (who holds the OPRF key) to the client.
2020-11-02 13:51:43 -08:00
pub(crate) fn evaluate<G: Group>(point: G, oprf_key: &G::Scalar) -> Result<G, InternalPakeError> {
2020-06-05 09:35:14 -07:00
Ok(point * oprf_key)
}
/// Computes the third step for the multiplicative blinding version of DH-OPRF, in which
/// the client unblinds the server's message.
2020-11-02 13:51:43 -08:00
pub(crate) fn unblind_and_finalize<G: Group, H: Hash>(
token: &Token<G>,
2020-06-05 09:35:14 -07:00
point: G,
2020-07-27 15:25:04 -07:00
) -> Result<GenericArray<u8, <H as Digest>::OutputSize>, InternalPakeError> {
2020-11-02 13:51:43 -08:00
let unblinded = point * &G::scalar_invert(&token.blind);
let ikm: Vec<u8> = [&unblinded.to_arr()[..], &token.data].concat();
// TODO: implement proper finalizing code here
2020-07-27 15:25:04 -07:00
let (prk, _) = Hkdf::<H>::extract(None, &ikm);
2020-06-05 09:35:14 -07:00
Ok(prk)
}
// Benchmarking shims
#[cfg(feature = "bench")]
#[inline]
2020-11-02 13:51:43 -08:00
pub fn blind_shim<R: RngCore + CryptoRng, G: GroupWithMapToCurve>(
input: &[u8],
blinding_factor_rng: &mut R,
2020-11-02 13:51:43 -08:00
) -> Result<(Token<G>, G), InternalPakeError> {
blind_with_postprocessing(input, blinding_factor_rng, std::convert::identity)
}
#[cfg(feature = "bench")]
#[inline]
2020-11-02 13:51:43 -08:00
pub fn evaluate_shim<G: Group>(point: G, oprf_key: &G::Scalar) -> Result<G, InternalPakeError> {
evaluate(point, oprf_key)
}
#[cfg(feature = "bench")]
#[inline]
2020-11-02 13:51:43 -08:00
pub fn unblind_and_finalize_shim<G: Group, H: Hash>(
token: &Token<G>,
point: G,
2020-09-03 15:54:46 -04:00
) -> Result<GenericArray<u8, <H as Digest>::OutputSize>, InternalPakeError> {
2020-11-02 13:51:43 -08:00
unblind_and_finalize::<G, H>(token, point)
}
2020-06-05 09:35:14 -07:00
// Tests
// =====
#[cfg(test)]
mod tests {
use super::*;
use crate::group::Group;
use curve25519_dalek::ristretto::RistrettoPoint;
use generic_array::{arr, GenericArray};
2020-06-05 09:35:14 -07:00
use hkdf::Hkdf;
use rand_core::OsRng;
2020-07-27 15:25:04 -07:00
use sha2::{Sha256, Sha512};
2020-06-05 09:35:14 -07:00
fn prf(
input: &[u8],
oprf_key: &[u8; 32],
) -> GenericArray<u8, <RistrettoPoint as Group>::ElemLen> {
2020-11-02 13:51:43 -08:00
let (hashed_input, _) = Hkdf::<Sha512>::extract(Some(STR_VOPRF), &input);
let point = RistrettoPoint::hash_to_curve(GenericArray::from_slice(&hashed_input));
2020-06-05 09:35:14 -07:00
let scalar =
RistrettoPoint::from_scalar_slice(GenericArray::from_slice(&oprf_key[..])).unwrap();
let res = point * scalar;
let ikm: Vec<u8> = [&res.to_arr()[..], &input].concat();
2020-06-05 09:35:14 -07:00
let (prk, _) = Hkdf::<Sha256>::extract(None, &ikm);
prk
}
#[test]
fn oprf_retrieval() -> Result<(), InternalPakeError> {
let input = b"hunter2";
let mut rng = OsRng;
2020-11-02 13:51:43 -08:00
let (token, alpha) = blind_with_postprocessing::<_, RistrettoPoint>(
&input[..],
&mut rng,
std::convert::identity,
)?;
let oprf_key_bytes = arr![
2020-06-05 09:35:14 -07:00
u8; 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23,
24, 25, 26, 27, 28, 29, 30, 31, 32,
];
2020-11-02 13:51:43 -08:00
let oprf_key = RistrettoPoint::from_scalar_slice(&oprf_key_bytes)?;
let beta = evaluate::<RistrettoPoint>(alpha, &oprf_key)?;
let res = unblind_and_finalize::<RistrettoPoint, sha2::Sha256>(&token, beta)?;
let res2 = prf(&input[..], &oprf_key.as_bytes());
2020-06-05 09:35:14 -07:00
assert_eq!(res, res2);
Ok(())
}
#[test]
fn oprf_inversion_unsalted() {
let mut rng = OsRng;
let mut input = vec![0u8; 64];
rng.fill_bytes(&mut input);
2020-11-02 13:51:43 -08:00
let (token, alpha) = blind_with_postprocessing::<_, RistrettoPoint>(
&input,
&mut rng,
std::convert::identity,
)
.unwrap();
let res = unblind_and_finalize::<RistrettoPoint, sha2::Sha256>(&token, alpha).unwrap();
let (hashed_input, _) = Hkdf::<Sha512>::extract(Some(STR_VOPRF), &input);
let mut bits = [0u8; 64];
2020-07-27 15:25:04 -07:00
bits.copy_from_slice(&hashed_input);
let point = RistrettoPoint::from_uniform_bytes(&bits);
2020-06-05 09:35:14 -07:00
let mut ikm: Vec<u8> = Vec::new();
ikm.extend_from_slice(&point.to_arr());
2020-06-05 09:35:14 -07:00
ikm.extend_from_slice(&input);
let (prk, _) = Hkdf::<Sha256>::extract(None, &ikm);
assert_eq!(res, prk);
}
}