use std::collections::{BTreeMap, BTreeSet};
use hidpp::channel::RequestOutcome;
use openlogi_fixture::{
CassetteExchange, FIXTURE_SCHEMA_VERSION, HidCassette, ReportSupport, SyntheticIdentityKind,
classify_synthetic_identity_bytes,
};
use thiserror::Error;
use super::{RecordedChannel, RecordedChannelOpenOutcome, RecordedRequest, RecordedRequestFact};
mod sanitizer;
use sanitizer::{ProtocolSanitizer, unassociated_rejection};
#[cfg(test)]
mod tests;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct HidCassetteMetadata {
pub name: String,
pub channel: String,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub enum SanitizedIdentityKind {
ReceiverUniqueId,
ReceiverSerialNumber,
DeviceUnitId,
DeviceSerialNumber,
}
impl SanitizedIdentityKind {
const ALL: [Self; 4] = [
Self::ReceiverUniqueId,
Self::ReceiverSerialNumber,
Self::DeviceUnitId,
Self::DeviceSerialNumber,
];
const fn policy_kind(self) -> SyntheticIdentityKind {
match self {
Self::ReceiverUniqueId => SyntheticIdentityKind::BoltReceiverUid,
Self::ReceiverSerialNumber => SyntheticIdentityKind::UnifyingReceiverSerial,
Self::DeviceUnitId => SyntheticIdentityKind::DeviceUnitId,
Self::DeviceSerialNumber => SyntheticIdentityKind::DeviceSerialNumber,
}
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct HidCassetteIdentityPlan {
replacements: BTreeMap<SanitizedIdentityKind, Vec<u8>>,
}
impl HidCassetteIdentityPlan {
pub fn insert(
&mut self,
kind: SanitizedIdentityKind,
synthetic_value: Vec<u8>,
) -> Result<(), HidCassetteIdentityPlanError> {
classify_synthetic_identity_bytes(kind.policy_kind(), &synthetic_value)
.map_err(|_| HidCassetteIdentityPlanError::InvalidValue { kind })?;
if self.replacements.contains_key(&kind) {
return Err(HidCassetteIdentityPlanError::DuplicateKind { kind });
}
self.replacements.insert(kind, synthetic_value);
Ok(())
}
pub(super) fn replacement(&self, kind: SanitizedIdentityKind) -> Option<&[u8]> {
self.replacements.get(&kind).map(Vec::as_slice)
}
}
#[derive(Clone, Debug, Error, PartialEq, Eq)]
pub enum HidCassetteIdentityPlanError {
#[error("identity plan repeats {kind:?}")]
DuplicateKind {
kind: SanitizedIdentityKind,
},
#[error("identity plan contains a noncanonical value for {kind:?}")]
InvalidValue {
kind: SanitizedIdentityKind,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct IdentityReplacement {
pub kind: SanitizedIdentityKind,
pub synthetic_value: Vec<u8>,
pub occurrences: usize,
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct HidCassetteAudit {
pub replacements: Vec<IdentityReplacement>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum CassetteRejectionReason {
ChannelNotOpened,
UnsupportedReportSupport,
IncompleteChannel,
DuplicateRequestId,
OutgoingReportCount {
actual: usize,
},
OutcomeCount {
actual: usize,
},
IncomingReportCount {
actual: usize,
},
RequestTimedOut,
RequestWriteFailed,
RequestLostResponse,
RequestCancelled,
UnsupportedRequestOutcome,
UnprovenFireAndForget,
UnmatchedIncomingReport,
MalformedIncomingReport,
UnsupportedObservation,
MalformedReport,
CorrelationMismatch,
UnsupportedCrossVersionPing,
UnsupportedHidpp10Register,
UnknownFeatureIndex {
device_index: u8,
feature_index: u8,
},
UnsupportedIdentityFeature {
feature_id: u16,
},
UnsupportedHidpp20Function {
feature_id: u16,
function_id: u8,
},
AmbiguousFeatureMapping,
MalformedIdentity,
SyntheticIdentitySpaceExhausted,
PlannedIdentityMatchesOriginal,
PairingTraffic,
EmptyCassette,
InvalidCassette {
message: String,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct CassetteRejection {
pub request_id: Option<u64>,
pub reason: CassetteRejectionReason,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct HidCassetteBuildReport {
pub cassette: Option<HidCassette>,
pub audit: HidCassetteAudit,
pub rejections: Vec<CassetteRejection>,
}
impl HidCassetteBuildReport {
#[must_use]
pub fn is_committable(&self) -> bool {
self.cassette.is_some() && self.rejections.is_empty()
}
}
impl RecordedChannel {
#[must_use]
pub fn build_hid_cassette(&self, metadata: HidCassetteMetadata) -> HidCassetteBuildReport {
self.build_hid_cassette_with_identity_plan(metadata, &HidCassetteIdentityPlan::default())
}
#[must_use]
pub fn build_hid_cassette_with_identity_plan(
&self,
metadata: HidCassetteMetadata,
identity_plan: &HidCassetteIdentityPlan,
) -> HidCassetteBuildReport {
CassetteBuilder::new(self, metadata, identity_plan).build()
}
}
struct CassetteBuilder<'a> {
channel: &'a RecordedChannel,
metadata: HidCassetteMetadata,
sanitizer: ProtocolSanitizer,
rejections: Vec<CassetteRejection>,
}
impl<'a> CassetteBuilder<'a> {
fn new(
channel: &'a RecordedChannel,
metadata: HidCassetteMetadata,
identity_plan: &HidCassetteIdentityPlan,
) -> Self {
Self {
channel,
metadata,
sanitizer: ProtocolSanitizer::with_identity_plan(identity_plan),
rejections: Vec::new(),
}
}
fn build(mut self) -> HidCassetteBuildReport {
let report_support = self.validate_channel();
self.validate_unassociated();
let mut requests: Vec<_> = self.channel.requests.iter().collect();
requests.sort_by_key(|request| request.request_id);
self.reject_duplicate_request_ids(&requests);
let mut exchanges = Vec::with_capacity(requests.len());
for request in requests {
if let Some(exchange) = self.build_exchange(request) {
exchanges.push(exchange);
}
}
if exchanges.is_empty() {
self.reject(None, CassetteRejectionReason::EmptyCassette);
}
let cassette = if self.rejections.is_empty() {
if let Some(report_support) = report_support {
let cassette = HidCassette {
schema_version: FIXTURE_SCHEMA_VERSION,
name: self.metadata.name.clone(),
channel: self.metadata.channel.clone(),
report_support,
exchanges,
};
match cassette.validate() {
Ok(()) => Some(cassette),
Err(error) => {
self.reject(
None,
CassetteRejectionReason::InvalidCassette {
message: error.to_string(),
},
);
None
}
}
} else {
None
}
} else {
None
};
let audit = self.sanitizer.finish();
HidCassetteBuildReport {
cassette,
audit,
rejections: self.rejections,
}
}
fn validate_channel(&mut self) -> Option<ReportSupport> {
let support = match self.channel.open_outcome {
RecordedChannelOpenOutcome::Opened {
supports_short: true,
supports_long: true,
} => Some(ReportSupport::ShortAndLong),
RecordedChannelOpenOutcome::Opened {
supports_short: false,
supports_long: true,
} => Some(ReportSupport::LongOnly),
RecordedChannelOpenOutcome::Opened { .. } => {
self.reject(None, CassetteRejectionReason::UnsupportedReportSupport);
None
}
RecordedChannelOpenOutcome::NotHidpp
| RecordedChannelOpenOutcome::Failed(_)
| RecordedChannelOpenOutcome::Cancelled => {
self.reject(None, CassetteRejectionReason::ChannelNotOpened);
None
}
};
if self.channel.closed_at.is_none() {
self.reject(None, CassetteRejectionReason::IncompleteChannel);
}
support
}
fn validate_unassociated(&mut self) {
for evidence in &self.channel.unassociated {
self.reject(None, unassociated_rejection(&evidence.observation));
}
}
fn reject_duplicate_request_ids(&mut self, requests: &[&RecordedRequest]) {
let mut ids = BTreeSet::new();
for request in requests {
if !ids.insert(request.request_id) {
self.reject(
Some(request.request_id),
CassetteRejectionReason::DuplicateRequestId,
);
}
}
}
fn build_exchange(&mut self, request: &RecordedRequest) -> Option<CassetteExchange> {
let mut outgoing = Vec::new();
let mut incoming = Vec::new();
let mut outcomes = Vec::new();
for fact in &request.facts {
match fact {
RecordedRequestFact::OutgoingReport { report, .. } => outgoing.push(report),
RecordedRequestFact::IncomingReport { report, .. } => incoming.push(report),
RecordedRequestFact::Outcome { outcome, .. } => outcomes.push(*outcome),
}
}
let mut valid = true;
if outgoing.len() != 1 {
self.reject(
Some(request.request_id),
CassetteRejectionReason::OutgoingReportCount {
actual: outgoing.len(),
},
);
valid = false;
}
if outcomes.len() != 1 {
self.reject(
Some(request.request_id),
CassetteRejectionReason::OutcomeCount {
actual: outcomes.len(),
},
);
valid = false;
}
if !valid {
return None;
}
match outcomes[0] {
RequestOutcome::Succeeded => {}
RequestOutcome::TimedOut => {
self.reject(
Some(request.request_id),
CassetteRejectionReason::RequestTimedOut,
);
return None;
}
RequestOutcome::WriteFailed => {
self.reject(
Some(request.request_id),
CassetteRejectionReason::RequestWriteFailed,
);
return None;
}
RequestOutcome::NoResponse => {
self.reject(
Some(request.request_id),
CassetteRejectionReason::RequestLostResponse,
);
return None;
}
RequestOutcome::Cancelled => {
self.reject(
Some(request.request_id),
CassetteRejectionReason::RequestCancelled,
);
return None;
}
_ => {
self.reject(
Some(request.request_id),
CassetteRejectionReason::UnsupportedRequestOutcome,
);
return None;
}
}
if incoming.len() != 1 {
self.reject(
Some(request.request_id),
CassetteRejectionReason::IncomingReportCount {
actual: incoming.len(),
},
);
return None;
}
match self
.sanitizer
.exchange(outgoing[0].as_bytes(), incoming[0].as_bytes())
{
Ok(exchange) => Some(exchange),
Err(reason) => {
self.reject(Some(request.request_id), reason);
None
}
}
}
fn reject(&mut self, request_id: Option<u64>, reason: CassetteRejectionReason) {
self.rejections
.push(CassetteRejection { request_id, reason });
}
}