From 6b22064863b317a02c4cfd2547721148678f2c39 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Fran=C3=A7ois=20Garillot?= Date: Wed, 9 Dec 2020 12:49:59 -0800 Subject: [PATCH] Fold a few panics Removes a few panics we don't need by folding them in the Error case of their enclosing Result return. --- src/group.rs | 5 +-- src/keypair.rs | 9 ++++- src/slow_hash.rs | 3 +- src/tests/opaque_ke_test.rs | 66 ++++++++++++++----------------------- 4 files changed, 35 insertions(+), 48 deletions(-) diff --git a/src/group.rs b/src/group.rs index 78e3f61..f01aabf 100644 --- a/src/group.rs +++ b/src/group.rs @@ -196,10 +196,7 @@ mod tests { ]; fn deserialize_point(pt: &[u8]) -> Result { - let bytes: [u8; 32] = (&pt[..32]) - .try_into() - .expect("Slice pattern invariant broken"); - + let bytes: [u8; 32] = (&pt[..32]).try_into()?; curve25519_dalek::edwards::CompressedEdwardsY(bytes) .decompress() .ok_or_else(|| anyhow!("Point decompression failed!")) diff --git a/src/keypair.rs b/src/keypair.rs index 711175b..2014601 100644 --- a/src/keypair.rs +++ b/src/keypair.rs @@ -160,7 +160,14 @@ impl KeyPair for X25519KeyPair { } fn check_public_key(key: Self::Repr) -> Result { - let key_bytes: [u8; 32] = (&key[..]).try_into().expect("Key invariant broken"); + let key_bytes: [u8; 32] = + (&key[..]) + .try_into() + .map_err(|_| InternalPakeError::SizeError { + name: "key", + len: 32, + actual_len: key.len(), + })?; let point = ::curve25519_dalek::montgomery::MontgomeryPoint(key_bytes) .to_edwards(1) .ok_or(InternalPakeError::PointError)?; diff --git a/src/slow_hash.rs b/src/slow_hash.rs index 8c94acf..ec23e97 100644 --- a/src/slow_hash.rs +++ b/src/slow_hash.rs @@ -35,7 +35,8 @@ impl SlowHash for scrypt::ScryptParams { fn hash( input: GenericArray::OutputSize>, ) -> Result, InternalPakeError> { - let params = scrypt::ScryptParams::new(15, 8, 1).unwrap(); + let params = + scrypt::ScryptParams::new(15, 8, 1).map_err(|_| InternalPakeError::SlowHashError)?; let mut output = vec![0u8; ::OutputSize::to_usize()]; scrypt::scrypt(&input, &[], ¶ms, &mut output) .map_err(|_| InternalPakeError::SlowHashError)?; diff --git a/src/tests/opaque_ke_test.rs b/src/tests/opaque_ke_test.rs index 3272d3b..19bf305 100644 --- a/src/tests/opaque_ke_test.rs +++ b/src/tests/opaque_ke_test.rs @@ -413,7 +413,7 @@ fn postprocess_blinding_factor(_: G::Scalar) -> G::Scalar { } #[test] -fn test_r1() -> Result<(), PakeError> { +fn test_r1() -> Result<(), ProtocolError> { let parameters = populate_test_vectors(&serde_json::from_str(TEST_VECTOR).unwrap()); let mut rng = OsRng; let (r1, client_registration) = ClientRegistration::::start( @@ -421,8 +421,7 @@ fn test_r1() -> Result<(), PakeError> { ClientRegistrationStartParameters::WithIdentifiers(parameters.id_u, parameters.id_s), &mut rng, postprocess_blinding_factor::<::Group>, - ) - .unwrap(); + )?; assert_eq!(hex::encode(¶meters.r1), hex::encode(r1.serialize())); assert_eq!( hex::encode(¶meters.client_registration_state), @@ -432,15 +431,14 @@ fn test_r1() -> Result<(), PakeError> { } #[test] -fn test_r2() -> Result<(), PakeError> { +fn test_r2() -> Result<(), ProtocolError> { let parameters = populate_test_vectors(&serde_json::from_str(TEST_VECTOR).unwrap()); let mut oprf_key_rng = CycleRng::new(parameters.oprf_key); let (r2, server_registration) = ServerRegistration::::start( RegisterFirstMessage::deserialize(¶meters.r1[..]).unwrap(), &Key::try_from(¶meters.server_s_pk[..]).unwrap(), &mut oprf_key_rng, - ) - .unwrap(); + )?; assert_eq!(hex::encode(parameters.r2), hex::encode(r2.serialize())); assert_eq!( hex::encode(¶meters.server_registration_state), @@ -450,7 +448,7 @@ fn test_r2() -> Result<(), PakeError> { } #[test] -fn test_r3() -> Result<(), PakeError> { +fn test_r3() -> Result<(), ProtocolError> { let parameters = populate_test_vectors(&serde_json::from_str(TEST_VECTOR).unwrap()); let client_s_sk_and_nonce: Vec = @@ -458,14 +456,11 @@ fn test_r3() -> Result<(), PakeError> { let mut finish_registration_rng = CycleRng::new(client_s_sk_and_nonce); let (r3, export_key_registration) = ClientRegistration::::try_from( ¶meters.client_registration_state[..], - ) - .unwrap() + )? .finish( RegisterSecondMessage::deserialize(¶meters.r2[..]).unwrap(), &mut finish_registration_rng, - ) - .unwrap(); - + )?; assert_eq!(hex::encode(parameters.r3), hex::encode(r3.serialize())); assert_eq!( hex::encode(parameters.export_key), @@ -476,17 +471,14 @@ fn test_r3() -> Result<(), PakeError> { } #[test] -fn test_password_file() -> Result<(), PakeError> { +fn test_password_file() -> Result<(), ProtocolError> { let parameters = populate_test_vectors(&serde_json::from_str(TEST_VECTOR).unwrap()); let server_registration = ServerRegistration::::try_from( ¶meters.server_registration_state[..], - ) - .unwrap(); + )?; let password_file = server_registration - .finish(RegisterThirdMessage::deserialize(¶meters.r3[..]).unwrap()) - .unwrap(); - + .finish(RegisterThirdMessage::deserialize(¶meters.r3[..]).unwrap())?; assert_eq!( hex::encode(parameters.password_file), hex::encode(password_file.to_bytes()) @@ -495,7 +487,7 @@ fn test_password_file() -> Result<(), PakeError> { } #[test] -fn test_l1() -> Result<(), PakeError> { +fn test_l1() -> Result<(), ProtocolError> { let parameters = populate_test_vectors(&serde_json::from_str(TEST_VECTOR).unwrap()); let client_login_start = [ @@ -514,8 +506,7 @@ fn test_l1() -> Result<(), PakeError> { parameters.id_s, ), postprocess_blinding_factor::<::Group>, - ) - .unwrap(); + )?; assert_eq!( hex::encode(¶meters.l1), hex::encode(client_login_start_result.credential_request.serialize()) @@ -528,7 +519,7 @@ fn test_l1() -> Result<(), PakeError> { } #[test] -fn test_l2() -> Result<(), PakeError> { +fn test_l2() -> Result<(), ProtocolError> { let parameters = populate_test_vectors(&serde_json::from_str(TEST_VECTOR).unwrap()); let mut server_e_sk_rng = CycleRng::new(parameters.server_e_sk); @@ -538,9 +529,7 @@ fn test_l2() -> Result<(), PakeError> { LoginFirstMessage::::deserialize(¶meters.l1[..]).unwrap(), &mut server_e_sk_rng, ServerLoginStartParameters::WithInfo(parameters.info2.to_vec(), parameters.einfo2.to_vec()), - ) - .unwrap(); - + )?; assert_eq!( hex::encode(¶meters.info1), hex::encode(server_login_start_result.plain_info), @@ -561,21 +550,17 @@ fn test_l2() -> Result<(), PakeError> { } #[test] -fn test_l3() -> Result<(), PakeError> { +fn test_l3() -> Result<(), ProtocolError> { let parameters = populate_test_vectors(&serde_json::from_str(TEST_VECTOR).unwrap()); let client_login_finish_result = - ClientLogin::::try_from(¶meters.client_login_state[..]) - .unwrap() - .finish( - LoginSecondMessage::::deserialize(¶meters.l2[..]).unwrap(), - ClientLoginFinishParameters::WithInfo( - parameters.info3.to_vec(), - parameters.einfo3.to_vec(), - ), - ) - .unwrap(); - + ClientLogin::::try_from(¶meters.client_login_state[..])?.finish( + LoginSecondMessage::::deserialize(¶meters.l2[..]).unwrap(), + ClientLoginFinishParameters::WithInfo( + parameters.info3.to_vec(), + parameters.einfo3.to_vec(), + ), + )?; assert_eq!( hex::encode(¶meters.info2), hex::encode(&client_login_finish_result.plain_info) @@ -610,11 +595,8 @@ fn test_server_login_finish() -> Result<(), ProtocolError> { let parameters = populate_test_vectors(&serde_json::from_str(TEST_VECTOR).unwrap()); let server_login_result = - ServerLogin::::try_from(¶meters.server_login_state[..]) - .unwrap() - .finish(LoginThirdMessage::try_from(¶meters.l3[..])?) - .unwrap(); - + ServerLogin::::try_from(¶meters.server_login_state[..])? + .finish(LoginThirdMessage::try_from(¶meters.l3[..])?)?; assert_eq!( hex::encode(parameters.info3), hex::encode(server_login_result.plain_info)