use super::Error;
use super::ole::{self, AliceState, BobMsg};
use super::otext::{ExtReceiver, ExtSender, ExtendMsg1};
use super::secp::Scalar;
pub const MUL_CHECK_FAILED: &str =
"ole: Mul-then-check failed — Bob's β differs across parallel runs";
pub struct CheckedBobMsg {
pub msg1: BobMsg,
pub msg2: BobMsg,
pub z: Scalar,
}
pub struct CheckedAliceState {
state1: AliceState,
state2: AliceState,
}
fn sub_sid(sid: &[u8], tag: u8) -> Vec<u8> {
let mut out = Vec::with_capacity(sid.len() + 2);
out.extend_from_slice(sid);
out.push(b'|');
out.push(tag);
out
}
pub fn checked_alice_step1(
sid: &[u8],
ext_receiver: &ExtReceiver,
alpha: &Scalar,
) -> Result<(ExtendMsg1, ExtendMsg1, CheckedAliceState), Error> {
let sid1 = sub_sid(sid, b'1');
let sid2 = sub_sid(sid, b'2');
let (msg1, state1) = ole::alice_step1(&sid1, ext_receiver, alpha)?;
let (msg2, state2) = ole::alice_step1(&sid2, ext_receiver, alpha)?;
Ok((msg1, msg2, CheckedAliceState { state1, state2 }))
}
pub fn checked_bob_step1(
sid: &[u8],
ext_sender: &ExtSender,
beta: &Scalar,
alice_msg1: &ExtendMsg1,
alice_msg2: &ExtendMsg1,
) -> Result<(CheckedBobMsg, Scalar), Error> {
let sid1 = sub_sid(sid, b'1');
let sid2 = sub_sid(sid, b'2');
let (msg1, u_b1) = ole::bob_step1(&sid1, ext_sender, beta, alice_msg1)?;
let (msg2, u_b2) = ole::bob_step1(&sid2, ext_sender, beta, alice_msg2)?;
let z = u_b1.sub(&u_b2);
Ok((CheckedBobMsg { msg1, msg2, z }, u_b1))
}
pub fn checked_alice_step2(
state: &CheckedAliceState,
bob_msg: &CheckedBobMsg,
) -> Result<Scalar, Error> {
let u_a1 = ole::alice_step2(&state.state1, &bob_msg.msg1)?;
let u_a2 = ole::alice_step2(&state.state2, &bob_msg.msg2)?;
let z_a = u_a1.sub(&u_a2);
let sum = z_a.add(&bob_msg.z);
if !bool::from(sum.is_zero()) {
return Err(Error::Validation(MUL_CHECK_FAILED.into()));
}
Ok(u_a1)
}
#[cfg(test)]
mod tests {
use super::super::baseot;
use super::super::otext;
use super::super::secp;
use super::*;
use purecrypto::rng::{OsRng, RngCore as _};
fn ot_setup() -> (ExtSender, ExtReceiver) {
let sid = b"ole-check-base";
let mut delta = [0u8; otext::DELTA_BYTES];
OsRng.fill_bytes(&mut delta);
let (bs, m1) = baseot::Sender::new(sid, otext::KAPPA, &mut OsRng);
let (br, m2) = baseot::Receiver::new(sid, otext::KAPPA, &delta, &m1, &mut OsRng).unwrap();
let (k0, k1) = bs.finalize(&m2).unwrap();
let chosen = br.finalize();
(
ExtSender::from_base(&delta, &chosen).unwrap(),
ExtReceiver::from_base(&k0, &k1).unwrap(),
)
}
#[test]
fn checked_shares_reconstruct_product() {
let (ext_sender, ext_receiver) = ot_setup();
let sid = b"checked-correctness";
let alpha = secp::random_scalar(&mut OsRng);
let beta = secp::random_scalar(&mut OsRng);
let (m1, m2, state) = checked_alice_step1(sid, &ext_receiver, &alpha).unwrap();
let (bmsg, u_b) = checked_bob_step1(sid, &ext_sender, &beta, &m1, &m2).unwrap();
let u_a = checked_alice_step2(&state, &bmsg).unwrap();
let lhs = u_a.add(&u_b);
let rhs = alpha.mul(&beta);
assert!(
bool::from(lhs.ct_eq(&rhs)),
"checked OLE must reconstruct α·β"
);
}
#[test]
fn checked_detects_inconsistent_beta() {
let (ext_sender, ext_receiver) = ot_setup();
let sid = b"checked-inconsistent";
let alpha = secp::random_scalar(&mut OsRng);
let beta1 = secp::random_scalar(&mut OsRng);
let beta2 = secp::random_scalar(&mut OsRng);
let (m1, m2, state) = checked_alice_step1(sid, &ext_receiver, &alpha).unwrap();
let sid1 = sub_sid(sid, b'1');
let sid2 = sub_sid(sid, b'2');
let (bmsg1, u_b1) = ole::bob_step1(&sid1, &ext_sender, &beta1, &m1).unwrap();
let (bmsg2, u_b2) = ole::bob_step1(&sid2, &ext_sender, &beta2, &m2).unwrap();
let z = u_b1.sub(&u_b2);
let bad = CheckedBobMsg {
msg1: bmsg1,
msg2: bmsg2,
z,
};
match checked_alice_step2(&state, &bad) {
Err(Error::Validation(m)) => assert_eq!(m, MUL_CHECK_FAILED),
Err(e) => panic!("expected MUL_CHECK_FAILED, got {e}"),
Ok(_) => panic!("checked_alice_step2 must reject inconsistent β"),
}
}
#[test]
fn checked_detects_tampered_z() {
let (ext_sender, ext_receiver) = ot_setup();
let sid = b"checked-tampered-z";
let alpha = secp::random_scalar(&mut OsRng);
let beta = secp::random_scalar(&mut OsRng);
let (m1, m2, state) = checked_alice_step1(sid, &ext_receiver, &alpha).unwrap();
let (mut bmsg, _u_b) = checked_bob_step1(sid, &ext_sender, &beta, &m1, &m2).unwrap();
bmsg.z = bmsg.z.add(&Scalar::ONE);
match checked_alice_step2(&state, &bmsg) {
Err(Error::Validation(m)) => assert_eq!(m, MUL_CHECK_FAILED),
Err(e) => panic!("expected MUL_CHECK_FAILED, got {e}"),
Ok(_) => panic!("checked_alice_step2 must reject a tampered Z"),
}
}
}