Add SPLIT_BEFORE function

This commit is contained in:
Paul Masurel
2026-08-19 10:09:24 +02:00
parent 966aed48da
commit 9659d4e050
5 changed files with 342 additions and 7 deletions
+2 -5
View File
@@ -54,7 +54,7 @@ excluded from this pass, or is complex enough to warrant a separate implementati
| 40 | `LEFT` | done |
| 41 | `RIGHT` | done |
| 42 | `SUBSTRING` | done |
| 43 | `SPLIT_BEFORE` | out-of-scope |
| 43 | `SPLIT_BEFORE` | done |
| 44 | `SPLIT_AFTER` | out-of-scope |
| 50 | `REGEXP_EXTRACT` | done |
| 54 | `TRIM` | done |
@@ -76,7 +76,7 @@ excluded from this pass, or is complex enough to warrant a separate implementati
| 79 | `SUBSTRING_COUNT` | done |
| 80 | `REGEXP_LIKE` | done |
Progress: **32 / 32 in-scope** functions implemented; **23** functions are out-of-scope.
Progress: **33 / 33 in-scope** functions implemented; **22** functions are out-of-scope.
## Deferred implementation notes
@@ -102,9 +102,6 @@ Progress: **32 / 32 in-scope** functions implemented; **23** functions are out-o
- `SQRT` has a contradictory production type contract: the dd-go type checker returns the selected
input type (including integer), while the registry declares a `float64` output and both integer
and float kernels write `float64`. It is deferred until one of those contracts is chosen.
- `SPLIT_BEFORE` does not stop after handling a negative occurrence. Its scalar kernel continues
and can construct a negative slice bound, aborting the query. It is deferred until negative
occurrence semantics are defined explicitly.
- `SPLIT_AFTER` shares the negative-occurrence double-append defect with `SPLIT_BEFORE`; its scalar
result can contain both an empty value and the unsplit input even though the function is expected
to be scalar.
+2
View File
@@ -144,6 +144,7 @@ fn function_name(function: Function) -> &'static str {
Function::RegexpExtract => "REGEXP_EXTRACT",
Function::RegexpLike => "REGEXP_LIKE",
Function::Right => "RIGHT",
Function::SplitBefore => "SPLIT_BEFORE",
Function::Subtract => "SUBTRACT",
Function::Substring => "SUBSTRING",
Function::SubstringCount => "SUBSTRING_COUNT",
@@ -181,6 +182,7 @@ fn parse_function(name: &str, offset: usize) -> Result<Function, DeserializeErro
"REGEXP_EXTRACT" => Ok(Function::RegexpExtract),
"REGEXP_LIKE" => Ok(Function::RegexpLike),
"RIGHT" => Ok(Function::Right),
"SPLIT_BEFORE" => Ok(Function::SplitBefore),
"SUBTRACT" => Ok(Function::Subtract),
"SUBSTRING" => Ok(Function::Substring),
"SUBSTRING_COUNT" => Ok(Function::SubstringCount),
+13
View File
@@ -26,6 +26,7 @@ mod pow;
mod regexp_extract;
mod regexp_like;
mod right;
mod split_before;
mod substring;
mod substring_count;
mod subtract;
@@ -66,6 +67,7 @@ pub(crate) use self::pow::PowFnCall;
pub(crate) use self::regexp_extract::RegexpExtractFnCall;
pub(crate) use self::regexp_like::RegexpLikeFnCall;
pub(crate) use self::right::RightFnCall;
pub(crate) use self::split_before::SplitBeforeFnCall;
pub(crate) use self::substring::SubstringFnCall;
pub(crate) use self::substring_count::SubstringCountFnCall;
pub(crate) use self::subtract::SubtractFnCall;
@@ -131,6 +133,8 @@ pub enum Function {
RegexpLike,
/// Returns the last requested number of bytes from a string.
Right,
/// Returns the prefix before a selected occurrence of a literal separator.
SplitBefore,
/// Subtracts the second numeric argument from the first.
Subtract,
/// Returns a byte-indexed string slice.
@@ -205,6 +209,9 @@ impl Function {
Function::Right => {
<RightFnCall as FnCall>::call_with_types(args, target_type_set, context)
}
Function::SplitBefore => {
<SplitBeforeFnCall as FnCall>::call_with_types(args, target_type_set, context)
}
Function::Subtract => {
<SubtractFnCall as FnCall>::call_with_types(args, target_type_set, context)
}
@@ -285,6 +292,9 @@ impl Function {
Function::Right => {
<RightFnCall as FnCall>::infer_types(args, target_type, inferred_types)
}
Function::SplitBefore => {
<SplitBeforeFnCall as FnCall>::infer_types(args, target_type, inferred_types)
}
Function::Subtract => {
<SubtractFnCall as FnCall>::infer_types(args, target_type, inferred_types)
}
@@ -342,6 +352,7 @@ pub(crate) enum FnCallEnum {
RegexpExtract(RegexpExtractFnCall),
RegexpLike(RegexpLikeFnCall),
Right(RightFnCall),
SplitBefore(SplitBeforeFnCall),
Subtract(SubtractFnCall),
Substring(SubstringFnCall),
SubstringCount(SubstringCountFnCall),
@@ -379,6 +390,7 @@ impl FnCallEnum {
FnCallEnum::RegexpExtract(call) => call.args_mut(),
FnCallEnum::RegexpLike(call) => call.args_mut(),
FnCallEnum::Right(call) => call.args_mut(),
FnCallEnum::SplitBefore(call) => call.args_mut(),
FnCallEnum::Subtract(call) => call.args_mut(),
FnCallEnum::Substring(call) => call.args_mut(),
FnCallEnum::SubstringCount(call) => call.args_mut(),
@@ -424,6 +436,7 @@ impl FnCallEnum {
}
FnCallEnum::RegexpLike(call) => call.emit_cranelift_ir(return_type, context, builder),
FnCallEnum::Right(call) => call.emit_cranelift_ir(return_type, context, builder),
FnCallEnum::SplitBefore(call) => call.emit_cranelift_ir(return_type, context, builder),
FnCallEnum::Subtract(call) => call.emit_cranelift_ir(return_type, context, builder),
FnCallEnum::Substring(call) => call.emit_cranelift_ir(return_type, context, builder),
FnCallEnum::SubstringCount(call) => {
+9 -2
View File
@@ -2,8 +2,8 @@ use cranelift::codegen::ir::{FuncRef, Function as CraneliftFunction, Type};
use cranelift_jit::{JITBuilder, JITModule};
use super::{
comparison, concat, eq, int_mod, lower, pow, regexp_extract, regexp_like, substring,
substring_count, trim, upper,
comparison, concat, eq, int_mod, lower, pow, regexp_extract, regexp_like, split_before,
substring, substring_count, trim, upper,
};
use crate::compile::CompileError;
@@ -15,6 +15,7 @@ pub(crate) struct NativeFunctions {
string_trim: FuncRef,
substring_count: FuncRef,
substring: FuncRef,
split_before: FuncRef,
string_concat: FuncRef,
float_mod: FuncRef,
float_pow: FuncRef,
@@ -48,6 +49,10 @@ impl NativeFunctions {
self.substring
}
pub(crate) fn split_before(&self) -> FuncRef {
self.split_before
}
pub(crate) fn string_concat(&self) -> FuncRef {
self.string_concat
}
@@ -90,6 +95,7 @@ pub(crate) fn register_jit_symbols(jit_builder: &mut JITBuilder) {
trim::register_jit_symbol(jit_builder);
substring_count::register_jit_symbol(jit_builder);
substring::register_jit_symbol(jit_builder);
split_before::register_jit_symbol(jit_builder);
concat::register_jit_symbol(jit_builder);
int_mod::register_jit_symbol(jit_builder);
pow::register_jit_symbol(jit_builder);
@@ -111,6 +117,7 @@ pub(crate) fn declare_native_functions(
string_trim: trim::declare_native_function(module, function, pointer_type)?,
substring_count: substring_count::declare_native_function(module, function, pointer_type)?,
substring: substring::declare_native_function(module, function, pointer_type)?,
split_before: split_before::declare_native_function(module, function, pointer_type)?,
string_concat: concat::declare_native_function(module, function, pointer_type)?,
float_mod: int_mod::declare_native_function(module, function)?,
float_pow: pow::declare_native_function(module, function)?,
+316
View File
@@ -0,0 +1,316 @@
//! `SPLIT_BEFORE(input, separator, occurrence)` returns the part of a string before a separator.
//!
//! The separator must be a string constant. The optional occurrence must be a nonnegative integer
//! constant, uses zero-based indexing, and defaults to zero. Matches are literal, case-sensitive,
//! and non-overlapping. A missing occurrence or an empty separator returns a present empty string.
//! Null input, a null constant, or a negative occurrence returns null.
use std::collections::HashMap;
use std::sync::Arc;
use cranelift::codegen::ir::{FuncRef, Function as CraneliftFunction, Type, types};
use cranelift::frontend::FunctionBuilder;
use cranelift::prelude::{AbiParam, InstBuilder, IntCC};
use cranelift_jit::{JITBuilder, JITModule};
use cranelift_module::{Linkage, Module};
use crate::ast::{Function, InferredTypeSet, Literal, TypeError, UntypedExpr};
use crate::compile::{
CompileError, CompileFnBuilder, LoweredValue, LoweringContext, TypedExpr, TypedExprAst,
};
use crate::functions::{FnCall, FnCallEnum};
use crate::types::VarType;
const SYMBOL: &str = "jitexpr_split_before";
#[derive(Clone, Debug, PartialEq)]
pub(crate) struct SplitBeforeFnCall {
input: Box<TypedExpr>,
separator: Arc<str>,
occurrence: usize,
}
fn constant_occurrence(expression: Option<&UntypedExpr>) -> Option<usize> {
let Some(expression) = expression else {
return Some(0);
};
let UntypedExpr::Literal(literal) = expression else {
panic!("SPLIT_BEFORE occurrence must be constant");
};
match literal {
Literal::I64(value) => usize::try_from(*value).ok(),
Literal::U64(value) => usize::try_from(*value).ok(),
Literal::F64(value) => usize::try_from(*value as i64).ok(),
Literal::None => None,
Literal::Bool(_) | Literal::String(_) => {
unreachable!("type inference constrains SPLIT_BEFORE occurrence to an integer")
}
}
}
impl FnCall for SplitBeforeFnCall {
fn infer_types<'a>(
args: &'a [UntypedExpr],
target_type: InferredTypeSet,
inferred_types: &mut HashMap<&'a str, InferredTypeSet>,
) -> Result<InferredTypeSet, TypeError> {
if target_type.intersect(InferredTypeSet::STRING).is_none() {
return Err(TypeError::WrongFunctionReturnType {
function: Function::SplitBefore,
expected: target_type,
got: InferredTypeSet::STRING,
});
}
if !(2..=3).contains(&args.len()) {
return Err(TypeError::InvalidNumberOfArguments {
function: Function::SplitBefore,
expected: 3,
got: args.len(),
});
}
crate::ast::infer_types_aux(&args[0], InferredTypeSet::STRING, inferred_types)?;
crate::ast::infer_types_aux(&args[1], InferredTypeSet::STRING, inferred_types)?;
if let Some(occurrence) = args.get(2) {
crate::ast::infer_types_aux(occurrence, InferredTypeSet::I64, inferred_types)?;
}
Ok(InferredTypeSet::STRING)
}
fn call_with_types(
args: &[UntypedExpr],
_target_type_set: InferredTypeSet,
context: &mut CompileFnBuilder<'_, '_>,
) -> Result<TypedExpr, CompileError> {
assert!(
(2..=3).contains(&args.len()),
"expected 2 or 3 args for SPLIT_BEFORE"
);
let input = context.apply_types(&args[0], InferredTypeSet::STRING)?;
let separator = match &args[1] {
UntypedExpr::Literal(Literal::String(separator)) => Arc::clone(separator),
UntypedExpr::Literal(Literal::None) => return Ok(TypedExpr::none()),
_ => panic!("SPLIT_BEFORE separator must be a string constant"),
};
let Some(occurrence) = constant_occurrence(args.get(2)) else {
return Ok(TypedExpr::none());
};
Ok(TypedExpr {
return_type: VarType::Str,
ast: TypedExprAst::from_call(SplitBeforeFnCall {
input: Box::new(input),
separator,
occurrence,
}),
})
}
fn args_mut(&mut self) -> &mut [TypedExpr] {
std::slice::from_mut(&mut self.input)
}
fn emit_cranelift_ir(
&self,
return_type: VarType,
context: &mut LoweringContext<'_>,
builder: &mut FunctionBuilder<'_>,
) -> Result<LoweredValue, CompileError> {
debug_assert_eq!(return_type, VarType::Str);
let input = context.compile_expr(&self.input, builder)?;
let null = builder.ins().iconst(context.pointer_type(), 0);
let input_ptr = builder.ins().select(input.is_present, input.value, null);
let separator_ptr = builder
.ins()
.iconst(context.pointer_type(), self.separator.as_ptr() as i64);
let separator_len = builder
.ins()
.iconst(types::I64, self.separator.len() as i64);
let occurrence = builder.ins().iconst(types::I64, self.occurrence as i64);
let call = builder.ins().call(
context.native_functions().split_before(),
&[
input_ptr,
input.string_len,
separator_ptr,
separator_len,
occurrence,
],
);
let value = builder.inst_results(call)[0];
let string_len = builder.inst_results(call)[1];
let is_present = builder.ins().icmp_imm_u(IntCC::NotEqual, value, 0);
Ok(LoweredValue {
value,
is_present,
string_len,
})
}
}
pub(super) fn register_jit_symbol(builder: &mut JITBuilder) {
builder.symbol(SYMBOL, split_before as *const u8);
}
pub(super) fn declare_native_function(
module: &mut JITModule,
function: &mut CraneliftFunction,
pointer_type: Type,
) -> Result<FuncRef, CompileError> {
let mut signature = module.make_signature();
signature.params.extend([
AbiParam::new(pointer_type),
AbiParam::new(types::I64),
AbiParam::new(pointer_type),
AbiParam::new(types::I64),
AbiParam::new(types::I64),
]);
signature
.returns
.extend([AbiParam::new(pointer_type), AbiParam::new(types::I64)]);
let function_id = module.declare_function(SYMBOL, Linkage::Import, &signature)?;
Ok(module.declare_func_in_func(function_id, function))
}
#[repr(C)]
struct RawStr {
ptr: *const u8,
len: usize,
}
impl RawStr {
fn none() -> Self {
Self {
ptr: std::ptr::null(),
len: 0,
}
}
fn some(value: &str) -> Self {
Self {
ptr: value.as_ptr(),
len: value.len(),
}
}
}
unsafe extern "C" fn split_before(
input_ptr: *const u8,
input_len: usize,
separator_ptr: *const u8,
separator_len: usize,
occurrence: usize,
) -> RawStr {
if input_ptr.is_null() {
return RawStr::none();
}
// SAFETY: CompiledFn supplies a live UTF-8 input and the typed call owns the separator.
let input =
unsafe { std::str::from_utf8_unchecked(std::slice::from_raw_parts(input_ptr, input_len)) };
// SAFETY: The separator pointer and length come from the call's live Arc<str>.
let separator = unsafe {
std::str::from_utf8_unchecked(std::slice::from_raw_parts(separator_ptr, separator_len))
};
if separator.is_empty() {
return RawStr::some(&input[..0]);
}
let Some((separator_start, _)) = input.match_indices(separator).nth(occurrence) else {
return RawStr::some(&input[..0]);
};
RawStr::some(&input[..separator_start])
}
impl From<SplitBeforeFnCall> for FnCallEnum {
fn from(call: SplitBeforeFnCall) -> Self {
FnCallEnum::SplitBefore(call)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ast::{deserialize, infer_types};
use crate::compile::compile;
use crate::types::VariableValue;
fn eval(expression: &str) -> Option<String> {
let expression = deserialize(expression).unwrap();
let mut compiled = compile(&expression, &HashMap::new()).unwrap();
// SAFETY: The expression has no inputs and returns a nullable string.
unsafe { compiled.call(&[]).as_str().map(str::to_owned) }
}
#[test]
fn test_signature_and_occurrences() {
for expression in [
"(SPLIT_BEFORE \"a.b\")",
"(SPLIT_BEFORE \"a.b\" \".\" 0i64 1i64)",
] {
let expression = deserialize(expression).unwrap();
assert!(matches!(
infer_types(&expression),
Err(TypeError::InvalidNumberOfArguments {
function: Function::SplitBefore,
expected: 3,
..
})
));
}
assert_eq!(eval("(SPLIT_BEFORE \"a.b.c\" \".\")"), Some("a".into()));
assert_eq!(
eval("(SPLIT_BEFORE \"a.b.c\" \".\" 0i64)"),
Some("a".into())
);
assert_eq!(
eval("(SPLIT_BEFORE \"a.b.c\" \".\" 1i64)"),
Some("a.b".into())
);
}
#[test]
fn test_literal_non_overlapping_and_unicode_matches() {
assert_eq!(
eval("(SPLIT_BEFORE \"......\" \"...\" 1i64)"),
Some("...".into())
);
assert_eq!(
eval("(SPLIT_BEFORE \"a...\" \"..\" 0i64)"),
Some("a".into())
);
assert_eq!(
eval("(SPLIT_BEFORE \"α→β→γ\" \"→\" 1i64)"),
Some("α→β".into())
);
}
#[test]
fn test_empty_missing_and_invalid_occurrences() {
assert_eq!(
eval("(SPLIT_BEFORE \"abc\" \".\" 0i64)"),
Some(String::new())
);
assert_eq!(
eval("(SPLIT_BEFORE \"abc\" \"\" 3i64)"),
Some(String::new())
);
assert_eq!(eval("(SPLIT_BEFORE \"a.b\" \".\" -1i64)"), None);
assert_eq!(eval("(SPLIT_BEFORE \"a.b\" \".\" none)"), None);
assert_eq!(eval("(SPLIT_BEFORE \"a.b\" none)"), None);
}
#[test]
fn test_runtime_null() {
let expression = deserialize("(SPLIT_BEFORE value \".\")").unwrap();
let mut compiled = compile(&expression, &HashMap::from([("value", VarType::Str)])).unwrap();
// SAFETY: The compiled expression expects one nullable string argument.
assert_eq!(
unsafe { compiled.call(&[VariableValue::some("a.b")]).as_str() },
Some("a")
);
// SAFETY: The compiled expression expects one nullable string argument.
assert_eq!(
unsafe { compiled.call(&[VariableValue::none()]).as_str() },
None
);
}
}