diff --git a/jitexpr/src/compile/compile_fn_builder.rs b/jitexpr/src/compile/compile_fn_builder.rs index 2823100b8..547246d61 100644 --- a/jitexpr/src/compile/compile_fn_builder.rs +++ b/jitexpr/src/compile/compile_fn_builder.rs @@ -8,7 +8,6 @@ use cranelift::codegen::ir::{MemFlagsData, UserFuncName}; use cranelift::prelude::*; use cranelift_jit::{JITBuilder, JITModule}; use cranelift_module::{FuncId, Module, ModuleError, default_libcall_names}; -use regex::Regex; use super::compiled_fn::JitEntry; use super::{ @@ -18,19 +17,9 @@ use crate::ast::{InferredTypeSet, Literal, UntypedExpr}; use crate::functions::{declare_native_functions, register_jit_symbols}; use crate::types::VarType; -#[derive(Clone, Copy, Debug, Eq, PartialEq)] -pub(crate) struct RegexRef(usize); - -impl RegexRef { - pub(crate) fn index(self) -> usize { - self.0 - } -} - pub(crate) struct CompileFnBuilder<'types, 'names> { variable_types: &'types HashMap<&'names str, VarType>, input_vars: Vec, - regexes: Vec, } struct LoweredFunction { @@ -38,7 +27,6 @@ struct LoweredFunction { context: CodegenContext, function_id: FuncId, input_vars: Vec, - regexes: Box<[Regex]>, expression: Box, } @@ -47,7 +35,6 @@ impl<'types, 'names> CompileFnBuilder<'types, 'names> { CompileFnBuilder { variable_types, input_vars: Vec::new(), - regexes: Vec::new(), } } @@ -55,12 +42,6 @@ impl<'types, 'names> CompileFnBuilder<'types, 'names> { self.variable_types } - pub(crate) fn register_regex(&mut self, regex: Regex) -> RegexRef { - let regex_ref = RegexRef(self.regexes.len()); - self.regexes.push(regex); - regex_ref - } - /// If a variable is missing from `variable_types`, it is treated as `None`. pub(crate) fn build_typed_expr( &mut self, @@ -196,12 +177,7 @@ impl<'types, 'names> CompileFnBuilder<'types, 'names> { } fn lower_typed_expr(self, expression: TypedExpr) -> Result { - let CompileFnBuilder { - input_vars, - regexes, - .. - } = self; - let regexes = regexes.into_boxed_slice(); + let CompileFnBuilder { input_vars, .. } = self; let expression = Box::new(expression); let mut jit_builder = @@ -211,12 +187,11 @@ impl<'types, 'names> CompileFnBuilder<'types, 'names> { let target_config = module.target_config(); let pointer_type = target_config.pointer_type(); - // The native entry point mirrors JitEntry: the two arguments point to the - // input slots and CompiledFn::regexes. VariableValue is returned as two - // integer-class values according to the native C ABI. + // The native entry point mirrors JitEntry: its argument points to the + // input slots, and VariableValue is returned as two integer-class values + // according to the native C ABI. let mut signature = module.make_signature(); signature.params.push(AbiParam::new(pointer_type)); - signature.params.push(AbiParam::new(pointer_type)); signature.returns.push(AbiParam::new(types::I64)); signature.returns.push(AbiParam::new(types::I64)); let function_id = module.declare_anonymous_function(&signature)?; @@ -237,10 +212,8 @@ impl<'types, 'names> CompileFnBuilder<'types, 'names> { builder.seal_block(entry_block); let args_ptr = builder.block_params(entry_block)[0]; - let regexes_ptr = builder.block_params(entry_block)[1]; let mut lowering_context = LoweringContext { args_ptr, - regexes_ptr, pointer_type, native_functions: &native_functions, }; @@ -268,7 +241,6 @@ impl<'types, 'names> CompileFnBuilder<'types, 'names> { context, function_id, input_vars, - regexes, expression, }) } @@ -297,7 +269,6 @@ impl LoweredFunction { mut context, function_id, input_vars, - regexes, expression, } = self; @@ -312,7 +283,6 @@ impl LoweredFunction { Ok(CompiledFn { entry, _module: module, - regexes, inputs: input_vars, _typed_expr: expression, }) @@ -482,7 +452,7 @@ mod tests { let lowered = builder.lower_typed_expr(expression).unwrap(); let signature = &lowered.context.func.signature; - assert_eq!(signature.params.len(), 2); + assert_eq!(signature.params.len(), 1); assert!( signature .params diff --git a/jitexpr/src/compile/compiled_fn.rs b/jitexpr/src/compile/compiled_fn.rs index 4eef73c2c..d6ece1da6 100644 --- a/jitexpr/src/compile/compiled_fn.rs +++ b/jitexpr/src/compile/compiled_fn.rs @@ -1,5 +1,4 @@ use cranelift_jit::JITModule; -use regex::Regex; use super::{TypedExpr, TypedVariable}; use crate::types::{VarType, VariableValue}; @@ -19,7 +18,7 @@ compile_error!( // types.rs, not an interface intended for C callers. #[allow(improper_ctypes_definitions)] pub(crate) type JitEntry = - for<'a> unsafe extern "C" fn(*const VariableValue<'a>, *const Regex) -> VariableValue<'a>; + for<'a> unsafe extern "C" fn(*const VariableValue<'a>) -> VariableValue<'a>; /// An expression compiled to native machine code. /// @@ -28,11 +27,9 @@ pub(crate) type JitEntry = pub struct CompiledFn { pub(crate) entry: JitEntry, pub(crate) _module: JITModule, - // Generated code selects a compiled regex by its index in this array. - pub(crate) regexes: Box<[Regex]>, /// Input slots in the exact order expected by [`CompiledFn::call`]. pub inputs: Vec, - // This AST owns the Arc-backed string literals embedded in generated code. + // This AST owns the Arc-backed literals and regexes embedded in generated code. pub(crate) _typed_expr: Box, } @@ -65,6 +62,6 @@ impl CompiledFn { let args: &[VariableValue<'output>] = args; // SAFETY: Guaranteed by the caller. Both the input and compiled-function // lifetimes outlive the lifetime selected for the returned value. - unsafe { (self.entry)(args.as_ptr(), self.regexes.as_ptr()) } + unsafe { (self.entry)(args.as_ptr()) } } } diff --git a/jitexpr/src/compile/mod.rs b/jitexpr/src/compile/mod.rs index 5925e95fd..15981944e 100644 --- a/jitexpr/src/compile/mod.rs +++ b/jitexpr/src/compile/mod.rs @@ -5,7 +5,7 @@ mod typed_expr; use std::collections::HashMap; -pub(crate) use compile_fn_builder::{CompileFnBuilder, RegexRef}; +pub(crate) use compile_fn_builder::CompileFnBuilder; pub use compiled_fn::CompiledFn; use cranelift::codegen::ir::{ InstBuilder as _, MemFlagsData, Type, Value as CraneliftValue, types as cranelift_types, @@ -40,7 +40,6 @@ pub fn compile_to_assembly( pub(crate) struct LoweringContext<'a> { args_ptr: CraneliftValue, - regexes_ptr: CraneliftValue, pointer_type: Type, native_functions: &'a NativeFunctions, } @@ -118,10 +117,6 @@ impl LoweringContext<'_> { self.pointer_type } - pub(crate) fn regexes_ptr(&self) -> CraneliftValue { - self.regexes_ptr - } - pub(crate) fn native_functions(&self) -> &NativeFunctions { self.native_functions } diff --git a/jitexpr/src/functions/mod.rs b/jitexpr/src/functions/mod.rs index 7a3393ec5..92ce305a9 100644 --- a/jitexpr/src/functions/mod.rs +++ b/jitexpr/src/functions/mod.rs @@ -136,7 +136,7 @@ pub(crate) trait FnCall: std::fmt::Debug + Into { /// /// `target_type_set` communicates the result types set accepted by the parent call. The /// implementation selects a concrete result type, applies compatible target types to its - /// arguments through `context`, and registers any compilation resources owned by the call. + /// arguments through `context`, and stores any compilation resources on the typed call. /// /// The type of the returned is given to the caller in the TypedExpr object. fn call_with_types( diff --git a/jitexpr/src/functions/regexp_extract.rs b/jitexpr/src/functions/regexp_extract.rs index f40d1c621..1b13c5cd1 100644 --- a/jitexpr/src/functions/regexp_extract.rs +++ b/jitexpr/src/functions/regexp_extract.rs @@ -12,6 +12,7 @@ // group is absent or did not participate in the match. use std::collections::HashMap; +use std::sync::Arc; use cranelift::codegen::ir::{FuncRef, Function as CraneliftFunction, Type, types}; use cranelift::frontend::FunctionBuilder; @@ -22,21 +23,28 @@ use regex::Regex; use crate::ast::{Function, InferredTypeSet, Literal, TypeError, UntypedExpr}; use crate::compile::{ - CompileError, CompileFnBuilder, LoweredValue, LoweringContext, RegexRef, TypedExpr, - TypedExprAst, + CompileError, CompileFnBuilder, LoweredValue, LoweringContext, TypedExpr, TypedExprAst, }; use crate::functions::{FnCall, FnCallEnum}; use crate::types::VarType; const SYMBOL: &str = "jitexpr_regexp_extract"; -#[derive(Clone, Debug, PartialEq)] +#[derive(Clone, Debug)] pub(crate) struct RegexpExtractFnCall { - regex_ref: RegexRef, + regex: Arc, haystack: Box, capture_index: u64, } +impl PartialEq for RegexpExtractFnCall { + fn eq(&self, other: &Self) -> bool { + self.regex.as_str() == other.regex.as_str() + && self.haystack == other.haystack + && self.capture_index == other.capture_index + } +} + impl FnCall for RegexpExtractFnCall { fn infer_types<'a>( args: &'a [UntypedExpr], @@ -79,11 +87,12 @@ impl FnCall for RegexpExtractFnCall { let UntypedExpr::Literal(Literal::String(pattern)) = &args[1] else { panic!("regexp_extract pattern must be a string literal"); }; - let regex = Regex::new(pattern).map_err(|source| CompileError::InvalidRegex { - pattern: pattern.to_string(), - source, - })?; - let regex_ref = context.register_regex(regex); + let regex = Arc::new( + Regex::new(pattern).map_err(|source| CompileError::InvalidRegex { + pattern: pattern.to_string(), + source, + })?, + ); let UntypedExpr::Literal(Literal::U64(capture_index)) = &args[2] else { panic!("regexp_extract capture index must be a u64 literal"); @@ -92,7 +101,7 @@ impl FnCall for RegexpExtractFnCall { Ok(TypedExpr { return_type: VarType::Str, ast: TypedExprAst::from_call(RegexpExtractFnCall { - regex_ref, + regex, haystack: Box::new(haystack), capture_index: *capture_index, }), @@ -117,19 +126,13 @@ impl FnCall for RegexpExtractFnCall { let haystack_ptr = builder .ins() .select(haystack.is_present, haystack.value, null); - let regex_index = builder + let regex_ptr = builder .ins() - .iconst(context.pointer_type(), self.regex_ref.index() as i64); + .iconst(context.pointer_type(), Arc::as_ptr(&self.regex) as i64); let capture_index = builder.ins().iconst(types::I64, self.capture_index as i64); let call = builder.ins().call( context.native_functions().regexp_extract(), - &[ - context.regexes_ptr(), - regex_index, - haystack_ptr, - haystack.string_len, - capture_index, - ], + &[regex_ptr, haystack_ptr, haystack.string_len, capture_index], ); let value = builder.inst_results(call)[0]; let string_len = builder.inst_results(call)[1]; @@ -156,7 +159,7 @@ pub(super) fn declare_native_function( let mut signature = module.make_signature(); signature .params - .extend(std::iter::repeat_n(AbiParam::new(pointer_type), 3)); + .extend(std::iter::repeat_n(AbiParam::new(pointer_type), 2)); signature.params.push(AbiParam::new(types::I64)); signature.params.push(AbiParam::new(types::I64)); signature.returns.push(AbiParam::new(pointer_type)); @@ -193,8 +196,7 @@ impl RawStr { /// The JIT forwards a nullable UTF-8 pointer and byte length. The returned /// pointer and length borrow directly from the haystack. unsafe extern "C" fn regexp_extract( - regexes: *const Regex, - regex_index: usize, + regex: *const Regex, haystack_ptr: *const u8, haystack_len: usize, capture_index: u64, @@ -205,9 +207,9 @@ unsafe extern "C" fn regexp_extract( let Ok(capture_index) = usize::try_from(capture_index) else { return RawStr::none(); }; - // SAFETY: Generated code passes CompiledFn::regexes and an index - // assigned while constructing that same array. - let regex = unsafe { &*regexes.add(regex_index) }; + // SAFETY: Generated code embeds a pointer to the Arc-owned Regex stored in + // the typed expression retained by CompiledFn. + let regex = unsafe { &*regex }; // SAFETY: The contract of CompiledFn::call requires a live UTF-8 string // pointer and its exact byte length for every present string input. let haystack = unsafe { @@ -260,8 +262,6 @@ mod tests { let input = [VariableValue::some(haystack)]; let output = unsafe { compiled.call(&input) }; - assert_eq!(compiled.regexes.len(), 1); - assert_eq!(compiled.regexes[0].as_str(), r"([a-z]+)-(\d+)"); let extracted = unsafe { output.as_str() }.unwrap(); assert_eq!(extracted, "user"); assert_eq!(extracted.as_ptr(), haystack[7..].as_ptr()); @@ -334,7 +334,7 @@ mod tests { } #[test] - fn test_compile_nested_calls_use_distinct_regexes() { + fn test_compile_nested_calls_use_their_own_regexes() { let expression = ast::deserialize( r#"(REGEXP_EXTRACT (REGEXP_EXTRACT message "([a-z]+-\\d+)" 1u64) @@ -347,7 +347,6 @@ mod tests { let input = [VariableValue::some("id=user-123!")]; let output = unsafe { compiled.call(&input) }; - assert_eq!(compiled.regexes.len(), 2); assert_eq!(unsafe { output.as_str() }, Some("user")); }