use std::{borrow::Cow, cell::OnceCell, cmp::Ordering, fmt::Write, iter, mem, str};
use fancy_regex::{CompileError, Error as RegexError, Regex, RegexBuilder};
use serde::{Deserialize, Deserializer, Serialize, Serializer, de};
use smallvec::SmallVec;
use crate::{
args::{ArgValues, FromArgs},
bytecode::{CallResult, VM},
defer_drop,
exception_private::{ExcType, RunError, RunResult},
heap::{Heap, HeapData, HeapId, HeapItem, HeapRead, HeapReadOutput},
intern::StaticStrings,
modules::re::{ASCII, DOTALL, IGNORECASE, MULTILINE},
resource::{ResourceTracker, check_estimated_size},
types::{
LazyHeapSet, List, PyTrait, ReMatch, Type, allocate_tuple,
str::{allocate_string, string_repr_fmt},
},
value::{EitherStr, Value},
};
#[derive(Debug, Clone)]
pub(crate) struct RePattern {
pattern: String,
flags: u16,
compiled: Regex,
compiled_match: OnceCell<Regex>,
compiled_fullmatch: OnceCell<Regex>,
delegate_size_limit: Option<usize>,
}
impl PartialEq for RePattern {
fn eq(&self, other: &Self) -> bool {
self.pattern == other.pattern && self.flags == other.flags
}
}
pub(crate) enum BoundedCompileError {
TooBig,
Invalid(RunError),
}
impl RePattern {
pub fn compile(pattern: String, flags: u16) -> RunResult<Self> {
Self::compile_inner(pattern, flags, None).map_err(ExcType::re_pattern_error)
}
pub(crate) fn compile_bounded(
pattern: String,
flags: u16,
delegate_size_limit: usize,
) -> Result<Self, BoundedCompileError> {
Self::compile_inner(pattern, flags, Some(delegate_size_limit)).map_err(|err| {
if is_size_limit_error(&err) {
BoundedCompileError::TooBig
} else {
BoundedCompileError::Invalid(ExcType::re_pattern_error(err))
}
})
}
fn compile_inner(pattern: String, flags: u16, delegate_size_limit: Option<usize>) -> Result<Self, RegexError> {
let compiled = compile_regex_limited(&pattern, flags, delegate_size_limit)?;
Ok(Self {
pattern,
flags,
compiled,
compiled_match: OnceCell::new(),
compiled_fullmatch: OnceCell::new(),
delegate_size_limit,
})
}
fn match_regex(&self) -> RunResult<&Regex> {
if let Some(regex) = self.compiled_match.get() {
return Ok(regex);
}
let compiled = compile_regex_limited(
&format!("\\A(?:{})", self.pattern),
self.flags,
self.delegate_size_limit,
)
.map_err(ExcType::re_pattern_error)?;
let _ = self.compiled_match.set(compiled);
Ok(self.compiled_match.get().expect("cell was just initialised"))
}
fn fullmatch_regex(&self) -> RunResult<&Regex> {
if let Some(regex) = self.compiled_fullmatch.get() {
return Ok(regex);
}
let compiled = compile_regex_limited(
&format!("\\A(?:{})\\z", self.pattern),
self.flags,
self.delegate_size_limit,
)
.map_err(ExcType::re_pattern_error)?;
let _ = self.compiled_fullmatch.set(compiled);
Ok(self.compiled_fullmatch.get().expect("cell was just initialised"))
}
fn build_match(
&self,
caps: &fancy_regex::Captures<'_>,
subject: &Value,
all_ascii: bool,
heap: &Heap<impl ResourceTracker>,
) -> RunResult<Value> {
let m = ReMatch::from_captures(caps, subject.clone_with_heap(heap), all_ascii, &self.compiled);
Ok(Value::Ref(heap.allocate(HeapData::ReMatch(m))?))
}
pub fn search(&self, subject: &Value, text: &str, heap: &Heap<impl ResourceTracker>) -> RunResult<Value> {
match self.compiled.captures(text) {
Ok(Some(caps)) => self.build_match(&caps, subject, text.is_ascii(), heap),
Ok(None) => Ok(Value::None),
Err(err) => Err(ExcType::re_pattern_error(err)),
}
}
pub fn match_start(&self, subject: &Value, text: &str, heap: &Heap<impl ResourceTracker>) -> RunResult<Value> {
match self.match_regex()?.captures(text) {
Ok(Some(caps)) => self.build_match(&caps, subject, text.is_ascii(), heap),
Ok(None) => Ok(Value::None),
Err(err) => Err(ExcType::re_pattern_error(err)),
}
}
pub fn fullmatch(&self, subject: &Value, text: &str, heap: &Heap<impl ResourceTracker>) -> RunResult<Value> {
match self.fullmatch_regex()?.captures(text) {
Ok(Some(caps)) => self.build_match(&caps, subject, text.is_ascii(), heap),
Ok(None) => Ok(Value::None),
Err(err) => Err(ExcType::re_pattern_error(err)),
}
}
pub fn findall(&self, text: &str, heap: &Heap<impl ResourceTracker>) -> RunResult<Value> {
let cap_count = self.compiled.captures_len();
let mut results = Vec::new();
match cap_count {
0 | 1 => {
for m in self.compiled.find_iter(text) {
let val = m.map_err(ExcType::re_pattern_error)?.as_str();
results.push(allocate_string(val, heap)?);
}
}
2 => {
for caps in self.compiled.captures_iter(text) {
let caps = caps.map_err(ExcType::re_pattern_error)?;
let val = caps.get(1).map_or("", |m| m.as_str());
results.push(allocate_string(val, heap)?);
}
}
_ => {
for caps in self.compiled.captures_iter(text) {
let caps = caps.map_err(ExcType::re_pattern_error)?;
let mut elements: SmallVec<[Value; 3]> = SmallVec::with_capacity(cap_count - 1);
for cap in caps.iter().skip(1) {
let val = cap.map_or("", |m| m.as_str());
elements.push(allocate_string(val, heap)?);
}
results.push(allocate_tuple(elements, heap)?);
}
}
}
let list = List::new(results);
Ok(Value::Ref(heap.allocate(HeapData::List(list))?))
}
pub fn sub(&self, repl: &str, text: &str, count: usize, heap: &Heap<impl ResourceTracker>) -> RunResult<Value> {
let rust_repl = translate_replacement(repl);
let effective_count = if count == 0 { usize::MAX } else { count };
let mut result = String::new();
let mut last_end = 0;
for caps in self.compiled.captures_iter(text).take(effective_count) {
let caps = caps.map_err(ExcType::re_pattern_error)?;
let m = caps.get(0).expect("capture group 0 always exists");
result.push_str(&text[last_end..m.start()]);
caps.expand(rust_repl.as_ref(), &mut result);
last_end = m.end();
check_estimated_size(result.len() + (text.len() - last_end), heap.tracker())?;
}
result.push_str(&text[last_end..]);
Ok(allocate_string(result, heap)?)
}
pub fn split(&self, text: &str, maxsplit: i64, heap: &Heap<impl ResourceTracker>) -> RunResult<Value> {
let pieces: Vec<&str> = match maxsplit.cmp(&0) {
Ordering::Less => vec![text],
Ordering::Equal => self
.compiled
.split(text)
.collect::<Result<Vec<_>, _>>()
.map_err(ExcType::re_pattern_error)?,
Ordering::Greater => {
let limit = usize::try_from(maxsplit).unwrap_or(usize::MAX).saturating_add(1);
self.compiled
.splitn(text, limit)
.collect::<Result<Vec<_>, _>>()
.map_err(ExcType::re_pattern_error)?
}
};
let mut results = Vec::with_capacity(pieces.len());
for piece in pieces {
results.push(allocate_string(piece, heap)?);
}
let list = List::new(results);
Ok(Value::Ref(heap.allocate(HeapData::List(list))?))
}
pub fn finditer(&self, subject: &Value, text: &str, heap: &Heap<impl ResourceTracker>) -> RunResult<Value> {
let all_ascii = text.is_ascii();
let mut results = Vec::new();
for caps in self.compiled.captures_iter(text) {
let caps = caps.map_err(ExcType::re_pattern_error)?;
results.push(self.build_match(&caps, subject, all_ascii, heap)?);
}
let list = List::new(results);
Ok(Value::Ref(heap.allocate(HeapData::List(list))?))
}
}
impl<'h> PyTrait<'h> for HeapRead<'h, RePattern> {
fn py_type(&self, _vm: &VM<'h, impl ResourceTracker>) -> Type {
Type::RePattern
}
fn py_len(&self, _vm: &VM<'h, impl ResourceTracker>) -> Option<usize> {
None
}
fn py_eq_impl(&self, other: &Value, vm: &mut VM<'h, impl ResourceTracker>) -> RunResult<Option<bool>> {
let Some(HeapReadOutput::RePattern(other)) = other.read_heap(vm) else {
return Ok(None);
};
Ok(Some(self.get(vm.heap) == other.get(vm.heap)))
}
fn py_bool(&self, _vm: &mut VM<'h, impl ResourceTracker>) -> bool {
true
}
fn py_repr_fmt(
&self,
f: &mut impl Write,
vm: &mut VM<'h, impl ResourceTracker>,
_heap_ids: &mut LazyHeapSet,
) -> RunResult<()> {
let this = self.get(vm.heap);
write!(f, "re.compile(")?;
string_repr_fmt(&this.pattern, f)?;
if this.flags != 0 {
let mut flag_parts = smallvec::SmallVec::<[&'static str; 4]>::new();
if this.flags & IGNORECASE != 0 {
flag_parts.push("re.IGNORECASE");
}
if this.flags & MULTILINE != 0 {
flag_parts.push("re.MULTILINE");
}
if this.flags & DOTALL != 0 {
flag_parts.push("re.DOTALL");
}
if this.flags & ASCII != 0 {
flag_parts.push("re.ASCII");
}
write!(f, ", {}", flag_parts.join("|"))?;
}
Ok(write!(f, ")")?)
}
fn py_getattr(&self, attr: &EitherStr, vm: &mut VM<'h, impl ResourceTracker>) -> RunResult<Option<CallResult>> {
match attr.static_string() {
Some(StaticStrings::PatternAttr) => {
let v = allocate_string(self.get(vm.heap).pattern.as_str(), vm.heap)?;
Ok(Some(CallResult::Value(v)))
}
Some(StaticStrings::Flags) => Ok(Some(CallResult::Value(Value::Int(i64::from(self.get(vm.heap).flags))))),
_ => Err(ExcType::attribute_error(Type::RePattern, attr.as_str(vm.interns))),
}
}
fn py_call_attr(
&mut self,
_self_id: HeapId,
vm: &mut VM<'h, impl ResourceTracker>,
attr: &EitherStr,
args: ArgValues,
) -> RunResult<CallResult> {
let result = match attr.static_string() {
Some(StaticStrings::Search) => {
let arg = args.get_one_arg("Pattern.search", vm.heap)?;
defer_drop!(arg, vm);
let text = arg.to_str(vm)?;
self.get(vm.heap).search(arg, text, vm.heap)
}
Some(StaticStrings::Match) => {
let arg = args.get_one_arg("Pattern.match", vm.heap)?;
defer_drop!(arg, vm);
let text = arg.to_str(vm)?;
self.get(vm.heap).match_start(arg, text, vm.heap)
}
Some(StaticStrings::Fullmatch) => {
let arg = args.get_one_arg("Pattern.fullmatch", vm.heap)?;
defer_drop!(arg, vm);
let text = arg.to_str(vm)?;
self.get(vm.heap).fullmatch(arg, text, vm.heap)
}
Some(StaticStrings::Findall) => {
let arg = args.get_one_arg("Pattern.findall", vm.heap)?;
defer_drop!(arg, vm);
let text = arg.to_str(vm)?;
self.get(vm.heap).findall(text, vm.heap)
}
Some(StaticStrings::Sub) => call_pattern_sub(self, args, vm),
Some(StaticStrings::Split) => call_pattern_split(self, args, vm),
Some(StaticStrings::Finditer) => {
let arg = args.get_one_arg("Pattern.finditer", vm.heap)?;
defer_drop!(arg, vm);
let text = arg.to_str(vm)?;
self.get(vm.heap).finditer(arg, text, vm.heap)
}
_ => {
return Err(ExcType::attribute_error(Type::RePattern, attr.as_str(vm.interns)));
}
}?;
Ok(CallResult::Value(result))
}
}
impl HeapItem for RePattern {
fn py_estimate_size(&self) -> usize {
mem::size_of::<Self>() + self.pattern.len()
}
fn py_dec_ref_ids(&mut self, _stack: &mut Vec<HeapId>) {
}
}
fn call_pattern_sub<'h>(
pattern: &HeapRead<'h, RePattern>,
args: ArgValues,
vm: &mut VM<'h, impl ResourceTracker>,
) -> RunResult<Value> {
let PatternSubArgs {
repl: repl_val,
string: string_val,
count: count_val,
} = PatternSubArgs::from_args(args, vm)?;
defer_drop!(repl_val, vm);
defer_drop!(string_val, vm);
let count = extract_count(count_val, vm)?;
if !repl_val.is_str(vm.heap) {
return Err(ExcType::type_error(
"callable replacement is not yet supported in re.sub()",
));
}
let Some(count) = count else {
let _ = string_val.to_str(vm)?;
return Ok(string_val.clone_with_heap(vm.heap));
};
let repl = repl_val.to_str(vm)?.to_owned();
let text = string_val.to_str(vm)?.to_owned();
pattern.get(vm.heap).sub(&repl, &text, count, vm.heap)
}
fn call_pattern_split<'h>(
pattern: &HeapRead<'h, RePattern>,
args: ArgValues,
vm: &mut VM<'h, impl ResourceTracker>,
) -> RunResult<Value> {
let PatternSplitArgs {
string: string_val,
maxsplit: maxsplit_val,
} = PatternSplitArgs::from_args(args, vm)?;
defer_drop!(string_val, vm);
let maxsplit = extract_maxsplit(maxsplit_val, vm)?;
let text = string_val.to_str(vm)?.to_owned();
pattern.get(vm.heap).split(&text, maxsplit, vm.heap)
}
#[derive(FromArgs)]
#[from_args(name = "sub", style = c_named, at_most_total)]
struct PatternSubArgs {
repl: Value,
#[from_args(static_string = "StringAttr")]
string: Value,
#[from_args(default)]
count: Option<Value>,
}
#[derive(FromArgs)]
#[from_args(name = "split", style = c_named, at_most_total)]
struct PatternSplitArgs {
#[from_args(static_string = "StringAttr")]
string: Value,
#[from_args(default)]
maxsplit: Option<Value>,
}
pub(crate) fn extract_maxsplit(val: Option<Value>, vm: &mut VM<'_, impl ResourceTracker>) -> RunResult<i64> {
match val {
None => Ok(0),
Some(Value::Int(n)) => Ok(n),
Some(Value::Bool(b)) => Ok(i64::from(b)),
Some(other) => {
let t = other.py_type_name(vm);
other.drop_with_heap(vm);
Err(ExcType::type_error(format!(
"'{t}' object cannot be interpreted as an integer"
)))
}
}
}
pub(crate) fn extract_count(val: Option<Value>, vm: &mut VM<'_, impl ResourceTracker>) -> RunResult<Option<usize>> {
match val {
None => Ok(Some(0)),
Some(Value::Int(n)) if n >= 0 => Ok(Some(usize::try_from(n).unwrap_or(usize::MAX))),
Some(Value::Bool(b)) => Ok(Some(usize::from(b))),
Some(Value::Int(_)) => Ok(None),
Some(other) => {
let t = other.py_type_name(vm);
other.drop_with_heap(vm);
Err(ExcType::type_error(format!(
"'{t}' object cannot be interpreted as an integer"
)))
}
}
}
fn compile_regex_limited(pattern: &str, flags: u16, delegate_size_limit: Option<usize>) -> Result<Regex, RegexError> {
let mut prefix = String::new();
if flags & IGNORECASE != 0 {
prefix.push('i');
}
if flags & MULTILINE != 0 {
prefix.push('m');
}
if flags & DOTALL != 0 {
prefix.push('s');
}
let full_pattern = if prefix.is_empty() {
pattern.to_owned()
} else {
format!("(?{prefix}){pattern}")
};
let mut builder = RegexBuilder::new(&full_pattern);
if let Some(limit) = delegate_size_limit {
builder.delegate_size_limit(limit);
}
builder.build()
}
fn is_size_limit_error(err: &RegexError) -> bool {
match err {
RegexError::CompileError(compile_error) => match &**compile_error {
CompileError::InnerError(inner) => inner.size_limit().is_some(),
_ => false,
},
_ => false,
}
}
fn translate_replacement(repl: &str) -> Cow<'_, str> {
if !repl.contains('\\') && !repl.contains('$') {
return Cow::Borrowed(repl);
}
let mut result = String::with_capacity(repl.len());
let mut chars = repl.chars().peekable();
while let Some(c) = chars.next() {
if c == '\\' {
match chars.peek() {
Some(&d) if d.is_ascii_digit() => {
result.push('$');
result.push(d);
chars.next();
}
Some(&'g') => {
chars.next(); translate_g_backref(&mut chars, &mut result);
}
Some(&'\\') => {
result.push('\\');
chars.next();
}
_ => {
result.push('\\');
}
}
} else if c == '$' {
result.push('$');
result.push('$');
} else {
result.push(c);
}
}
Cow::Owned(result)
}
fn translate_g_backref(chars: &mut iter::Peekable<str::Chars<'_>>, result: &mut String) {
if chars.peek() != Some(&'<') {
result.push('\\');
result.push('g');
return;
}
chars.next();
let mut name = String::new();
loop {
match chars.next() {
Some('>') => break,
Some(ch) => name.push(ch),
None => {
result.push('\\');
result.push('g');
result.push('<');
result.push_str(&name);
return;
}
}
}
result.push('$');
result.push('{');
result.push_str(&name);
result.push('}');
}
impl Serialize for RePattern {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
(&self.pattern, self.flags).serialize(serializer)
}
}
impl<'de> Deserialize<'de> for RePattern {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let (pattern, flags): (String, u16) = Deserialize::deserialize(deserializer)?;
Self::compile(pattern, flags).map_err(|e| de::Error::custom(format!("{e:?}")))
}
}