mirror of
https://github.com/quickwit-oss/tantivy.git
synced 2026-10-06 11:52:40 +00:00
Embed regex
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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()) }
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"));
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user