use std::rc::Rc;
use ahash::RandomState;
use crate::{
args::{ArgValues, FromArgs},
builtins::Builtins,
bytecode::{CallResult, VM},
defer_drop, defer_drop_mut,
exception_private::{ExcType, RunResult},
heap::{ContainsHeap, DropWithHeap, Heap, HeapData, HeapId},
intern::StaticStrings,
modules::ModuleFunctions,
resource::{ResourceError, ResourceTracker},
types::{
BoundedCompileError, Module, RePattern, Type,
re_pattern::{extract_count, extract_maxsplit},
str::allocate_string,
},
value::Value,
};
pub(crate) const NOFLAG: u16 = 0;
pub(crate) const IGNORECASE: u16 = 2;
pub(crate) const MULTILINE: u16 = 8;
pub(crate) const DOTALL: u16 = 16;
pub(crate) const ASCII: u16 = 256;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, strum::Display, serde::Serialize, serde::Deserialize)]
#[strum(serialize_all = "lowercase")]
pub(crate) enum ReFunctions {
Compile,
Search,
Match,
Fullmatch,
Findall,
Sub,
Split,
Finditer,
Escape,
}
pub fn create_module(vm: &mut VM<'_, impl ResourceTracker>) -> Result<HeapId, ResourceError> {
let mut module = Module::new(StaticStrings::Re);
module.set_attr(
StaticStrings::Compile,
Value::ModuleFunction(ModuleFunctions::Re(ReFunctions::Compile)),
vm,
);
module.set_attr(
StaticStrings::Search,
Value::ModuleFunction(ModuleFunctions::Re(ReFunctions::Search)),
vm,
);
module.set_attr(
StaticStrings::Match,
Value::ModuleFunction(ModuleFunctions::Re(ReFunctions::Match)),
vm,
);
module.set_attr(
StaticStrings::Fullmatch,
Value::ModuleFunction(ModuleFunctions::Re(ReFunctions::Fullmatch)),
vm,
);
module.set_attr(
StaticStrings::Findall,
Value::ModuleFunction(ModuleFunctions::Re(ReFunctions::Findall)),
vm,
);
module.set_attr(
StaticStrings::Sub,
Value::ModuleFunction(ModuleFunctions::Re(ReFunctions::Sub)),
vm,
);
module.set_attr(
StaticStrings::Split,
Value::ModuleFunction(ModuleFunctions::Re(ReFunctions::Split)),
vm,
);
module.set_attr(
StaticStrings::Finditer,
Value::ModuleFunction(ModuleFunctions::Re(ReFunctions::Finditer)),
vm,
);
module.set_attr(
StaticStrings::Escape,
Value::ModuleFunction(ModuleFunctions::Re(ReFunctions::Escape)),
vm,
);
module.set_attr(StaticStrings::NoFlag, Value::Int(i64::from(NOFLAG)), vm);
module.set_attr(StaticStrings::Ignorecase, Value::Int(i64::from(IGNORECASE)), vm);
module.set_attr(StaticStrings::I, Value::Int(i64::from(IGNORECASE)), vm);
module.set_attr(StaticStrings::MultilineFlag, Value::Int(i64::from(MULTILINE)), vm);
module.set_attr(StaticStrings::M, Value::Int(i64::from(MULTILINE)), vm);
module.set_attr(StaticStrings::DotallFlag, Value::Int(i64::from(DOTALL)), vm);
module.set_attr(StaticStrings::S, Value::Int(i64::from(DOTALL)), vm);
module.set_attr(StaticStrings::AsciiFlag, Value::Int(i64::from(ASCII)), vm);
module.set_attr(StaticStrings::A, Value::Int(i64::from(ASCII)), vm);
module.set_attr(
StaticStrings::PatternError,
Value::Builtin(Builtins::ExcType(ExcType::RePatternError)),
vm,
);
module.set_attr(
StaticStrings::Error,
Value::Builtin(Builtins::ExcType(ExcType::RePatternError)),
vm,
);
module.set_attr(
StaticStrings::PatternClass,
Value::Builtin(Builtins::Type(Type::RePattern)),
vm,
);
module.set_attr(
StaticStrings::MatchClass,
Value::Builtin(Builtins::Type(Type::ReMatch)),
vm,
);
vm.heap.allocate(HeapData::Module(module))
}
pub(super) fn call(
vm: &mut VM<'_, impl ResourceTracker>,
function: ReFunctions,
args: ArgValues,
) -> RunResult<CallResult> {
match function {
ReFunctions::Compile => call_compile(vm, args).map(CallResult::Value),
ReFunctions::Search => call_search(vm, args).map(CallResult::Value),
ReFunctions::Match => call_match(vm, args).map(CallResult::Value),
ReFunctions::Fullmatch => call_fullmatch(vm, args).map(CallResult::Value),
ReFunctions::Findall => call_findall(vm, args).map(CallResult::Value),
ReFunctions::Sub => call_sub(vm, args).map(CallResult::Value),
ReFunctions::Split => call_split(vm, args).map(CallResult::Value),
ReFunctions::Finditer => call_finditer(vm, args).map(CallResult::Value),
ReFunctions::Escape => call_escape(vm, args).map(CallResult::Value),
}
}
fn call_compile(vm: &mut VM<'_, impl ResourceTracker>, args: ArgValues) -> RunResult<Value> {
let ReCompileArgs { pattern, flags } = ReCompileArgs::from_args(args, vm)?;
match resolve_pattern(pattern, flags, vm)? {
ResolvedPattern::Cached(compiled) => Ok(Value::Ref(
vm.heap.allocate(HeapData::RePattern(Box::new((*compiled).clone())))?,
)),
ResolvedPattern::Heap(value) => Ok(value),
}
}
fn call_search(vm: &mut VM<'_, impl ResourceTracker>, args: ArgValues) -> RunResult<Value> {
let ReSearchArgs { pattern, string, flags } = ReSearchArgs::from_args(args, vm)?;
defer_drop!(string, vm);
let resolved = resolve_pattern(pattern, flags, vm)?;
defer_drop!(resolved, vm);
resolved.get(vm.heap).search(string, subject_str(string, vm)?, vm.heap)
}
fn call_match(vm: &mut VM<'_, impl ResourceTracker>, args: ArgValues) -> RunResult<Value> {
let ReMatchArgs { pattern, string, flags } = ReMatchArgs::from_args(args, vm)?;
defer_drop!(string, vm);
let resolved = resolve_pattern(pattern, flags, vm)?;
defer_drop!(resolved, vm);
resolved
.get(vm.heap)
.match_start(string, subject_str(string, vm)?, vm.heap)
}
fn call_fullmatch(vm: &mut VM<'_, impl ResourceTracker>, args: ArgValues) -> RunResult<Value> {
let ReFullmatchArgs { pattern, string, flags } = ReFullmatchArgs::from_args(args, vm)?;
defer_drop!(string, vm);
let resolved = resolve_pattern(pattern, flags, vm)?;
defer_drop!(resolved, vm);
resolved
.get(vm.heap)
.fullmatch(string, subject_str(string, vm)?, vm.heap)
}
fn call_findall(vm: &mut VM<'_, impl ResourceTracker>, args: ArgValues) -> RunResult<Value> {
let ReFindallArgs { pattern, string, flags } = ReFindallArgs::from_args(args, vm)?;
defer_drop!(string, vm);
let resolved = resolve_pattern(pattern, flags, vm)?;
defer_drop!(resolved, vm);
resolved.get(vm.heap).findall(subject_str(string, vm)?, vm.heap)
}
fn call_sub(vm: &mut VM<'_, impl ResourceTracker>, args: ArgValues) -> RunResult<Value> {
let ReSubArgs {
pattern,
repl,
string,
count,
flags,
} = ReSubArgs::from_args(args, vm)?;
defer_drop!(string, vm);
defer_drop!(repl, vm);
defer_drop_mut!(count, vm);
let resolved = resolve_pattern(pattern, flags, vm)?;
defer_drop!(resolved, vm);
let count = extract_count(count.take(), vm)?;
if !repl.is_str(vm.heap) {
return Err(ExcType::type_error(
"callable replacement is not yet supported in re.sub()",
));
}
let Some(count) = count else {
let _ = subject_str(string, vm)?;
return Ok(string.clone_with_heap(vm.heap));
};
resolved
.get(vm.heap)
.sub(repl.to_str(vm)?, subject_str(string, vm)?, count, vm.heap)
}
fn call_split(vm: &mut VM<'_, impl ResourceTracker>, args: ArgValues) -> RunResult<Value> {
let ReSplitArgs {
pattern,
string,
maxsplit,
flags,
} = ReSplitArgs::from_args(args, vm)?;
defer_drop!(string, vm);
defer_drop_mut!(maxsplit, vm);
let resolved = resolve_pattern(pattern, flags, vm)?;
defer_drop!(resolved, vm);
let maxsplit = extract_maxsplit(maxsplit.take(), vm)?;
resolved.get(vm.heap).split(subject_str(string, vm)?, maxsplit, vm.heap)
}
#[derive(FromArgs)]
#[from_args(name = "sub", style = def)]
struct ReSubArgs {
#[from_args(static_string = "PatternAttr")]
pattern: Value,
repl: Value,
#[from_args(static_string = "StringAttr")]
string: Value,
#[from_args(default)]
count: Option<Value>,
#[from_args(default = Value::Int(0))]
flags: Value,
}
#[derive(FromArgs)]
#[from_args(name = "split", style = def)]
struct ReSplitArgs {
#[from_args(static_string = "PatternAttr")]
pattern: Value,
#[from_args(static_string = "StringAttr")]
string: Value,
#[from_args(default)]
maxsplit: Option<Value>,
#[from_args(default = Value::Int(0))]
flags: Value,
}
#[derive(FromArgs)]
#[from_args(name = "compile", style = def)]
struct ReCompileArgs {
#[from_args(static_string = "PatternAttr")]
pattern: Value,
#[from_args(default = Value::Int(0))]
flags: Value,
}
#[derive(FromArgs)]
#[from_args(name = "escape", style = def)]
struct ReEscapeArgs {
#[from_args(static_string = "PatternAttr")]
pattern: Value,
}
#[derive(FromArgs)]
#[from_args(name = "search", style = def)]
struct ReSearchArgs {
#[from_args(static_string = "PatternAttr")]
pattern: Value,
#[from_args(static_string = "StringAttr")]
string: Value,
#[from_args(default = Value::Int(0))]
flags: Value,
}
#[derive(FromArgs)]
#[from_args(name = "match", style = def)]
struct ReMatchArgs {
#[from_args(static_string = "PatternAttr")]
pattern: Value,
#[from_args(static_string = "StringAttr")]
string: Value,
#[from_args(default = Value::Int(0))]
flags: Value,
}
#[derive(FromArgs)]
#[from_args(name = "fullmatch", style = def)]
struct ReFullmatchArgs {
#[from_args(static_string = "PatternAttr")]
pattern: Value,
#[from_args(static_string = "StringAttr")]
string: Value,
#[from_args(default = Value::Int(0))]
flags: Value,
}
#[derive(FromArgs)]
#[from_args(name = "findall", style = def)]
struct ReFindallArgs {
#[from_args(static_string = "PatternAttr")]
pattern: Value,
#[from_args(static_string = "StringAttr")]
string: Value,
#[from_args(default = Value::Int(0))]
flags: Value,
}
#[derive(FromArgs)]
#[from_args(name = "finditer", style = def)]
struct ReFinditerArgs {
#[from_args(static_string = "PatternAttr")]
pattern: Value,
#[from_args(static_string = "StringAttr")]
string: Value,
#[from_args(default = Value::Int(0))]
flags: Value,
}
fn subject_str<'a>(value: &'a Value, vm: &'a VM<'_, impl ResourceTracker>) -> RunResult<&'a str> {
if value.is_str(vm.heap) {
value.to_str(vm)
} else if value.py_type_heap(vm.heap) == Type::Bytes {
Err(ExcType::type_error(
"cannot use a string pattern on a bytes-like object",
))
} else {
Err(ExcType::type_error(format!(
"expected string or bytes-like object, got '{}'",
value.py_type_name(vm)
)))
}
}
pub(crate) enum PatternArg {
Str(String),
Compiled(Value),
}
impl PatternArg {
fn extract(value: Value, vm: &mut VM<'_, impl ResourceTracker>) -> RunResult<Self> {
if let Some(either) = value.as_either_str(vm.heap) {
let pattern = either.into_string(vm.interns);
value.drop_with_heap(vm);
Ok(Self::Str(pattern))
} else if value.py_type_heap(vm.heap) == Type::RePattern {
Ok(Self::Compiled(value))
} else {
value.drop_with_heap(vm);
Err(ExcType::type_error("first argument must be string or compiled pattern"))
}
}
}
impl DropWithHeap for PatternArg {
fn drop_with_heap<H: ContainsHeap>(self, heap: &mut H) {
if let Self::Compiled(value) = self {
value.drop_with_heap(heap);
}
}
}
#[derive(Debug, Clone, Copy, Default)]
pub(crate) struct ReFlags(u16);
impl ReFlags {
fn extract(value: Value, vm: &mut VM<'_, impl ResourceTracker>) -> RunResult<Self> {
let result = match value {
Value::Int(n) => u16::try_from(n)
.map(Self)
.map_err(|_| ExcType::type_error("flags must be a non-negative integer")),
Value::Bool(b) => Ok(Self(u16::from(b))),
_ => Err(ExcType::binary_type_error(
"&",
value.py_type_heap(vm.heap),
value.py_type_name(vm),
"int",
)),
};
value.drop_with_heap(vm);
result
}
}
impl ReFlags {
fn get(self) -> u16 {
self.0
}
}
enum ResolvedPattern {
Cached(Rc<RePattern>),
Heap(Value),
}
impl ResolvedPattern {
fn get<'a>(&'a self, heap: &'a Heap<impl ResourceTracker>) -> &'a RePattern {
match self {
Self::Cached(pattern) => pattern,
Self::Heap(value) => {
let Value::Ref(heap_id) = value else {
unreachable!("ResolvedPattern::Heap always holds a heap ref")
};
match heap.get(*heap_id) {
HeapData::RePattern(pattern) => pattern,
_ => unreachable!("ResolvedPattern::Heap always points at a re.Pattern"),
}
}
}
}
}
impl DropWithHeap for ResolvedPattern {
fn drop_with_heap<H: ContainsHeap>(self, heap: &mut H) {
if let Self::Heap(value) = self {
value.drop_with_heap(heap);
}
}
}
type ReCacheEntry = Option<(u64, Box<str>, u16, Rc<RePattern>)>;
const CACHE_CAPACITY: usize = 256;
const CACHED_SIZE_LIMIT: usize = 64 * 1024;
#[derive(Default)]
pub(crate) struct RePatternCache(Option<(Box<[ReCacheEntry]>, RandomState)>);
impl RePatternCache {
fn get_or_compile(&mut self, pattern: &str, flags: u16) -> RunResult<Rc<RePattern>> {
let (entries, hash_builder) = self
.0
.get_or_insert_with(|| (vec![None; CACHE_CAPACITY].into_boxed_slice(), RandomState::default()));
let hash = hash_builder.hash_one((pattern, flags));
let hash_index = usize::try_from(hash % CACHE_CAPACITY as u64).expect("index < CACHE_CAPACITY");
let mut empty_slot = None;
for index in hash_index..hash_index + 5 {
match entries.get(index) {
Some(Some((entry_hash, entry_pattern, entry_flags, compiled))) => {
if *entry_hash == hash && *entry_flags == flags && &**entry_pattern == pattern {
return Ok(Rc::clone(compiled));
}
}
Some(None) => {
empty_slot = Some(index);
break;
}
None => break,
}
}
let compiled = match RePattern::compile_bounded(pattern.to_owned(), flags, CACHED_SIZE_LIMIT) {
Ok(compiled) => compiled,
Err(BoundedCompileError::TooBig) => {
return Ok(Rc::new(RePattern::compile(pattern.to_owned(), flags)?));
}
Err(BoundedCompileError::Invalid(err)) => return Err(err),
};
let compiled = Rc::new(compiled);
let slot = empty_slot.unwrap_or(hash_index);
entries[slot] = Some((hash, Box::from(pattern), flags, Rc::clone(&compiled)));
Ok(compiled)
}
}
fn resolve_pattern(pattern: Value, flags: Value, vm: &mut VM<'_, impl ResourceTracker>) -> RunResult<ResolvedPattern> {
let pattern = match PatternArg::extract(pattern, vm) {
Ok(pattern) => pattern,
Err(e) => {
flags.drop_with_heap(vm);
return Err(e);
}
};
let flags = match ReFlags::extract(flags, vm) {
Ok(flags) => flags,
Err(e) => {
pattern.drop_with_heap(vm);
return Err(e);
}
};
match pattern {
PatternArg::Str(pattern) => Ok(ResolvedPattern::Cached(
vm.re_pattern_cache.get_or_compile(&pattern, flags.get())?,
)),
PatternArg::Compiled(value) => {
if flags.get() == 0 {
Ok(ResolvedPattern::Heap(value))
} else {
value.drop_with_heap(vm);
Err(ExcType::value_error(
"cannot process flags argument with a compiled pattern",
))
}
}
}
}
fn call_finditer(vm: &mut VM<'_, impl ResourceTracker>, args: ArgValues) -> RunResult<Value> {
let ReFinditerArgs { pattern, string, flags } = ReFinditerArgs::from_args(args, vm)?;
defer_drop!(string, vm);
let resolved = resolve_pattern(pattern, flags, vm)?;
defer_drop!(resolved, vm);
resolved
.get(vm.heap)
.finditer(string, subject_str(string, vm)?, vm.heap)
}
fn call_escape(vm: &mut VM<'_, impl ResourceTracker>, args: ArgValues) -> RunResult<Value> {
let ReEscapeArgs { pattern } = ReEscapeArgs::from_args(args, vm)?;
defer_drop!(pattern, vm);
let Ok(text) = pattern.to_str(vm) else {
let t = pattern.py_type_name(vm);
return Err(ExcType::type_error(format!(
"decoding to str: need a bytes-like object, {t} found"
)));
};
let mut result = String::with_capacity(text.len() * 2);
for c in text.chars() {
if should_escape(c) {
result.push('\\');
}
result.push(c);
}
Ok(allocate_string(result, vm.heap)?)
}
fn should_escape(c: char) -> bool {
matches!(
c,
'\t' | '\n'
| '\x0b'
| '\x0c'
| '\r'
| ' '
| '#'
| '$'
| '&'
| '('
| ')'
| '*'
| '+'
| '-'
| '.'
| '?'
| '['
| '\\'
| ']'
| '^'
| '{'
| '|'
| '}'
| '~'
)
}
#[cfg(test)]
mod tests {
use super::*;
fn occupied(cache: &RePatternCache) -> usize {
cache
.0
.as_ref()
.map_or(0, |(entries, _)| entries.iter().filter(|e| e.is_some()).count())
}
#[test]
fn hit_shares_one_entry() {
let mut cache = RePatternCache::default();
let first = cache.get_or_compile(r"\s+", 0).unwrap();
let second = cache.get_or_compile(r"\s+", 0).unwrap();
assert!(Rc::ptr_eq(&first, &second));
assert_eq!(occupied(&cache), 1);
}
#[test]
fn flags_key_the_entry() {
let mut cache = RePatternCache::default();
let plain = cache.get_or_compile("abc", 0).unwrap();
let ignorecase = cache.get_or_compile("abc", IGNORECASE).unwrap();
assert!(!Rc::ptr_eq(&plain, &ignorecase));
assert_eq!(occupied(&cache), 2);
}
#[test]
fn compile_error_caches_nothing() {
let mut cache = RePatternCache::default();
assert!(cache.get_or_compile("(", 0).is_err());
assert_eq!(occupied(&cache), 0);
}
const OVERSIZE_PATTERN: &str = "a{5000}";
#[test]
fn oversize_pattern_compiles_but_is_not_retained() {
let mut cache = RePatternCache::default();
let first = cache.get_or_compile(OVERSIZE_PATTERN, 0).unwrap();
let second = cache.get_or_compile(OVERSIZE_PATTERN, 0).unwrap();
assert!(!Rc::ptr_eq(&first, &second));
assert_eq!(occupied(&cache), 0);
}
#[test]
fn oversize_pattern_does_not_disturb_cached_entries() {
let mut cache = RePatternCache::default();
let small = cache.get_or_compile("abc", 0).unwrap();
cache.get_or_compile(OVERSIZE_PATTERN, 0).unwrap();
let again = cache.get_or_compile("abc", 0).unwrap();
assert!(Rc::ptr_eq(&small, &again));
assert_eq!(occupied(&cache), 1);
}
#[test]
fn occupancy_is_bounded_by_capacity() {
let mut cache = RePatternCache::default();
for i in 0..CACHE_CAPACITY * 4 {
cache.get_or_compile(&format!("pat{i}"), 0).unwrap();
}
assert!(occupied(&cache) <= CACHE_CAPACITY);
}
}