use crate::atomic::{OncePtr, PyAtomic, Radium};
use crate::format::CharLen;
use crate::wtf8::{CodePoint, Wtf8, Wtf8Buf};
use crate::wtf8_index::Wtf8Index;
use alloc::borrow::Cow;
use ascii::{AsciiChar, AsciiStr, AsciiString};
use core::fmt;
use core::ops::{Bound, RangeBounds};
use core::sync::atomic::Ordering::Relaxed;
#[allow(non_camel_case_types)]
pub type wchar_t = cfg_select! {
target_arch = "wasm32" => u32,
_ => libc::wchar_t,
};
#[derive(Debug, Copy, Clone, PartialEq, Eq, PartialOrd, Ord)]
pub enum StrKind {
Ascii,
Utf8,
Wtf8,
}
impl core::ops::BitOr for StrKind {
type Output = Self;
fn bitor(self, other: Self) -> Self {
match (self, other) {
(Self::Wtf8, _) | (_, Self::Wtf8) => Self::Wtf8,
(Self::Utf8, _) | (_, Self::Utf8) => Self::Utf8,
(Self::Ascii, Self::Ascii) => Self::Ascii,
}
}
}
impl StrKind {
#[must_use]
pub const fn is_ascii(&self) -> bool {
matches!(self, Self::Ascii)
}
#[must_use]
pub const fn is_utf8(&self) -> bool {
matches!(self, Self::Ascii | Self::Utf8)
}
#[inline(always)]
#[must_use]
pub fn can_encode(&self, code: CodePoint) -> bool {
match self {
Self::Ascii => code.is_ascii(),
Self::Utf8 => code.to_char().is_some(),
Self::Wtf8 => true,
}
}
}
pub trait DeduceStrKind {
fn str_kind(&self) -> StrKind;
}
impl DeduceStrKind for str {
fn str_kind(&self) -> StrKind {
if self.is_ascii() {
StrKind::Ascii
} else {
StrKind::Utf8
}
}
}
impl DeduceStrKind for Wtf8 {
fn str_kind(&self) -> StrKind {
if self.is_ascii() {
StrKind::Ascii
} else if self.is_utf8() {
StrKind::Utf8
} else {
StrKind::Wtf8
}
}
}
impl DeduceStrKind for String {
fn str_kind(&self) -> StrKind {
(**self).str_kind()
}
}
impl DeduceStrKind for Wtf8Buf {
fn str_kind(&self) -> StrKind {
(**self).str_kind()
}
}
impl<T: DeduceStrKind + ?Sized> DeduceStrKind for &T {
fn str_kind(&self) -> StrKind {
(**self).str_kind()
}
}
impl<T: DeduceStrKind + ?Sized> DeduceStrKind for Box<T> {
fn str_kind(&self) -> StrKind {
(**self).str_kind()
}
}
#[derive(Debug)]
pub enum PyKindStr<'a> {
Ascii(&'a AsciiStr),
Utf8(&'a str),
Wtf8(&'a Wtf8),
}
const MAX_WALK_TO_INDEX: usize = 4;
#[derive(Debug, Clone)]
pub struct StrData {
data: Box<Wtf8>,
kind: StrKind,
len: StrLen,
index: Wtf8IndexSlot,
}
#[derive(Default)]
struct Wtf8IndexSlot(OncePtr<Wtf8Index>);
impl Wtf8IndexSlot {
#[inline(always)]
fn new() -> Self {
Self(OncePtr::new())
}
#[inline]
fn get_or_build(&self, data: &Wtf8, char_len: usize) -> &Wtf8Index {
let index = self
.0
.get_or_init(|| Box::new(Wtf8Index::new(data, char_len)));
unsafe { index.as_ref() }
}
}
impl fmt::Debug for Wtf8IndexSlot {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self.0.get() {
Some(_) => f.write_str("<built>"),
None => f.write_str("<unbuilt>"),
}
}
}
impl Clone for Wtf8IndexSlot {
fn clone(&self) -> Self {
Self::new()
}
}
impl Drop for Wtf8IndexSlot {
fn drop(&mut self) {
if let Some(index) = self.0.get() {
drop(unsafe { Box::from_raw(index.as_ptr()) });
}
}
}
struct StrLen(PyAtomic<usize>);
impl From<usize> for StrLen {
#[inline(always)]
fn from(value: usize) -> Self {
Self(Radium::new(value))
}
}
impl fmt::Debug for StrLen {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
let len = self.0.load(Relaxed);
if len == usize::MAX {
f.write_str("<uncomputed>")
} else {
len.fmt(f)
}
}
}
impl StrLen {
#[inline(always)]
fn zero() -> Self {
0usize.into()
}
#[inline(always)]
fn uncomputed() -> Self {
usize::MAX.into()
}
}
impl Clone for StrLen {
fn clone(&self) -> Self {
Self(self.0.load(Relaxed).into())
}
}
impl Default for StrData {
fn default() -> Self {
Self {
data: <Box<Wtf8>>::default(),
kind: StrKind::Ascii,
len: StrLen::zero(),
index: Wtf8IndexSlot::new(),
}
}
}
impl From<Box<Wtf8>> for StrData {
fn from(value: Box<Wtf8>) -> Self {
let kind = value.str_kind();
unsafe { Self::new_str_unchecked(value, kind) }
}
}
impl From<Box<str>> for StrData {
#[inline]
fn from(value: Box<str>) -> Self {
let kind = value.str_kind();
unsafe { Self::new_str_unchecked(value.into(), kind) }
}
}
impl From<Box<AsciiStr>> for StrData {
#[inline]
fn from(value: Box<AsciiStr>) -> Self {
Self {
len: value.len().into(),
data: value.into(),
kind: StrKind::Ascii,
index: Wtf8IndexSlot::new(),
}
}
}
impl From<AsciiChar> for StrData {
fn from(ch: AsciiChar) -> Self {
AsciiString::from(ch).into_boxed_ascii_str().into()
}
}
impl From<char> for StrData {
fn from(ch: char) -> Self {
if let Ok(ch) = ascii::AsciiChar::from_ascii(ch) {
ch.into()
} else {
Self {
data: ch.to_string().into(),
kind: StrKind::Utf8,
len: 1.into(),
index: Wtf8IndexSlot::new(),
}
}
}
}
impl From<CodePoint> for StrData {
fn from(ch: CodePoint) -> Self {
if let Some(ch) = ch.to_char() {
ch.into()
} else {
Self {
data: Wtf8Buf::from(ch).into(),
kind: StrKind::Wtf8,
len: 1.into(),
index: Wtf8IndexSlot::new(),
}
}
}
}
impl StrData {
#[must_use]
pub unsafe fn new_str_unchecked(data: Box<Wtf8>, kind: StrKind) -> Self {
let len = match kind {
StrKind::Ascii => data.len().into(),
_ => StrLen::uncomputed(),
};
Self {
data,
kind,
len,
index: Wtf8IndexSlot::new(),
}
}
#[must_use]
pub unsafe fn new_with_char_len(data: Box<Wtf8>, kind: StrKind, char_len: usize) -> Self {
Self {
data,
kind,
len: char_len.into(),
index: Wtf8IndexSlot::new(),
}
}
#[inline]
pub const fn as_wtf8(&self) -> &Wtf8 {
&self.data
}
#[inline]
pub fn as_str(&self) -> Option<&str> {
self.kind
.is_utf8()
.then(|| unsafe { core::str::from_utf8_unchecked(self.data.as_bytes()) })
}
pub fn as_ascii(&self) -> Option<&AsciiStr> {
self.kind
.is_ascii()
.then(|| unsafe { AsciiStr::from_ascii_unchecked(self.data.as_bytes()) })
}
pub const fn kind(&self) -> StrKind {
self.kind
}
#[inline]
pub fn as_str_kind(&self) -> PyKindStr<'_> {
match self.kind {
StrKind::Ascii => {
PyKindStr::Ascii(unsafe { AsciiStr::from_ascii_unchecked(self.data.as_bytes()) })
}
StrKind::Utf8 => {
PyKindStr::Utf8(unsafe { core::str::from_utf8_unchecked(self.data.as_bytes()) })
}
StrKind::Wtf8 => PyKindStr::Wtf8(&self.data),
}
}
#[inline]
pub fn len(&self) -> usize {
self.data.len()
}
pub fn is_empty(&self) -> bool {
self.data.is_empty()
}
#[inline]
pub fn char_len(&self) -> usize {
match self.len.0.load(Relaxed) {
usize::MAX => self._compute_char_len(),
len => len,
}
}
#[cold]
fn _compute_char_len(&self) -> usize {
let len = if let Some(s) = self.as_str() {
s.chars().count()
} else {
self.data.code_points().count()
};
self.len.0.store(len, Relaxed);
len
}
pub fn char_index_to_byte(&self, index: usize) -> usize {
if self.kind.is_ascii() {
return index.min(self.data.len());
}
let char_len = self.char_len();
if index >= char_len {
return self.data.len();
}
self.index
.get_or_build(&self.data, char_len)
.byte_offset(&self.data, index)
}
fn char_index_to_byte_once(&self, index: usize) -> usize {
if index <= MAX_WALK_TO_INDEX {
return self
.data
.code_point_indices()
.nth(index)
.map_or(self.data.len(), |(byte, _)| byte);
}
let from_end = self.char_len() - index;
if from_end <= MAX_WALK_TO_INDEX {
return self
.data
.code_point_indices()
.nth_back(from_end - 1)
.map_or(self.data.len(), |(byte, _)| byte);
}
self.char_index_to_byte(index)
}
#[must_use]
pub fn char_range_to_bytes(&self, range: core::ops::Range<usize>) -> core::ops::Range<usize> {
if self.kind.is_ascii() {
return range;
}
let from_end = self.char_len() - range.end;
if range.start <= MAX_WALK_TO_INDEX && from_end <= MAX_WALK_TO_INDEX {
let start = self
.data
.code_point_indices()
.nth(range.start)
.map_or(self.data.len(), |(byte, _)| byte);
let end = match from_end {
0 => self.data.len(),
n => self
.data
.code_point_indices()
.nth_back(n - 1)
.map_or(self.data.len(), |(byte, _)| byte),
};
return start..end;
}
self.char_index_to_byte(range.start)..self.char_index_to_byte(range.end)
}
pub fn byte_to_char_index(&self, bytepos: usize) -> usize {
if self.kind.is_ascii() {
return bytepos;
}
let char_len = self.char_len();
self.index
.get_or_build(&self.data, char_len)
.char_index_at_byte(&self.data, bytepos, char_len)
}
pub fn nth_char(&self, index: usize) -> CodePoint {
match self.as_str_kind() {
PyKindStr::Ascii(s) => s[index].into(),
_ => self.data[self.char_index_to_byte_once(index)..]
.code_points()
.next()
.unwrap(),
}
}
}
impl core::fmt::Display for StrData {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
self.data.fmt(f)
}
}
impl CharLen for StrData {
fn char_len(&self) -> usize {
self.char_len()
}
}
pub fn try_get_chars(s: &str, range: impl RangeBounds<usize>) -> Option<&str> {
let mut chars = s.chars();
let start = match range.start_bound() {
Bound::Included(&i) => i,
Bound::Excluded(&i) => i + 1,
Bound::Unbounded => 0,
};
for _ in 0..start {
chars.next()?;
}
let s = chars.as_str();
let range_len = match range.end_bound() {
Bound::Included(&i) => i + 1 - start,
Bound::Excluded(&i) => i - start,
Bound::Unbounded => return Some(s),
};
char_range_end(s, range_len).map(|end| &s[..end])
}
pub fn get_chars(s: &str, range: impl RangeBounds<usize>) -> &str {
try_get_chars(s, range).unwrap()
}
#[inline]
#[must_use]
pub fn char_range_end(s: &str, n_chars: usize) -> Option<usize> {
let i = match n_chars.checked_sub(1) {
Some(last_char_index) => {
let (index, c) = s.char_indices().nth(last_char_index)?;
index + c.len_utf8()
}
None => 0,
};
Some(i)
}
pub fn try_get_codepoints(w: &Wtf8, range: impl RangeBounds<usize>) -> Option<&Wtf8> {
let mut chars = w.code_points();
let start = match range.start_bound() {
Bound::Included(&i) => i,
Bound::Excluded(&i) => i + 1,
Bound::Unbounded => 0,
};
for _ in 0..start {
chars.next()?;
}
let s = chars.as_wtf8();
let range_len = match range.end_bound() {
Bound::Included(&i) => i + 1 - start,
Bound::Excluded(&i) => i - start,
Bound::Unbounded => return Some(s),
};
codepoint_range_end(s, range_len).map(|end| &s[..end])
}
pub fn get_codepoints(w: &Wtf8, range: impl RangeBounds<usize>) -> &Wtf8 {
try_get_codepoints(w, range).unwrap()
}
#[inline]
#[must_use]
pub fn codepoint_range_end(s: &Wtf8, n_chars: usize) -> Option<usize> {
let i = match n_chars.checked_sub(1) {
Some(last_char_index) => {
let (index, c) = s.code_point_indices().nth(last_char_index)?;
index + c.len_wtf8()
}
None => 0,
};
Some(i)
}
#[must_use]
pub fn zfill(bytes: &[u8], width: usize) -> Option<Vec<u8>> {
if width <= bytes.len() {
return Some(bytes.to_vec());
}
let (sign, s) = match bytes.first() {
Some(_sign @ (b'+' | b'-')) => (unsafe { bytes.get_unchecked(..1) }, &bytes[1..]),
_ => (&b""[..], bytes),
};
let mut filled = Vec::new();
filled.try_reserve_exact(width).ok()?;
filled.extend_from_slice(sign);
filled.extend(core::iter::repeat_n(b'0', width - bytes.len()));
filled.extend_from_slice(s);
Some(filled)
}
#[must_use]
pub fn to_ascii(value: &Wtf8) -> AsciiString {
let mut ascii = Vec::new();
for cp in value.code_points() {
if cp.is_ascii() {
ascii.push(cp.to_u32() as u8);
} else {
let c = cp.to_u32();
let hex = if c < 0x100 {
format!("\\x{c:02x}")
} else if c < 0x10000 {
format!("\\u{c:04x}")
} else {
format!("\\U{c:08x}")
};
ascii.append(&mut hex.into_bytes());
}
}
unsafe { AsciiString::from_ascii_unchecked(ascii) }
}
#[derive(Clone, Copy)]
pub struct UnicodeEscapeCodepoint(pub CodePoint);
impl fmt::Display for UnicodeEscapeCodepoint {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let c = self.0.to_u32();
if c >= 0x10000 {
write!(f, "\\U{c:08x}")
} else if c >= 0x100 {
write!(f, "\\u{c:04x}")
} else {
write!(f, "\\x{c:02x}")
}
}
}
pub mod levenshtein {
pub const MOVE_COST: usize = 2;
const CASE_COST: usize = 1;
const MAX_STRING_SIZE: usize = 40;
const fn substitution_cost(mut a: u8, mut b: u8) -> usize {
if (a & 31) != (b & 31) {
return MOVE_COST;
}
if a == b {
return 0;
}
if a.is_ascii_uppercase() {
a += b'a' - b'A';
}
if b.is_ascii_uppercase() {
b += b'a' - b'A';
}
if a == b { CASE_COST } else { MOVE_COST }
}
#[must_use]
pub fn levenshtein_distance(a: &[u8], b: &[u8], max_cost: usize) -> usize {
if a == b {
return 0;
}
let (mut a_bytes, mut b_bytes) = (a, b);
let (mut a_begin, mut a_end) = (0usize, a.len());
let (mut b_begin, mut b_end) = (0usize, b.len());
while a_end > 0 && b_end > 0 && (a_bytes[a_begin] == b_bytes[b_begin]) {
a_begin += 1;
b_begin += 1;
a_end -= 1;
b_end -= 1;
}
while a_end > 0
&& b_end > 0
&& (a_bytes[a_begin + a_end - 1] == b_bytes[b_begin + b_end - 1])
{
a_end -= 1;
b_end -= 1;
}
if a_end == 0 || b_end == 0 {
return (a_end + b_end) * MOVE_COST;
}
if a_end > MAX_STRING_SIZE || b_end > MAX_STRING_SIZE {
return max_cost + 1;
}
if b_end < a_end {
core::mem::swap(&mut a_bytes, &mut b_bytes);
core::mem::swap(&mut a_begin, &mut b_begin);
core::mem::swap(&mut a_end, &mut b_end);
}
if (b_end - a_end) * MOVE_COST > max_cost {
return max_cost + 1;
}
let mut buffer = [0usize; MAX_STRING_SIZE];
for (i, x) in buffer.iter_mut().take(a_end).enumerate() {
*x = (i + 1) * MOVE_COST;
}
let mut result = 0usize;
for (b_index, b_code) in b_bytes[b_begin..(b_begin + b_end)].iter().enumerate() {
result = b_index * MOVE_COST;
let mut distance = result;
let mut minimum = usize::MAX;
for (a_index, a_code) in a_bytes[a_begin..(a_begin + a_end)].iter().enumerate() {
let substitute = distance + substitution_cost(*b_code, *a_code);
distance = buffer[a_index];
let insert_delete = usize::min(result, distance) + MOVE_COST;
result = usize::min(insert_delete, substitute);
buffer[a_index] = result;
if result < minimum {
minimum = result;
}
}
if minimum > max_cost {
return max_cost + 1;
}
}
result
}
}
#[must_use]
pub fn expandtabs(input: &Wtf8, tab_size: usize) -> Wtf8Buf {
if tab_size == 0 {
return input.code_points().filter(|ch| *ch != '\t').collect();
}
let tab_stop = tab_size;
let mut expanded_str = Wtf8Buf::with_capacity(input.len());
let mut tab_size = tab_stop;
let mut col_count = 0usize;
for ch in input.code_points() {
if ch == '\t' {
let num_spaces = tab_size - col_count;
col_count += num_spaces;
expanded_str.push_str(&" ".repeat(num_spaces));
} else {
expanded_str.push(ch);
if ch == '\r' || ch == '\n' {
col_count = 0;
tab_size = 0;
} else {
col_count += 1;
}
}
if col_count >= tab_size {
tab_size += tab_stop;
}
}
expanded_str
}
#[macro_export]
macro_rules! ascii {
($x:expr $(,)?) => {{
let s = const {
let s: &str = $x;
assert!(s.is_ascii(), "ascii!() argument is not an ascii string");
s
};
unsafe { $crate::vendored::ascii::AsciiStr::from_ascii_unchecked(s.as_bytes()) }
}};
}
pub use ascii;
#[must_use]
pub fn char_to_decimal(ch: char) -> Option<u8> {
let value = rustpython_unicode::Ucd::new(true).decimal(CodePoint::from(ch))?;
u8::try_from(value).ok()
}
#[must_use]
pub fn transform_decimal_and_space_to_ascii(s: &str) -> Cow<'_, str> {
if s.is_ascii() {
return Cow::Borrowed(s);
}
let mut out = String::with_capacity(s.len());
for c in s.chars() {
if (c as u32) < 127 {
out.push(c);
} else if c.is_whitespace() {
out.push(' ');
} else if let Some(n) = char_to_decimal(c) {
out.push(char::from_digit(n.into(), 10).unwrap());
} else {
out.push('?');
break;
}
}
debug_assert!(out.is_ascii());
Cow::Owned(out)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn char_to_decimal_values() {
assert_eq!(char_to_decimal('7'), Some(7));
assert_eq!(char_to_decimal('٣'), Some(3)); assert_eq!(char_to_decimal('𐴵'), Some(5)); assert_eq!(char_to_decimal('🯵'), Some(5)); assert_eq!(char_to_decimal('½'), None);
assert_eq!(char_to_decimal('가'), None);
}
#[test]
fn transform_decimal_and_space() {
assert!(matches!(
transform_decimal_and_space_to_ascii("123"),
Cow::Borrowed("123")
));
assert_eq!(transform_decimal_and_space_to_ascii("١٢٣"), "123");
assert_eq!(transform_decimal_and_space_to_ascii("12३"), "123");
assert_eq!(transform_decimal_and_space_to_ascii("1٢3"), "123");
assert_eq!(transform_decimal_and_space_to_ascii("\u{3000}٣"), " 3");
assert_eq!(transform_decimal_and_space_to_ascii("0x١f"), "0x1f");
assert_eq!(transform_decimal_and_space_to_ascii("-١_٢"), "-1_2");
assert_eq!(transform_decimal_and_space_to_ascii("½가"), "?");
assert_eq!(transform_decimal_and_space_to_ascii("١٢가٣"), "12?");
assert_eq!(transform_decimal_and_space_to_ascii("١\u{7f}"), "1?");
}
#[test]
fn get_chars_basic() {
let s = "0123456789";
assert_eq!(get_chars(s, 3..7), "3456");
assert_eq!(get_chars(s, 3..7), &s[3..7]);
let s = "0유니코드 문자열9";
assert_eq!(get_chars(s, 3..7), "코드 문");
let s = "0😀😃😄😁😆😅😂🤣9";
assert_eq!(get_chars(s, 3..7), "😄😁😆😅");
}
fn expandtabs(input: &str, tab_size: usize) -> Wtf8Buf {
super::expandtabs(Wtf8::new(input), tab_size)
}
#[test]
fn expandtabs_with_zero_tab_size_drops_tabs() {
assert_eq!(expandtabs("a\tb", 0), Wtf8Buf::from("ab"));
assert_eq!(expandtabs("ab\tcd\tef", 0), Wtf8Buf::from("abcdef"));
assert_eq!(expandtabs("a\nb\tc", 0), Wtf8Buf::from("a\nbc"));
assert_eq!(expandtabs("á\tb", 0), Wtf8Buf::from("áb"));
assert_eq!(expandtabs("\ta", 0), Wtf8Buf::from("a"));
assert_eq!(expandtabs("\t", 0), Wtf8Buf::from(""));
assert_eq!(expandtabs("", 0), Wtf8Buf::from(""));
assert_eq!(expandtabs("no tabs", 0), Wtf8Buf::from("no tabs"));
}
#[test]
fn expandtabs_with_a_real_tab_size_is_unchanged() {
assert_eq!(expandtabs("a\tb", 8), Wtf8Buf::from("a b"));
assert_eq!(expandtabs("a\tb", 1), Wtf8Buf::from("a b"));
assert_eq!(expandtabs("abcd\te", 4), Wtf8Buf::from("abcd e"));
assert_eq!(expandtabs("a\nb\tc", 4), Wtf8Buf::from("a\nb c"));
assert_eq!(expandtabs("\ta", 4), Wtf8Buf::from(" a"));
}
}