use serde::{Deserialize, Serialize};
use crate::WireError;
pub const PAYLOAD_TYPE_OPERATOR_REKEY: u8 = 0x31;
#[cfg(feature = "handshake")]
const OPERATOR_REKEY_AEAD_LABEL: &[u8] = b"xenia/operator/rekey/aead";
const OPERATOR_REKEY_EPOCH_CONTEXT_SCHEMA: &str = "xenia-operator-rekey-epoch-context-v1";
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum OperatorRekeyReason {
Interval,
Manual,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
struct OperatorRekeyEpochContext {
schema: String,
key_epoch: u64,
base_transcript_hash: [u8; 32],
previous_epoch_hash: [u8; 32],
reason: OperatorRekeyReason,
}
impl OperatorRekeyEpochContext {
fn new(
key_epoch: u64,
base_transcript_hash: [u8; 32],
previous_epoch_hash: [u8; 32],
reason: OperatorRekeyReason,
) -> Self {
Self {
schema: OPERATOR_REKEY_EPOCH_CONTEXT_SCHEMA.to_string(),
key_epoch,
base_transcript_hash,
previous_epoch_hash,
reason,
}
}
fn epoch_hash(&self) -> Result<[u8; 32], WireError> {
let bytes = bincode::serialize(self).map_err(WireError::encode)?;
Ok(*blake3::hash(&bytes).as_bytes())
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum OperatorRekeyMessage {
Proposal {
key_epoch: u64,
base_transcript_hash: [u8; 32],
previous_epoch_hash: [u8; 32],
reason: OperatorRekeyReason,
epoch_hash: [u8; 32],
},
Ack {
key_epoch: u64,
epoch_hash: [u8; 32],
},
}
impl OperatorRekeyMessage {
pub fn encode(&self) -> Result<Vec<u8>, WireError> {
bincode::serialize(self).map_err(WireError::encode)
}
pub fn decode(bytes: &[u8]) -> Result<Self, WireError> {
bincode::deserialize(bytes).map_err(WireError::decode)
}
}
pub fn propose(
key_epoch: u64,
base_transcript_hash: [u8; 32],
previous_epoch_hash: [u8; 32],
reason: OperatorRekeyReason,
) -> Result<OperatorRekeyMessage, WireError> {
let epoch_hash = OperatorRekeyEpochContext::new(
key_epoch,
base_transcript_hash,
previous_epoch_hash,
reason,
)
.epoch_hash()?;
Ok(OperatorRekeyMessage::Proposal {
key_epoch,
base_transcript_hash,
previous_epoch_hash,
reason,
epoch_hash,
})
}
pub fn verify_proposal_epoch_hash(
key_epoch: u64,
base_transcript_hash: [u8; 32],
previous_epoch_hash: [u8; 32],
reason: OperatorRekeyReason,
claimed_epoch_hash: [u8; 32],
) -> Result<[u8; 32], WireError> {
let computed = OperatorRekeyEpochContext::new(
key_epoch,
base_transcript_hash,
previous_epoch_hash,
reason,
)
.epoch_hash()?;
if computed != claimed_epoch_hash {
return Err(WireError::decode(
"operator rekey proposal epoch_hash does not match its own context",
));
}
Ok(computed)
}
#[cfg(feature = "handshake")]
pub fn derive_operator_rekey_key(rekey_root: &[u8; 32], epoch_hash: &[u8; 32]) -> [u8; 32] {
crate::handshake::derive_labeled_session_key(rekey_root, epoch_hash, OPERATOR_REKEY_AEAD_LABEL)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn proposal_round_trips_through_bincode() {
let msg = propose(3, [0x11; 32], [0x22; 32], OperatorRekeyReason::Interval).unwrap();
let bytes = msg.encode().unwrap();
assert_eq!(OperatorRekeyMessage::decode(&bytes).unwrap(), msg);
}
#[test]
fn ack_round_trips_through_bincode() {
let msg = OperatorRekeyMessage::Ack {
key_epoch: 3,
epoch_hash: [0x33; 32],
};
let bytes = msg.encode().unwrap();
assert_eq!(OperatorRekeyMessage::decode(&bytes).unwrap(), msg);
}
#[test]
fn decode_rejects_garbage() {
assert!(OperatorRekeyMessage::decode(&[0xff; 4]).is_err());
}
#[test]
fn payload_type_is_distinct_from_application_min() {
assert_ne!(
PAYLOAD_TYPE_OPERATOR_REKEY,
crate::payload_types::PAYLOAD_TYPE_APPLICATION_MIN
);
}
#[test]
fn propose_then_verify_round_trips() {
let OperatorRekeyMessage::Proposal {
key_epoch,
base_transcript_hash,
previous_epoch_hash,
reason,
epoch_hash,
} = propose(1, [0xaa; 32], [0xbb; 32], OperatorRekeyReason::Manual).unwrap()
else {
unreachable!()
};
let verified = verify_proposal_epoch_hash(
key_epoch,
base_transcript_hash,
previous_epoch_hash,
reason,
epoch_hash,
)
.unwrap();
assert_eq!(verified, epoch_hash);
}
#[test]
fn verify_rejects_a_tampered_epoch_hash() {
let OperatorRekeyMessage::Proposal {
key_epoch,
base_transcript_hash,
previous_epoch_hash,
reason,
..
} = propose(1, [0xaa; 32], [0xbb; 32], OperatorRekeyReason::Manual).unwrap()
else {
unreachable!()
};
let wrong_hash = [0xff; 32];
assert!(
verify_proposal_epoch_hash(
key_epoch,
base_transcript_hash,
previous_epoch_hash,
reason,
wrong_hash,
)
.is_err()
);
}
#[test]
fn verify_rejects_a_desynced_epoch_number() {
let msg = propose(1, [0xaa; 32], [0xbb; 32], OperatorRekeyReason::Manual).unwrap();
let OperatorRekeyMessage::Proposal {
base_transcript_hash,
previous_epoch_hash,
reason,
epoch_hash,
..
} = msg
else {
unreachable!()
};
assert!(
verify_proposal_epoch_hash(
2,
base_transcript_hash,
previous_epoch_hash,
reason,
epoch_hash,
)
.is_err()
);
}
#[test]
#[cfg(feature = "handshake")]
fn different_epochs_derive_different_keys() {
let root = [0x42; 32];
let hash_a = propose(1, [0; 32], [0; 32], OperatorRekeyReason::Interval)
.unwrap()
.encode()
.unwrap();
let hash_b = propose(2, [0; 32], [0; 32], OperatorRekeyReason::Interval)
.unwrap()
.encode()
.unwrap();
assert_ne!(hash_a, hash_b);
let OperatorRekeyMessage::Proposal {
epoch_hash: eh_a, ..
} = OperatorRekeyMessage::decode(&hash_a).unwrap()
else {
unreachable!()
};
let OperatorRekeyMessage::Proposal {
epoch_hash: eh_b, ..
} = OperatorRekeyMessage::decode(&hash_b).unwrap()
else {
unreachable!()
};
assert_ne!(
derive_operator_rekey_key(&root, &eh_a),
derive_operator_rekey_key(&root, &eh_b)
);
}
}