use super::*;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PeerRpcValidationConfig {
pub local_tenant_id: TenantId,
pub local_cluster_id: ClusterId,
pub local_core_id: CoreId,
pub max_payload_bytes: usize,
pub nonce_window_ms: u64,
}
#[derive(Debug, Clone)]
pub struct PeerRpcValidator {
config: PeerRpcValidationConfig,
local_protocol_version: ProtocolVersion,
nonce_store: Arc<dyn crate::PeerNonceStore>,
}
impl PeerRpcValidator {
pub fn new(config: PeerRpcValidationConfig) -> Self {
Self {
config,
local_protocol_version: ProtocolVersion::default(),
nonce_store: Arc::new(crate::InMemoryPeerNonceStore::default()),
}
}
pub fn with_protocol_version(mut self, protocol_version: ProtocolVersion) -> Self {
self.local_protocol_version = protocol_version;
self
}
pub fn with_nonce_store(mut self, nonce_store: Arc<dyn crate::PeerNonceStore>) -> Self {
self.nonce_store = nonce_store;
self
}
pub(crate) fn max_envelope_bytes(&self) -> usize {
self.config
.max_payload_bytes
.saturating_mul(4)
.saturating_add(MAX_ENVELOPE_OVERHEAD_BYTES)
}
pub fn validate(&self, envelope: &PeerRpcEnvelope, now_ms: u64) -> Result<(), PeerRpcError> {
validate_envelope_identifiers(envelope)?;
if envelope.payload.len() > self.config.max_payload_bytes {
return Err(PeerRpcError::PayloadTooLarge);
}
if envelope.tenant_id != self.config.local_tenant_id {
return Err(PeerRpcError::TenantMismatch);
}
if envelope.cluster_id != self.config.local_cluster_id {
return Err(PeerRpcError::ClusterMismatch);
}
if envelope.target_core_id != self.config.local_core_id {
return Err(PeerRpcError::TargetMismatch);
}
if !self
.local_protocol_version
.is_compatible_with(envelope.protocol_version)
{
return Err(PeerRpcError::ProtocolMismatch);
}
let window_ms = self.config.nonce_window_ms.max(1);
if envelope.timestamp_ms >= envelope.expires_at_ms
|| envelope.expires_at_ms <= now_ms
|| envelope.timestamp_ms > now_ms.saturating_add(window_ms)
|| now_ms > envelope.timestamp_ms.saturating_add(window_ms)
{
return Err(PeerRpcError::Expired);
}
if envelope.body_hash != payload_hash(&envelope.payload) {
return Err(PeerRpcError::InvalidBodyHash);
}
if let Some(trace) = &envelope.trace {
if trace.trace_id != envelope.trace_id
|| trace.tenant_id != envelope.tenant_id
|| trace.current_core_id != envelope.source_core_id
{
return Err(PeerRpcError::InvalidEnvelope(
"trace_context_mismatch".to_string(),
));
}
}
let nonce_expires_at_ms = envelope.expires_at_ms.min(now_ms.saturating_add(window_ms));
self.nonce_store
.check_and_record(&envelope.nonce, nonce_expires_at_ms, now_ms)?;
Ok(())
}
}
fn validate_envelope_identifiers(envelope: &PeerRpcEnvelope) -> Result<(), PeerRpcError> {
for (kind, value) in [
("PeerRequestId", envelope.request_id.as_str()),
("TraceId", envelope.trace_id.as_str()),
("PeerNonce", envelope.nonce.as_str()),
] {
validate_identifier(kind, value)
.map_err(|_| PeerRpcError::InvalidEnvelope("invalid_identifier".to_string()))?;
}
if let Some(idempotency_key) = &envelope.idempotency_key {
validate_identifier("IdempotencyKey", idempotency_key)
.map_err(|_| PeerRpcError::InvalidEnvelope("invalid_idempotency_key".to_string()))?;
}
envelope
.source_core_id
.validate()
.and_then(|_| envelope.target_core_id.validate())
.and_then(|_| envelope.tenant_id.validate())
.and_then(|_| envelope.cluster_id.validate())
.and_then(|_| envelope.capability.validate())
.map_err(|_| PeerRpcError::InvalidEnvelope("invalid_identifier".to_string()))
}
pub fn payload_hash(payload: &[u8]) -> String {
let digest = Sha256::digest(payload);
hex_encode(&digest)
}
pub fn envelope_signing_hash(envelope: &PeerRpcEnvelope) -> String {
let mut hasher = Sha256::new();
hash_field(&mut hasher, envelope.request_id.as_bytes());
hash_field(&mut hasher, envelope.trace_id.as_bytes());
hasher.update(envelope.protocol_version.as_u16().to_be_bytes());
hash_field(&mut hasher, envelope.source_core_id.as_str().as_bytes());
hash_field(&mut hasher, envelope.target_core_id.as_str().as_bytes());
hash_field(&mut hasher, envelope.tenant_id.as_str().as_bytes());
hash_field(&mut hasher, envelope.cluster_id.as_str().as_bytes());
hasher.update(envelope.timestamp_ms.to_be_bytes());
hasher.update(envelope.expires_at_ms.to_be_bytes());
hash_field(&mut hasher, envelope.nonce.as_bytes());
hash_field(&mut hasher, envelope.capability.as_str().as_bytes());
hash_field(&mut hasher, envelope.body_hash.as_bytes());
hash_optional_field(&mut hasher, envelope.idempotency_key.as_deref());
if let Some(trace) = &envelope.trace {
hasher.update([1]);
hash_field(&mut hasher, trace.trace_id.as_bytes());
hash_field(&mut hasher, trace.span_id.as_bytes());
hash_optional_field(&mut hasher, trace.parent_span_id.as_deref());
hash_field(&mut hasher, trace.originating_core_id.as_str().as_bytes());
hash_field(&mut hasher, trace.current_core_id.as_str().as_bytes());
hash_field(&mut hasher, trace.tenant_id.as_str().as_bytes());
hash_optional_field(&mut hasher, trace.command_id.as_deref());
} else {
hasher.update([0]);
}
hex_encode(&hasher.finalize())
}
fn hash_field(hasher: &mut Sha256, value: &[u8]) {
hasher.update((value.len() as u64).to_be_bytes());
hasher.update(value);
}
fn hash_optional_field(hasher: &mut Sha256, value: Option<&str>) {
if let Some(value) = value {
hasher.update([1]);
hash_field(hasher, value.as_bytes());
} else {
hasher.update([0]);
}
}
pub fn route_for_query() -> &'static str {
PEER_QUERY_PATH
}
pub fn route_for_command() -> &'static str {
PEER_COMMAND_PATH
}
fn hex_encode(bytes: &[u8]) -> String {
const HEX: &[u8; 16] = b"0123456789abcdef";
let mut out = String::with_capacity(bytes.len() * 2);
for byte in bytes {
out.push(HEX[(byte >> 4) as usize] as char);
out.push(HEX[(byte & 0x0f) as usize] as char);
}
out
}