#![forbid(unsafe_code)]
use serde::de::{self, SeqAccess, Visitor};
use serde::ser::SerializeTuple;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use std::fmt;
pub const FEATURE_COUNT: usize = 35;
const MAX_KEY_BYTES: usize = 128;
const PERCEIVED_AGE_INDEX: usize = 11;
const FEATURE_BITS: u64 = (1_u64 << FEATURE_COUNT) - 1;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct TypeError {
kind: ErrorKind,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ErrorKind {
InvalidKey,
InvalidObjectId,
FeatureValue {
index: usize,
value: u8,
maximum: u8,
},
EmptyFeatureMask,
UnknownFeatureBits,
SegmentOrdinal,
SegmentTime,
}
impl TypeError {
const fn new(kind: ErrorKind) -> Self {
Self { kind }
}
}
impl fmt::Display for TypeError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self.kind {
ErrorKind::InvalidKey => formatter.write_str(
"key must contain 1 to 128 ASCII letters, digits, or . _ - / : characters",
),
ErrorKind::InvalidObjectId => {
formatter.write_str("object ID must contain exactly eight base64url characters")
}
ErrorKind::FeatureValue {
index,
value,
maximum,
} => write!(
formatter,
"feature at index {index} has value {value}, exceeding {maximum}"
),
ErrorKind::EmptyFeatureMask => {
formatter.write_str("feature mask must contain at least one bit")
}
ErrorKind::UnknownFeatureBits => {
formatter.write_str("feature mask may contain only bits 0 through 34")
}
ErrorKind::SegmentOrdinal => {
formatter.write_str("segment ordinal must be less than segment count")
}
ErrorKind::SegmentTime => {
formatter.write_str("segment start must be earlier than segment end")
}
}
}
}
impl std::error::Error for TypeError {}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize)]
#[serde(transparent)]
pub struct Key(String);
impl Key {
pub fn parse(value: &str) -> Result<Self, TypeError> {
if valid_key(value) {
Ok(Self(value.to_owned()))
} else {
Err(TypeError::new(ErrorKind::InvalidKey))
}
}
}
fn valid_key(value: &str) -> bool {
let bytes = value.as_bytes();
(1..=MAX_KEY_BYTES).contains(&bytes.len())
&& bytes.iter().all(|byte| {
byte.is_ascii_alphanumeric() || matches!(*byte, b'.' | b'_' | b'-' | b'/' | b':')
})
}
impl AsRef<str> for Key {
fn as_ref(&self) -> &str {
&self.0
}
}
impl<'de> Deserialize<'de> for Key {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let value = String::deserialize(deserializer)?;
Self::parse(&value).map_err(de::Error::custom)
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize)]
#[serde(transparent)]
pub struct ObjectId(String);
impl ObjectId {
pub fn parse(value: &str) -> Result<Self, TypeError> {
let bytes = value.as_bytes();
if bytes.len() == 8
&& bytes
.iter()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(*byte, b'-' | b'_'))
{
Ok(Self(value.to_owned()))
} else {
Err(TypeError::new(ErrorKind::InvalidObjectId))
}
}
}
impl AsRef<str> for ObjectId {
fn as_ref(&self) -> &str {
&self.0
}
}
impl<'de> Deserialize<'de> for ObjectId {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let value = String::deserialize(deserializer)?;
Self::parse(&value).map_err(de::Error::custom)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct FeatureVector([u8; FEATURE_COUNT]);
impl FeatureVector {
pub fn new(values: [u8; FEATURE_COUNT]) -> Result<Self, TypeError> {
for (index, value) in values.iter().copied().enumerate() {
let maximum = if index == PERCEIVED_AGE_INDEX {
120
} else {
100
};
if value > maximum {
return Err(TypeError::new(ErrorKind::FeatureValue {
index,
value,
maximum,
}));
}
}
Ok(Self(values))
}
}
impl AsRef<[u8; FEATURE_COUNT]> for FeatureVector {
fn as_ref(&self) -> &[u8; FEATURE_COUNT] {
&self.0
}
}
impl Serialize for FeatureVector {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
let mut tuple = serializer.serialize_tuple(FEATURE_COUNT)?;
for value in self.0 {
tuple.serialize_element(&value)?;
}
tuple.end()
}
}
struct FeatureVectorVisitor;
impl<'de> Visitor<'de> for FeatureVectorVisitor {
type Value = FeatureVector;
fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(formatter, "an array of {FEATURE_COUNT} feature values")
}
fn visit_seq<A>(self, mut sequence: A) -> Result<Self::Value, A::Error>
where
A: SeqAccess<'de>,
{
let mut values = [0; FEATURE_COUNT];
for (index, slot) in values.iter_mut().enumerate() {
*slot = sequence
.next_element()?
.ok_or_else(|| de::Error::invalid_length(index, &self))?;
}
if sequence.next_element::<de::IgnoredAny>()?.is_some() {
return Err(de::Error::invalid_length(FEATURE_COUNT + 1, &self));
}
FeatureVector::new(values).map_err(de::Error::custom)
}
}
impl<'de> Deserialize<'de> for FeatureVector {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_tuple(FEATURE_COUNT, FeatureVectorVisitor)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct FeatureMask(u64);
impl FeatureMask {
pub const fn from_bits(bits: u64) -> Result<Self, TypeError> {
if bits == 0 {
Err(TypeError::new(ErrorKind::EmptyFeatureMask))
} else if bits & !FEATURE_BITS != 0 {
Err(TypeError::new(ErrorKind::UnknownFeatureBits))
} else {
Ok(Self(bits))
}
}
}
impl From<FeatureMask> for u64 {
fn from(mask: FeatureMask) -> Self {
mask.0
}
}
impl Serialize for FeatureMask {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
serializer.serialize_u64(self.0)
}
}
impl<'de> Deserialize<'de> for FeatureMask {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let bits = u64::deserialize(deserializer)?;
Self::from_bits(bits).map_err(de::Error::custom)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum RecordingKind {
VoiceNote,
Meeting,
Call,
Other,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct SegmentRef {
pub source_object: ObjectId,
pub clip_object: ObjectId,
pub ordinal: u16,
pub segment_count: u16,
pub start_ms: u64,
pub end_ms: u64,
pub policy: Key,
}
impl SegmentRef {
pub fn validate(&self) -> Result<(), TypeError> {
if self.ordinal >= self.segment_count {
Err(TypeError::new(ErrorKind::SegmentOrdinal))
} else if self.start_ms >= self.end_ms {
Err(TypeError::new(ErrorKind::SegmentTime))
} else {
Ok(())
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct LabeledSample {
pub sample_id: Key,
pub attempt_id: Key,
pub speaker_id: Key,
pub cohort_id: Key,
pub group_id: Key,
pub clip_object: ObjectId,
pub primary_language: Key,
pub recording_kind: RecordingKind,
pub usable_speech_ms: u32,
pub recording_quality: u8,
pub features: FeatureVector,
}
#[cfg(test)]
mod tests {
use super::*;
use serde::de::IntoDeserializer;
use serde::ser::{Impossible, SerializeStruct};
const CONST_FEATURE_MASK: FeatureMask = match FeatureMask::from_bits(1) {
Ok(mask) => mask,
Err(_) => panic!("one feature bit is valid"),
};
#[derive(Debug)]
struct CodecError(String);
impl fmt::Display for CodecError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str(&self.0)
}
}
impl std::error::Error for CodecError {}
impl serde::ser::Error for CodecError {
fn custom<T: fmt::Display>(message: T) -> Self {
Self(message.to_string())
}
}
impl serde::de::Error for CodecError {
fn custom<T: fmt::Display>(message: T) -> Self {
Self(message.to_string())
}
}
#[derive(Debug)]
enum Value {
U8(u8),
U16(u16),
U32(u32),
U64(u64),
String(String),
Sequence(Vec<Value>),
Struct(Vec<(String, Value)>),
Variant(String),
}
struct ValueSerializer;
impl Serializer for ValueSerializer {
type Ok = Value;
type Error = CodecError;
type SerializeSeq = Impossible<Value, CodecError>;
type SerializeTuple = TupleBuilder;
type SerializeTupleStruct = Impossible<Value, CodecError>;
type SerializeTupleVariant = Impossible<Value, CodecError>;
type SerializeMap = Impossible<Value, CodecError>;
type SerializeStruct = StructBuilder;
type SerializeStructVariant = Impossible<Value, CodecError>;
fn serialize_u8(self, value: u8) -> Result<Self::Ok, Self::Error> {
Ok(Value::U8(value))
}
fn serialize_u16(self, value: u16) -> Result<Self::Ok, Self::Error> {
Ok(Value::U16(value))
}
fn serialize_u32(self, value: u32) -> Result<Self::Ok, Self::Error> {
Ok(Value::U32(value))
}
fn serialize_u64(self, value: u64) -> Result<Self::Ok, Self::Error> {
Ok(Value::U64(value))
}
fn serialize_str(self, value: &str) -> Result<Self::Ok, Self::Error> {
Ok(Value::String(value.to_owned()))
}
fn serialize_unit_variant(
self,
_name: &'static str,
_variant_index: u32,
variant: &'static str,
) -> Result<Self::Ok, Self::Error> {
Ok(Value::Variant(variant.to_owned()))
}
fn serialize_newtype_struct<T>(
self,
_name: &'static str,
value: &T,
) -> Result<Self::Ok, Self::Error>
where
T: ?Sized + Serialize,
{
value.serialize(self)
}
fn serialize_some<T>(self, value: &T) -> Result<Self::Ok, Self::Error>
where
T: ?Sized + Serialize,
{
value.serialize(self)
}
fn serialize_tuple(self, length: usize) -> Result<Self::SerializeTuple, Self::Error> {
Ok(TupleBuilder(Vec::with_capacity(length)))
}
fn serialize_struct(
self,
_name: &'static str,
length: usize,
) -> Result<Self::SerializeStruct, Self::Error> {
Ok(StructBuilder(Vec::with_capacity(length)))
}
fn serialize_bool(self, _value: bool) -> Result<Self::Ok, Self::Error> {
unsupported()
}
fn serialize_i8(self, _value: i8) -> Result<Self::Ok, Self::Error> {
unsupported()
}
fn serialize_i16(self, _value: i16) -> Result<Self::Ok, Self::Error> {
unsupported()
}
fn serialize_i32(self, _value: i32) -> Result<Self::Ok, Self::Error> {
unsupported()
}
fn serialize_i64(self, _value: i64) -> Result<Self::Ok, Self::Error> {
unsupported()
}
fn serialize_i128(self, _value: i128) -> Result<Self::Ok, Self::Error> {
unsupported()
}
fn serialize_u128(self, _value: u128) -> Result<Self::Ok, Self::Error> {
unsupported()
}
fn serialize_f32(self, _value: f32) -> Result<Self::Ok, Self::Error> {
unsupported()
}
fn serialize_f64(self, _value: f64) -> Result<Self::Ok, Self::Error> {
unsupported()
}
fn serialize_char(self, _value: char) -> Result<Self::Ok, Self::Error> {
unsupported()
}
fn serialize_bytes(self, _value: &[u8]) -> Result<Self::Ok, Self::Error> {
unsupported()
}
fn serialize_none(self) -> Result<Self::Ok, Self::Error> {
unsupported()
}
fn serialize_unit(self) -> Result<Self::Ok, Self::Error> {
unsupported()
}
fn serialize_unit_struct(self, _name: &'static str) -> Result<Self::Ok, Self::Error> {
unsupported()
}
fn serialize_newtype_variant<T>(
self,
_name: &'static str,
_variant_index: u32,
_variant: &'static str,
_value: &T,
) -> Result<Self::Ok, Self::Error>
where
T: ?Sized + Serialize,
{
unsupported()
}
fn serialize_seq(self, _length: Option<usize>) -> Result<Self::SerializeSeq, Self::Error> {
unsupported()
}
fn serialize_tuple_struct(
self,
_name: &'static str,
_length: usize,
) -> Result<Self::SerializeTupleStruct, Self::Error> {
unsupported()
}
fn serialize_tuple_variant(
self,
_name: &'static str,
_variant_index: u32,
_variant: &'static str,
_length: usize,
) -> Result<Self::SerializeTupleVariant, Self::Error> {
unsupported()
}
fn serialize_map(self, _length: Option<usize>) -> Result<Self::SerializeMap, Self::Error> {
unsupported()
}
fn serialize_struct_variant(
self,
_name: &'static str,
_variant_index: u32,
_variant: &'static str,
_length: usize,
) -> Result<Self::SerializeStructVariant, Self::Error> {
unsupported()
}
}
struct TupleBuilder(Vec<Value>);
impl serde::ser::SerializeTuple for TupleBuilder {
type Ok = Value;
type Error = CodecError;
fn serialize_element<T>(&mut self, value: &T) -> Result<(), Self::Error>
where
T: ?Sized + Serialize,
{
self.0.push(value.serialize(ValueSerializer)?);
Ok(())
}
fn end(self) -> Result<Self::Ok, Self::Error> {
Ok(Value::Sequence(self.0))
}
}
struct StructBuilder(Vec<(String, Value)>);
impl SerializeStruct for StructBuilder {
type Ok = Value;
type Error = CodecError;
fn serialize_field<T>(&mut self, key: &'static str, value: &T) -> Result<(), Self::Error>
where
T: ?Sized + Serialize,
{
self.0
.push((key.to_owned(), value.serialize(ValueSerializer)?));
Ok(())
}
fn end(self) -> Result<Self::Ok, Self::Error> {
Ok(Value::Struct(self.0))
}
}
fn unsupported<T>() -> Result<T, CodecError> {
Err(CodecError("unsupported test value".to_owned()))
}
impl<'de> IntoDeserializer<'de, CodecError> for Value {
type Deserializer = Self;
fn into_deserializer(self) -> Self::Deserializer {
self
}
}
impl<'de> Deserializer<'de> for Value {
type Error = CodecError;
fn deserialize_any<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: Visitor<'de>,
{
match self {
Self::U8(value) => visitor.visit_u8(value),
Self::U16(value) => visitor.visit_u16(value),
Self::U32(value) => visitor.visit_u32(value),
Self::U64(value) => visitor.visit_u64(value),
Self::String(value) => visitor.visit_string(value),
Self::Sequence(values) => {
visitor.visit_seq(serde::de::value::SeqDeserializer::<_, CodecError>::new(
values.into_iter(),
))
}
Self::Struct(fields) => {
visitor.visit_map(serde::de::value::MapDeserializer::<_, CodecError>::new(
fields.into_iter(),
))
}
Self::Variant(name) => visitor.visit_enum(serde::de::value::StringDeserializer::<
CodecError,
>::new(name)),
}
}
serde::forward_to_deserialize_any! {
bool i8 i16 i32 i64 i128 u8 u16 u32 u64 u128 f32 f64 char str string
bytes byte_buf option unit unit_struct newtype_struct seq tuple tuple_struct
map struct enum identifier ignored_any
}
}
fn round_trip<T>(value: &T) -> T
where
T: Serialize + for<'de> Deserialize<'de>,
{
let encoded = value.serialize(ValueSerializer).expect("serialize");
T::deserialize(encoded).expect("deserialize")
}
fn key(value: &str) -> Key {
Key::parse(value).expect("valid key")
}
fn object(value: &str) -> ObjectId {
ObjectId::parse(value).expect("valid object ID")
}
fn segment() -> SegmentRef {
SegmentRef {
source_object: object("source_1"),
clip_object: object("clip___1"),
ordinal: 0,
segment_count: 1,
start_ms: 10,
end_ms: 20,
policy: key("policy:v1"),
}
}
#[test]
fn accepts_all_key_characters_and_boundaries() {
for value in ["a", "Az09._-/:", &"x".repeat(MAX_KEY_BYTES)] {
assert!(Key::parse(value).is_ok(), "{value:?}");
}
}
#[test]
fn rejects_invalid_key_shapes() {
for value in [
"",
&"x".repeat(MAX_KEY_BYTES + 1),
"has space",
"line\nbreak",
"é",
"key+value",
] {
assert!(Key::parse(value).is_err(), "{value:?}");
}
}
#[test]
fn accepts_and_rejects_object_id_shapes() {
for value in ["Ab09-_xy", "________", "00000000"] {
assert!(ObjectId::parse(value).is_ok(), "{value:?}");
}
for value in [
"",
"1234567",
"123456789",
"abc+defg",
"abc/defg",
"abc=defg",
"pending:",
"pending:N",
"é1234567",
] {
assert!(ObjectId::parse(value).is_err(), "{value:?}");
}
}
#[test]
fn validates_feature_boundaries_and_age_exception() {
let mut values = [100; FEATURE_COUNT];
values[PERCEIVED_AGE_INDEX] = 120;
assert_eq!(
FeatureVector::new(values)
.expect("boundary values")
.as_ref(),
&values
);
let mut invalid_age = [0; FEATURE_COUNT];
invalid_age[PERCEIVED_AGE_INDEX] = 121;
assert!(FeatureVector::new(invalid_age).is_err());
let mut invalid_other = [0; FEATURE_COUNT];
invalid_other[0] = 101;
assert!(FeatureVector::new(invalid_other).is_err());
invalid_other[0] = 0;
invalid_other[FEATURE_COUNT - 1] = 101;
assert!(FeatureVector::new(invalid_other).is_err());
}
#[test]
fn validates_feature_mask_bits() {
assert_eq!(u64::from(CONST_FEATURE_MASK), 1);
assert!(FeatureMask::from_bits(0).is_err());
assert!(FeatureMask::from_bits(1_u64 << FEATURE_COUNT).is_err());
assert!(FeatureMask::from_bits(u64::MAX).is_err());
let mask = FeatureMask::from_bits(FEATURE_BITS).expect("all defined bits");
assert_eq!(u64::from(mask), FEATURE_BITS);
}
#[test]
fn validates_segment_bounds() {
let mut value = segment();
assert!(value.validate().is_ok());
value.ordinal = value.segment_count;
assert!(value.validate().is_err());
value.ordinal = 0;
value.segment_count = 0;
assert!(value.validate().is_err());
value.segment_count = 1;
value.start_ms = value.end_ms;
assert!(value.validate().is_err());
value.start_ms = value.end_ms + 1;
assert!(value.validate().is_err());
}
#[test]
fn serde_round_trips_public_values() {
let feature = FeatureVector::new([42; FEATURE_COUNT]).expect("features");
assert_eq!(round_trip(&key("sample:1")), key("sample:1"));
assert_eq!(round_trip(&object("object_1")), object("object_1"));
assert_eq!(round_trip(&feature), feature);
let mask = FeatureMask::from_bits(5).expect("mask");
assert_eq!(round_trip(&mask), mask);
for kind in [
RecordingKind::VoiceNote,
RecordingKind::Meeting,
RecordingKind::Call,
RecordingKind::Other,
] {
assert_eq!(round_trip(&kind), kind);
}
}
#[test]
fn serde_round_trips_shared_records() {
let segment = segment();
assert_eq!(round_trip(&segment), segment);
let sample = LabeledSample {
sample_id: key("sample:1"),
attempt_id: key("attempt:1"),
speaker_id: key("speaker:1"),
cohort_id: key("cohort:1"),
group_id: key("group:1"),
clip_object: object("clip___1"),
primary_language: key("en-US"),
recording_kind: RecordingKind::VoiceNote,
usable_speech_ms: 12_345,
recording_quality: 88,
features: FeatureVector::new([50; FEATURE_COUNT]).expect("features"),
};
assert_eq!(round_trip(&sample), sample);
}
#[test]
fn serde_rejects_invalid_validated_values() {
assert!(Key::deserialize(Value::String("not valid".to_owned())).is_err());
assert!(ObjectId::deserialize(Value::String("pending:".to_owned())).is_err());
assert!(FeatureMask::deserialize(Value::U64(0)).is_err());
let mut values = (0..FEATURE_COUNT).map(|_| Value::U8(0)).collect::<Vec<_>>();
values[PERCEIVED_AGE_INDEX] = Value::U8(121);
assert!(FeatureVector::deserialize(Value::Sequence(values)).is_err());
}
}