kcode-k1-audio-classification-fragment-format 0.2.2

Fragment binary formats for K1 audio classification
Documentation
use super::*;

fn txid(seed: u8) -> TxId {
    TxId::from_bytes([seed; 12])
}

fn staged_fragment() -> StagedFragmentV1 {
    StagedFragmentV1 {
        analysis_txid: txid(20),
        transcript: "hello".into(),
        speakers: vec![StagedSpeakerV1 {
            speaker: LocalSpeakerLabel::new(1).expect("speaker label"),
            language: "English".into(),
            features: FeatureVector24::default(),
            usable_for_training: true,
        }],
    }
}

fn final_fragment() -> FinalFragmentV1 {
    FinalFragmentV1 {
        analysis_txid: txid(30),
        confirmation_txid: txid(40),
        transcript: "hello".into(),
        speakers: vec![FinalSpeakerV1 {
            speaker: LocalSpeakerLabel::new(1).expect("speaker label"),
            person_id: Some(PersonId::from_tx_id(txid(7))),
            language: "English".into(),
            features: FeatureVector24::default(),
            usable_for_training: false,
        }],
    }
}

#[test]
fn v2_identity_byte_golden_and_fragments_roundtrip() {
    let staged = staged_fragment();
    let staged_bytes = encode_staged_fragment(&staged).expect("encode staged");
    assert_eq!(decode_staged_fragment(&staged_bytes), Ok(staged));

    let final_value = final_fragment();
    let final_bytes = encode_final_fragment(&final_value).expect("encode final");
    let identity_offset = BODY_OFFSET
        + 24
        + final_value.transcript.len()
        + final_value.speakers[0].speaker.to_string().len();
    assert_eq!(&final_bytes[..2], &[2, 2]);
    let mut identity_golden = [7; 13];
    identity_golden[0] = 1;
    assert_eq!(
        &final_bytes[identity_offset..identity_offset + 13],
        &identity_golden
    );
    assert_eq!(decode_final_fragment(&final_bytes), Ok(final_value.clone()));

    let mut unknown = final_value;
    unknown.speakers[0].person_id = None;
    let unknown_bytes = encode_final_fragment(&unknown).expect("encode unknown");
    assert_eq!(unknown_bytes[identity_offset], 0);
    assert_eq!(final_bytes.len(), unknown_bytes.len() + 12);
    assert_eq!(decode_final_fragment(&unknown_bytes), Ok(unknown));

    let mut invalid_identity = final_bytes;
    invalid_identity[identity_offset] = 2;
    assert_eq!(
        decode_final_fragment(&invalid_identity),
        Err(FormatError::InvalidBoolean(2))
    );
    assert_eq!(
        TxIdSlot::decode(&[0; 15]),
        Err(FormatError::InvalidTxIdSlotLength(15))
    );
}

#[test]
fn malformed_headers_slots_and_bodies_are_rejected() {
    let staged = staged_fragment();
    let bytes = encode_staged_fragment(&staged).expect("encode staged");
    let transcript = BODY_OFFSET + 8;
    let label = transcript + staged.transcript.len() + 16;
    let feature_length = label + staged.speakers[0].speaker.to_string().len() + 15;
    let feature = feature_length + 8;
    let boolean = feature
        + postcard::to_allocvec(&staged.speakers[0].features)
            .expect("features")
            .len();
    for (index, value, error) in [
        (0, 1, FormatError::UnsupportedVersion(1)),
        (1, 2, FormatError::InvalidFragmentKind(2)),
        (2, 1, FormatError::NonZeroReserved),
        (28, 1, FormatError::NonZeroPadding),
        (32, 1, FormatError::NonZeroStagedConfirmation),
        (transcript, 0xff, FormatError::InvalidUtf8),
        (label, b'x', FormatError::InvalidSpeakerLabel),
        (feature, 2, FormatError::InvalidFeatureBody),
        (boolean, 2, FormatError::InvalidBoolean(2)),
    ] {
        let mut malformed = bytes.clone();
        malformed[index] = value;
        assert_eq!(decode_staged_fragment(&malformed), Err(error));
    }
    assert_eq!(
        decode_staged_fragment(&bytes[..47]),
        Err(FormatError::Truncated)
    );
    let mut trailing = bytes;
    trailing.push(0);
    assert_eq!(
        decode_staged_fragment(&trailing),
        Err(FormatError::TrailingBytes)
    );
}

#[test]
fn every_path_shard_roundtrips_and_noncanonical_paths_fail() {
    for (index, expected) in BASE64_ALPHABET.iter().enumerate() {
        let mut bytes = [0; 12];
        bytes[0] = (index as u8) << 2;
        let txid = TxId::from_bytes(bytes);
        let path = txid_path(txid);
        assert_eq!(path.to_string_lossy().as_bytes()[0], *expected);
        assert_eq!(txid_from_path(&path), Ok(txid));
    }
    for path in [
        "/A/AAAAAAAAAAAAAAA.dat",
        "AA/AAAAAAAAAAAAAAA.dat",
        "A/AAAAAAAAAAAAAA!.dat",
        "A/AAAAAAAAAAAAAAA.dat/extra",
    ] {
        assert_eq!(txid_from_path(path), Err(FormatError::InvalidPath));
    }
}