Embed regex

This commit is contained in:
Paul Masurel
2026-08-18 11:47:43 +02:00
parent 3ea070a994
commit a6cec0585e
5 changed files with 38 additions and 77 deletions
+5 -35
View File
@@ -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<TypedVariable>,
regexes: Vec<Regex>,
}
struct LoweredFunction {
@@ -38,7 +27,6 @@ struct LoweredFunction {
context: CodegenContext,
function_id: FuncId,
input_vars: Vec<TypedVariable>,
regexes: Box<[Regex]>,
expression: Box<TypedExpr>,
}
@@ -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<LoweredFunction, CompileError> {
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
+3 -6
View File
@@ -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<TypedVariable>,
// 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<TypedExpr>,
}
@@ -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()) }
}
}
+1 -6
View File
@@ -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
}
+1 -1
View File
@@ -136,7 +136,7 @@ pub(crate) trait FnCall: std::fmt::Debug + Into<FnCallEnum> {
///
/// `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(
+28 -29
View File
@@ -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<Regex>,
haystack: Box<TypedExpr>,
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"));
}