use crate::courierust_error::Error;
const RANGES: [(u8, u8); 4] = [(0xA0, 0xBF), (0x80, 0x9F), (0x90, 0xBF), (0x80, 0x8F)];
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct Utf8Validator {
state: u8,
}
impl Utf8Validator {
#[inline]
pub const fn new() -> Self {
Self { state: 0 }
}
#[inline]
pub fn reset(&mut self) {
self.state = 0;
}
#[inline]
pub fn is_complete(&self) -> bool {
self.state == 0
}
#[inline]
pub fn at_boundary(&self) -> bool {
self.state == 0
}
pub fn feed(&mut self, bytes: &[u8]) -> core::result::Result<(), Utf8Error> {
let mut start = 0usize;
if self.state != 0 {
start = self.finish_open_character(bytes)?;
}
match core::str::from_utf8(&bytes[start..]) {
Ok(_) => Ok(()),
Err(e) => {
let mut i = start + e.valid_up_to();
while i < bytes.len() {
self.feed_byte(bytes[i], i)?;
i += 1;
}
Ok(())
}
}
}
fn finish_open_character(&mut self, bytes: &[u8]) -> core::result::Result<usize, Utf8Error> {
let mut i = 0usize;
while self.state != 0 && i < bytes.len() {
self.feed_byte(bytes[i], i)?;
i += 1;
}
Ok(i)
}
#[inline]
fn feed_byte(&mut self, b: u8, offset: usize) -> core::result::Result<(), Utf8Error> {
match self.state {
0 => {
self.state = match b {
0x00..=0x7F => 0,
0xC2..=0xDF => 1,
0xE1..=0xEC | 0xEE..=0xEF => 2,
0xE0 => 4,
0xED => 5,
0xF0 => 6,
0xF1..=0xF3 => 3,
0xF4 => 7,
_ => {
return Err(Utf8Error {
offset,
reason: "invalid lead byte",
})
}
};
}
1..=3 => {
if !(0x80..=0xBF).contains(&b) {
return Err(Utf8Error {
offset,
reason: "invalid continuation byte",
});
}
self.state -= 1;
}
_ => {
let idx = (self.state - 4) as usize;
let (lo, hi) = RANGES[idx];
if b < lo || b > hi {
return Err(Utf8Error {
offset,
reason: if idx == 1 {
"UTF-16 surrogate half is not a scalar value"
} else if idx == 3 {
"codepoint above U+10FFFF"
} else {
"overlong encoding"
},
});
}
self.state = if idx < 2 { 1 } else { 2 };
}
}
Ok(())
}
pub fn validate(bytes: &[u8]) -> bool {
let mut v = Self::new();
v.feed(bytes).is_ok() && v.is_complete()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Utf8Error {
pub offset: usize,
pub reason: &'static str,
}
impl Utf8Error {
pub fn into_error(self) -> Error {
Error::protocol(alloc::format!(
"websocket: invalid UTF-8 at byte {}: {}",
self.offset,
self.reason
))
}
}
impl core::fmt::Display for Utf8Error {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(f, "invalid UTF-8 at byte {}: {}", self.offset, self.reason)
}
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::vec::Vec;
#[test]
fn boundary_scalars_are_accepted() {
let cases: &[&str] = &[
"",
"a",
"\u{7F}",
"\u{80}",
"\u{7FF}",
"\u{800}",
"\u{D7FF}",
"\u{E000}",
"\u{FFFF}",
"\u{10000}",
"\u{10FFFF}",
"日本語テキスト 🚀 mixed ascii",
];
for s in cases {
assert!(Utf8Validator::validate(s.as_bytes()), "{s:?} must be valid");
}
}
#[test]
fn rejects_the_unicode_spec_falsehoods() {
let cases: &[&[u8]] = &[
&[0x80], &[0xC0, 0x80], &[0xC1, 0xBF], &[0xE0, 0x80, 0x80], &[0xE0, 0x9F, 0xBF], &[0xED, 0xA0, 0x80], &[0xED, 0xBF, 0xBF], &[0xF0, 0x80, 0x80, 0x80], &[0xF0, 0x8F, 0xBF, 0xBF], &[0xF4, 0x90, 0x80, 0x80], &[0xF5, 0x80, 0x80, 0x80], &[0xFF], &[0xC2], &[0xE2, 0x82], &[0xF0, 0x9F, 0x98], &[0xE2, 0x28, 0xA1], ];
for c in cases {
let mut v = Utf8Validator::new();
let ok = v.feed(c).is_ok() && v.is_complete();
assert!(!ok, "{c:02x?} must be rejected");
}
}
#[test]
fn agrees_with_std_for_all_two_byte_inputs() {
for a in 0u16..=255 {
for b in 0u16..=255 {
let bytes = [a as u8, b as u8];
let ours = Utf8Validator::validate(&bytes);
let theirs = core::str::from_utf8(&bytes).is_ok();
assert_eq!(ours, theirs, "disagreement on {bytes:02x?}");
}
}
}
#[test]
fn agrees_with_std_on_random_sequences() {
let mut state = 0x1234_5678u32;
let mut next = move || {
state ^= state << 13;
state ^= state >> 17;
state ^= state << 5;
state
};
for _ in 0..20_000 {
let len = (next() % 8) as usize + 1;
let mut buf = Vec::with_capacity(len);
for _ in 0..len {
let r = next();
buf.push(match r % 4 {
0 => (r >> 8) as u8,
1 => 0x80 | ((r >> 8) as u8 & 0x3f),
2 => 0xC0 | ((r >> 8) as u8 & 0x1f),
_ => 0xE0 | ((r >> 8) as u8 & 0x0f),
});
}
let ours = Utf8Validator::validate(&buf);
let theirs = core::str::from_utf8(&buf).is_ok();
assert_eq!(ours, theirs, "disagreement on {buf:02x?}");
}
}
#[test]
fn split_messages_keep_their_state() {
let text = "日本語テキスト🚀 end";
let bytes = text.as_bytes();
for split in 0..=bytes.len() {
let mut v = Utf8Validator::new();
let _ = v.feed(&bytes[..split]);
let _ = v.feed(&bytes[split..]);
assert!(v.is_complete(), "split at {split} lost state");
}
let emoji = "🚀".as_bytes();
let mut v = Utf8Validator::new();
assert!(v.feed(&emoji[..2]).is_ok());
assert!(!v.is_complete());
assert!(v.feed(&emoji[2..]).is_ok());
assert!(v.is_complete());
let mut v = Utf8Validator::new();
assert!(v.feed(&emoji[..2]).is_ok());
v.reset();
assert!(v.is_complete());
assert!(v.feed(b"ok").is_ok());
}
#[test]
fn error_reports_the_offending_offset() {
let mut v = Utf8Validator::new();
let e = v.feed(b"abc\xFF").unwrap_err();
assert_eq!(e.offset, 3);
assert!(!e.reason.is_empty());
}
}