From 9659d4e050a341d48995f4435b2f0a8e7bba1cc0 Mon Sep 17 00:00:00 2001 From: Paul Masurel Date: Wed, 19 Aug 2026 10:09:24 +0200 Subject: [PATCH] Add SPLIT_BEFORE function --- jitexpr/docs/FUNCTIONS.md | 7 +- jitexpr/src/ast/serialize.rs | 2 + jitexpr/src/functions/mod.rs | 13 + jitexpr/src/functions/native_function.rs | 11 +- jitexpr/src/functions/split_before.rs | 316 +++++++++++++++++++++++ 5 files changed, 342 insertions(+), 7 deletions(-) create mode 100644 jitexpr/src/functions/split_before.rs diff --git a/jitexpr/docs/FUNCTIONS.md b/jitexpr/docs/FUNCTIONS.md index f4606f2af..c1d65fc51 100644 --- a/jitexpr/docs/FUNCTIONS.md +++ b/jitexpr/docs/FUNCTIONS.md @@ -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. diff --git a/jitexpr/src/ast/serialize.rs b/jitexpr/src/ast/serialize.rs index e761b7c73..71fae4c6e 100644 --- a/jitexpr/src/ast/serialize.rs +++ b/jitexpr/src/ast/serialize.rs @@ -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 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), diff --git a/jitexpr/src/functions/mod.rs b/jitexpr/src/functions/mod.rs index f24ae62e0..f5c1d56d9 100644 --- a/jitexpr/src/functions/mod.rs +++ b/jitexpr/src/functions/mod.rs @@ -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 => { ::call_with_types(args, target_type_set, context) } + Function::SplitBefore => { + ::call_with_types(args, target_type_set, context) + } Function::Subtract => { ::call_with_types(args, target_type_set, context) } @@ -285,6 +292,9 @@ impl Function { Function::Right => { ::infer_types(args, target_type, inferred_types) } + Function::SplitBefore => { + ::infer_types(args, target_type, inferred_types) + } Function::Subtract => { ::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) => { diff --git a/jitexpr/src/functions/native_function.rs b/jitexpr/src/functions/native_function.rs index b533cf4d6..b1749c927 100644 --- a/jitexpr/src/functions/native_function.rs +++ b/jitexpr/src/functions/native_function.rs @@ -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)?, diff --git a/jitexpr/src/functions/split_before.rs b/jitexpr/src/functions/split_before.rs new file mode 100644 index 000000000..0d7f07ebc --- /dev/null +++ b/jitexpr/src/functions/split_before.rs @@ -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, + separator: Arc, + occurrence: usize, +} + +fn constant_occurrence(expression: Option<&UntypedExpr>) -> Option { + 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 { + 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 { + 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 { + 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 { + 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. + 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 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 { + 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 + ); + } +}