(Calculated fields) Add JIT expression framework with ADD, IS_NULL, and REGEXP_EXTRACT (#3078)

Add JIT expression framework with ADD, IS_NULL, and REGEXP_EXTRACT
This commit is contained in:
Paul Masurel
2026-09-10 15:30:02 +02:00
committed by GitHub
parent bc8a21cf12
commit f8ec38fd9f
23 changed files with 4711 additions and 0 deletions
+1
View File
@@ -165,6 +165,7 @@ members = [
"sstable",
"tokenizer-api",
"columnar",
"jitexpr",
]
# Following the "fail" crate best practises, we isolate
+12
View File
@@ -0,0 +1,12 @@
[package]
name = "jitexpr"
version = "0.1.0"
edition = "2024"
[dependencies]
cranelift = "0.134.3"
cranelift-jit = "0.134.3"
cranelift-module = "0.134.3"
cranelift-native = "0.134.3"
regex = "1"
thiserror = "2.0.1"
+12
View File
@@ -0,0 +1,12 @@
This is an expression compiler relying on Cranelift.
UntypedExpr
↓ injecting variable types, and type checking
TypedExpr
↓ lowering
Cranelift IR
↓ Cranelift code generation
Machine code (or assembly)
The project does not rely on cranelifts function call abstraction.
Instead it just manipulates expression, so everything is always inlined.
+46
View File
@@ -0,0 +1,46 @@
use std::collections::HashMap;
use std::error::Error;
use std::sync::Arc;
use jitexpr::ast::{Function, InferredTypeSet, UntypedExpr, infer_types};
use jitexpr::compile::{CompiledFn, CompiledFnCtx, compile};
use jitexpr::types::{VarType, VariableValue};
fn main() -> Result<(), Box<dyn Error>> {
// A simple expression that goes:
// my_col + 1
let untyped_expr = UntypedExpr::new_fn_call(
Function::Add,
vec![
UntypedExpr::variable("my_col"),
UntypedExpr::literal(1.0f64),
],
)?;
// Infer types does not return specific types, but instead a set of acceptable
// types for each variables.
let inferred_types: HashMap<&str, InferredTypeSet> = infer_types(&untyped_expr)?;
assert_eq!(
inferred_types.get("my_col").unwrap(),
&InferredTypeSet::NUMERICAL
);
// This is then up to us to decide the actual type for each variable.
// In tantivy, this means picking the first column with a type in inferred_types.
//
// If none match then we should use the VarType::None.
let variable_types: HashMap<&str, VarType> =
std::iter::once(("my_col", VarType::F64)).collect();
let compiled_fn: Arc<CompiledFn> = compile(&untyped_expr, &variable_types)?;
let mut compiled_fn_ctx = CompiledFnCtx::new(compiled_fn);
// We use a nullable wrapper around the value union to pass typed variables.
// For present values, it is up to us to populate the correct union member.
// Not doing so is UB.
let input: Box<[VariableValue]> = vec![VariableValue::from(1.2f64)].into_boxed_slice();
let output = unsafe { compiled_fn_ctx.call(&input[..]) };
assert_eq!(unsafe { output.as_f64() }, Some(1.2f64 + 1.0f64));
Ok(())
}
+309
View File
@@ -0,0 +1,309 @@
use std::collections::HashMap;
use std::collections::hash_map::Entry;
use crate::ast::{Function, Literal, UntypedExpr};
use crate::functions::InvalidFnCall;
use crate::types::VarType;
#[derive(Default, Copy, Clone, Debug, Eq, PartialEq)]
pub struct InferredTypeSet {
pub string: bool,
pub i64: bool,
pub u64: bool,
pub f64: bool,
pub boolean: bool,
}
impl InferredTypeSet {
pub const NONE: InferredTypeSet = InferredTypeSet {
string: false,
i64: false,
u64: false,
f64: false,
boolean: false,
};
pub const ALL: InferredTypeSet = InferredTypeSet {
string: true,
i64: true,
u64: true,
f64: true,
boolean: true,
};
pub const NUMERICAL: InferredTypeSet = InferredTypeSet {
i64: true,
u64: true,
f64: true,
boolean: false,
string: false,
};
pub const I64: InferredTypeSet = InferredTypeSet {
i64: true,
..Self::NONE
};
pub const U64: InferredTypeSet = InferredTypeSet {
u64: true,
..Self::NONE
};
pub const F64: InferredTypeSet = InferredTypeSet {
f64: true,
..Self::NONE
};
pub const STRING: InferredTypeSet = InferredTypeSet {
string: true,
..Self::NONE
};
pub const BOOLEAN: InferredTypeSet = InferredTypeSet {
boolean: true,
..Self::NONE
};
pub(crate) fn is_none(self) -> bool {
self == Self::NONE
}
pub fn singleton(var_type: VarType) -> InferredTypeSet {
match var_type {
VarType::Bool => Self::BOOLEAN,
VarType::F64 => Self::F64,
VarType::U64 => Self::U64,
VarType::I64 => Self::I64,
VarType::Str => Self::STRING,
VarType::None => Self::NONE,
}
}
pub(crate) fn intersect(self, target_inferred_type: InferredTypeSet) -> InferredTypeSet {
InferredTypeSet {
string: self.string && target_inferred_type.string,
i64: self.i64 && target_inferred_type.i64,
u64: self.u64 && target_inferred_type.u64,
f64: self.f64 && target_inferred_type.f64,
boolean: self.boolean && target_inferred_type.boolean,
}
}
pub fn contains(&self, var_type: VarType) -> bool {
match var_type {
VarType::Bool => self.boolean,
VarType::F64 => self.f64,
VarType::U64 => self.u64,
VarType::I64 => self.i64,
VarType::Str => self.string,
VarType::None => self.is_none(),
}
}
}
impl From<VarType> for InferredTypeSet {
fn from(var_type: VarType) -> Self {
match var_type {
VarType::Bool => Self::BOOLEAN,
VarType::F64 => Self::F64,
VarType::U64 => Self::U64,
VarType::I64 => Self::I64,
VarType::Str => Self::STRING,
VarType::None => Self::NONE,
}
}
}
impl std::fmt::Display for InferredTypeSet {
fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
let mut types = Vec::new();
if self.string {
types.push("string");
}
if self.i64 {
types.push("i64");
}
if self.u64 {
types.push("u64");
}
if self.f64 {
types.push("f64");
}
if self.boolean {
types.push("boolean");
}
write!(f, "{{{}}}", types.join(", "))
}
}
#[derive(Debug, thiserror::Error)]
pub enum TypeError {
#[error(transparent)]
InvalidFnCall(#[from] InvalidFnCall),
#[error("function `{function:?}` returns `{got}`, expected `{expected}`")]
WrongFunctionReturnType {
function: Function,
expected: InferredTypeSet,
got: InferredTypeSet,
},
#[error("function `{function:?}` expects `{expected}` args, was passed `{got}`")]
InvalidNumberOfArguments {
function: Function,
expected: usize,
got: usize,
},
#[error("expected `{expected}` , got `{literal:?}`")]
InvalidLiteralType {
literal: Literal,
expected: InferredTypeSet,
},
}
/// Infer the accepted types for the different variables present in the formula.
pub fn infer_types(expr: &UntypedExpr) -> Result<HashMap<&str, InferredTypeSet>, TypeError> {
infer_types_with_target(expr, InferredTypeSet::ALL)
}
/// Infer the accepted variable types while constraining the expression's result type.
pub fn infer_types_with_target(
expr: &UntypedExpr,
target_type: InferredTypeSet,
) -> Result<HashMap<&str, InferredTypeSet>, TypeError> {
let mut inferred_type_res = HashMap::default();
infer_types_aux(expr, target_type, &mut inferred_type_res)?;
Ok(inferred_type_res)
}
/// Infer the possible types of an UntypedExpr, meant to represent `target_inferred_type`.
///
/// As we call it recursively on the different nodes of the expression,
/// this method should mutate the inferred_types (found in the inferred_type_res map) of each
/// variable name encounterred, always restricting them.
pub(crate) fn infer_types_aux<'a>(
expr: &'a UntypedExpr,
target_inferred_type: InferredTypeSet,
inferred_types_res: &mut HashMap<&'a str, InferredTypeSet>,
) -> Result<InferredTypeSet, TypeError> {
match expr {
UntypedExpr::Literal(literal) => {
let literal_type: InferredTypeSet = target_inferred_type.intersect(literal.types());
if literal_type.is_none() {
return Err(TypeError::InvalidLiteralType {
literal: literal.clone(),
expected: target_inferred_type,
});
}
Ok(literal_type)
}
UntypedExpr::Variable(variable_name) => {
match inferred_types_res.entry(variable_name.as_ref()) {
Entry::Occupied(mut occupied_entry) => {
let inferred_types = occupied_entry.get().intersect(target_inferred_type);
occupied_entry.insert(inferred_types);
Ok(inferred_types)
}
Entry::Vacant(vacant_entry) => {
vacant_entry.insert_entry(target_inferred_type);
Ok(target_inferred_type)
}
}
}
UntypedExpr::FnCall { function, args } => {
function.infer_types(args, target_inferred_type, inferred_types_res)
}
}
}
pub(crate) fn infer_type_with_variable_types(
expr: &UntypedExpr,
target_inferred_type: InferredTypeSet,
variable_types: &HashMap<&str, VarType>,
) -> Result<InferredTypeSet, TypeError> {
let mut inferred_types = HashMap::new();
seed_variable_types(expr, variable_types, &mut inferred_types);
infer_types_aux(expr, target_inferred_type, &mut inferred_types)
}
/// Populate the inferred_types HashMap with the value types proved by the user.
fn seed_variable_types<'a>(
expr: &'a UntypedExpr,
variable_types: &HashMap<&str, VarType>,
inferred_types: &mut HashMap<&'a str, InferredTypeSet>,
) {
match expr {
UntypedExpr::Literal(_) => {}
UntypedExpr::Variable(variable_name) => {
// If the value is not provided by the user (for instance because we fed values from a
// columnar and no column with that column name exists), we treat it has being None.
let inferred_type = variable_types
.get(variable_name.as_ref())
.copied()
.map(InferredTypeSet::from)
.unwrap_or(InferredTypeSet::NONE);
inferred_types.insert(variable_name.as_ref(), inferred_type);
}
UntypedExpr::FnCall { args, .. } => {
for arg in args {
seed_variable_types(arg, variable_types, inferred_types);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_infer_types_bare_variable_accepts_all() {
// A lone variable should accept all types.
let expr = UntypedExpr::variable("a");
let inferred_types = infer_types(&expr).unwrap();
let a_types = inferred_types.get("a").unwrap();
assert_eq!(a_types, &InferredTypeSet::ALL);
}
#[test]
fn test_infer_type_uses_concrete_variable_types() {
let expr = UntypedExpr::new_fn_call(
Function::Add,
vec![UntypedExpr::variable("my_col"), UntypedExpr::literal(1i64)],
)
.unwrap();
let variable_types = HashMap::from([("my_col", VarType::U64)]);
let inferred_type =
infer_type_with_variable_types(&expr, InferredTypeSet::NUMERICAL, &variable_types)
.unwrap();
assert_eq!(inferred_type, InferredTypeSet::U64);
}
#[test]
fn test_infer_type_falls_back_to_f64_for_disjoint_numeric_types() {
let expr = UntypedExpr::new_fn_call(
Function::Add,
vec![
UntypedExpr::variable("unsigned"),
UntypedExpr::variable("signed"),
],
)
.unwrap();
let variable_types = HashMap::from([("unsigned", VarType::U64), ("signed", VarType::I64)]);
let inferred_type =
infer_type_with_variable_types(&expr, InferredTypeSet::NUMERICAL, &variable_types)
.unwrap();
assert_eq!(inferred_type, InferredTypeSet::F64);
}
#[test]
fn test_inferred_type_set_display_lists_concrete_numeric_types() {
assert_eq!(
InferredTypeSet::ALL.to_string(),
"{string, i64, u64, f64, boolean}"
);
assert_eq!(InferredTypeSet::NUMERICAL.to_string(), "{i64, u64, f64}");
}
}
+182
View File
@@ -0,0 +1,182 @@
use std::sync::Arc;
use crate::ast::InferredTypeSet;
use crate::types::VarType;
/// A literal supported by the first expression-language milestone.
#[derive(Clone, Debug, PartialEq)]
pub enum Literal {
None,
Bool(bool),
U64(u64),
I64(i64),
F64(f64),
String(Arc<str>),
}
impl Literal {
pub fn is_none(&self) -> bool {
matches!(self, Literal::None)
}
pub fn types(&self) -> InferredTypeSet {
match self {
Literal::None => InferredTypeSet::ALL,
Literal::Bool(_) => InferredTypeSet::BOOLEAN,
// A literal number represents a "real number". It can sometime be represented by a i64,
// a u64 or a f64. The choice of this representation is rather arbitrary. It
// can be the result of an implementation detail of serde_json for instance.
//
// Here we want to return the set of possible representation for the associated number.
Literal::I64(value) => InferredTypeSet {
i64: true,
u64: *value >= 0, // Any non-negative i64 can be represented as u64.
f64: true, // We always accept f64.
..InferredTypeSet::NONE
},
Literal::U64(value) => InferredTypeSet {
i64: *value <= i64::MAX as u64, // any u64 below i64::MAX can be represented as a
// i64.
u64: true,
f64: true, // We always accept f64
..InferredTypeSet::NONE
},
Literal::F64(value) => {
let is_integral: bool = value.fract() == 0.0;
InferredTypeSet {
i64: is_integral && *value >= i64::MIN as f64 && *value < -(i64::MIN as f64),
u64: is_integral && *value >= 0.0 && *value < u64::MAX as f64,
f64: true,
..InferredTypeSet::NONE
}
}
Literal::String(_) => InferredTypeSet::STRING,
}
}
#[cfg(test)]
pub fn r#type(&self) -> VarType {
match self {
Literal::None => VarType::None,
Literal::Bool(_) => VarType::Bool,
Literal::U64(_) => VarType::U64,
Literal::I64(_) => VarType::I64,
Literal::F64(_) => VarType::F64,
Literal::String(_) => VarType::Str,
}
}
}
impl From<bool> for Literal {
fn from(value: bool) -> Self {
Literal::Bool(value)
}
}
impl From<u64> for Literal {
fn from(value: u64) -> Self {
Literal::U64(value)
}
}
impl From<i64> for Literal {
fn from(value: i64) -> Self {
Literal::I64(value)
}
}
impl From<f64> for Literal {
fn from(value: f64) -> Self {
Literal::F64(value)
}
}
impl From<String> for Literal {
fn from(value: String) -> Self {
Literal::String(Arc::from(value))
}
}
impl From<&str> for Literal {
fn from(value: &str) -> Self {
Literal::String(Arc::from(value.to_string()))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_literal_types_depend_on_representable_value() {
let i64_f64 = InferredTypeSet {
i64: true,
f64: true,
..InferredTypeSet::NONE
};
let u64_f64 = InferredTypeSet {
u64: true,
f64: true,
..InferredTypeSet::NONE
};
assert_eq!(Literal::U64(1).types(), InferredTypeSet::NUMERICAL);
assert_eq!(Literal::I64(1).types(), InferredTypeSet::NUMERICAL);
assert_eq!(Literal::I64(-1).types(), i64_f64);
assert_eq!(Literal::U64(1 << 63).types(), u64_f64);
assert_eq!(Literal::F64(1.2).types(), InferredTypeSet::F64);
assert_eq!(Literal::F64(1.0).types(), InferredTypeSet::NUMERICAL);
}
#[test]
fn test_literal_types_accept_lossless_float_representation() {
assert_eq!(
Literal::I64((1 << 53) + 1).types(),
InferredTypeSet::NUMERICAL
);
// even though i64::MAX - 1i64 cannot be represented as f64 in a lossless manner...
assert_ne!(((i64::MAX - 1i64) as f64) as i64, (i64::MAX - 1));
// ... we list f64 as a valid inferred type.
assert_eq!(
Literal::I64(i64::MAX - 1).types(),
InferredTypeSet::NUMERICAL
);
assert_eq!(
Literal::U64(u64::MAX).types(),
InferredTypeSet {
u64: true,
f64: true,
..InferredTypeSet::NONE
}
);
assert_eq!(
Literal::I64(i64::MIN).types(),
InferredTypeSet {
i64: true,
f64: true,
..InferredTypeSet::NONE
}
);
}
#[test]
fn test_f64_literal_types_handle_integer_boundaries_and_special_values() {
assert_eq!(
Literal::F64(2f64.powi(63)).types(),
InferredTypeSet {
u64: true,
f64: true,
..InferredTypeSet::NONE
}
);
assert_eq!(Literal::F64(2f64.powi(64)).types(), InferredTypeSet::F64);
assert_eq!(Literal::F64(-0.0).types(), InferredTypeSet::NUMERICAL);
assert_eq!(Literal::F64(f64::NAN).types(), InferredTypeSet::F64);
assert_eq!(Literal::F64(f64::INFINITY).types(), InferredTypeSet::F64);
assert_eq!(
Literal::F64(f64::NEG_INFINITY).types(),
InferredTypeSet::F64
);
}
}
+13
View File
@@ -0,0 +1,13 @@
mod infer_types;
mod literal;
mod serde;
mod untyped_expr;
pub use infer_types::{InferredTypeSet, TypeError, infer_types, infer_types_with_target};
pub(crate) use infer_types::{infer_type_with_variable_types, infer_types_aux};
pub use literal::Literal;
pub(crate) use serde::format_variable_name;
pub use serde::{DeserializeError, deserialize, serialize};
pub use untyped_expr::UntypedExpr;
pub use crate::functions::{Function, InvalidFnCall};
+688
View File
@@ -0,0 +1,688 @@
//! De/Serialization for [`UntypedExpr`] using a small Lisp-like syntax.
//!
//! Calls are lists whose first item is a recognized uppercase function name.
//! Elsewhere, atoms name variables unless they match a literal. For example:
//!
//! ```text
//! (ADD 1i64 my_col)
//! ```
//!
//! Numerical literals always carry a type suffix. Parsing rejects non-finite
//! f64 literals (NaN, infinities, and overflow). The other literals are
//! none, true, false, and double-quoted strings. Backticks quote variable
//! names containing whitespace or syntax characters, or matching literals.
//! In most case, backticks quote are unnecessary.
//!
//! ```text
//! (ADD `text` 1u64)
//! ```
//! is the same as
//! ```text
//! (ADD text 1u64)
//! ```
//!
//! Quoted variables use escaping to including quotation marks.
use std::fmt;
use std::sync::Arc;
use crate::ast::{Function, Literal, UntypedExpr};
/// Serializes an untyped expression into its canonical Lisp-like form.
pub fn serialize(expr: &UntypedExpr) -> String {
expr.to_string()
}
/// Deserializes an untyped expression from its Lisp-like form.
///
/// Function calls are validated through [`UntypedExpr::call`]; type checking still happens later.
pub fn deserialize(input: &str) -> Result<UntypedExpr, DeserializeError> {
Parser::new(input).parse()
}
/// An error encountered while deserializing an [`UntypedExpr`].
#[derive(Clone, Debug, Eq, PartialEq, thiserror::Error)]
#[error("failed to deserialize expression at byte {offset}: {message}")]
pub struct DeserializeError {
offset: usize,
message: String,
}
impl DeserializeError {
fn new(offset: usize, message: impl Into<String>) -> Self {
Self {
offset,
message: message.into(),
}
}
/// Returns the byte offset at which parsing failed.
pub fn offset(&self) -> usize {
self.offset
}
/// Returns a description of the parsing failure.
pub fn message(&self) -> &str {
&self.message
}
}
impl fmt::Display for UntypedExpr {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
format_expr(self, formatter)
}
}
impl fmt::Debug for UntypedExpr {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
format_expr(self, formatter)
}
}
impl std::str::FromStr for UntypedExpr {
type Err = DeserializeError;
fn from_str(input: &str) -> Result<Self, Self::Err> {
deserialize(input)
}
}
fn format_expr(expr: &UntypedExpr, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match expr {
UntypedExpr::Literal(literal) => format_literal(literal, formatter),
UntypedExpr::Variable(variable_name) => format_variable_name(variable_name, formatter),
UntypedExpr::FnCall { function, args } => {
write!(formatter, "({}", function_name(*function))?;
for arg in args {
write!(formatter, " {arg}")?;
}
formatter.write_str(")")
}
}
}
/// Uses bare names only when unambiguous; otherwise backtick-quotes and escapes them.
pub(crate) fn format_variable_name(
variable_name: &str,
formatter: &mut fmt::Formatter<'_>,
) -> fmt::Result {
if can_format_bare_variable(variable_name) {
return formatter.write_str(variable_name);
}
format_quoted(variable_name, '`', formatter)
}
fn format_literal(literal: &Literal, formatter: &mut fmt::Formatter) -> fmt::Result {
match literal {
Literal::None => formatter.write_str("none"),
Literal::Bool(value) => write!(formatter, "{value}"),
Literal::U64(value) => write!(formatter, "{value}u64"),
Literal::I64(value) => write!(formatter, "{value}i64"),
Literal::F64(value) => write!(formatter, "{value}f64"),
Literal::String(value) => format_quoted(value, '"', formatter),
}
}
fn format_quoted(value: &str, quote: char, formatter: &mut fmt::Formatter) -> fmt::Result {
write!(formatter, "{quote}")?;
for character in value.chars() {
match character {
character if character == quote => write!(formatter, "\\{character}"),
'\\' => formatter.write_str("\\\\"),
'\n' => formatter.write_str("\\n"),
'\r' => formatter.write_str("\\r"),
'\t' => formatter.write_str("\\t"),
'\0' => formatter.write_str("\\0"),
character if character.is_control() => {
write!(formatter, "{}", character.escape_unicode())
}
character => write!(formatter, "{character}"),
}?;
}
write!(formatter, "{quote}")
}
fn function_name(function: Function) -> &'static str {
match function {
Function::Add => "ADD",
Function::IsNull => "IS_NULL",
Function::RegexpExtract => "REGEXP_EXTRACT",
}
}
fn parse_function(name: &str, offset: usize) -> Result<Function, DeserializeError> {
match name {
"ADD" => Ok(Function::Add),
"IS_NULL" => Ok(Function::IsNull),
"REGEXP_EXTRACT" => Ok(Function::RegexpExtract),
_ if !is_function_name(name) => Err(DeserializeError::new(
offset,
format!("function name `{name}` must be uppercase"),
)),
_ => Err(DeserializeError::new(
offset,
format!("unknown function `{name}`"),
)),
}
}
fn is_function_name(name: &str) -> bool {
let mut chars = name.chars();
matches!(chars.next(), Some(first) if first.is_ascii_uppercase())
&& chars.all(|character| {
character.is_ascii_uppercase() || character.is_ascii_digit() || character == '_'
})
}
fn can_format_bare_variable(name: &str) -> bool {
// Quote all recognized literal forms, including non-finite floats that the
// expression parser rejects, so variable names always round-trip.
!name.is_empty()
&& name
.chars()
.all(|c: char| c.is_ascii_alphabetic() || c == '_' || c == '.')
&& parse_literal_atom(name).is_none()
}
// Recognition intentionally includes non-finite floats: parse_atom must reject
// them rather than fall back to variables, and serialization must quote those names.
fn parse_literal_atom(atom: &str) -> Option<Literal> {
match atom {
"none" => return Some(Literal::None),
"true" => return Some(Literal::Bool(true)),
"false" => return Some(Literal::Bool(false)),
_ => {}
}
if let Some(value_str) = atom.strip_suffix("u64") {
let val = value_str.parse::<u64>().ok()?;
return Some(Literal::U64(val));
}
if let Some(value_str) = atom.strip_suffix("i64") {
let val = value_str.parse::<i64>().ok()?;
return Some(Literal::I64(val));
}
if let Some(value_str) = atom.strip_suffix("f64") {
let val = value_str.parse::<f64>().ok()?;
return Some(Literal::F64(val));
}
None
}
struct Parser<'a> {
input: &'a str,
offset: usize,
}
impl<'a> Parser<'a> {
fn new(input: &'a str) -> Self {
Self { input, offset: 0 }
}
fn parse(mut self) -> Result<UntypedExpr, DeserializeError> {
self.skip_whitespace();
let expr = self.parse_expr()?;
self.skip_whitespace();
if self.peek().is_some() {
return Err(DeserializeError::new(
self.offset,
"unexpected characters after expression",
));
}
Ok(expr)
}
fn parse_expr(&mut self) -> Result<UntypedExpr, DeserializeError> {
self.skip_whitespace();
match self.peek() {
Some('(') => self.parse_fn_call(),
Some('"') => self
.parse_quoted('"', "string literal")
.map(|value| UntypedExpr::Literal(Literal::String(Arc::from(value)))),
Some('`') => self
.parse_quoted('`', "quoted variable")
.map(|value| UntypedExpr::Variable(Arc::from(value))),
Some(')') => Err(DeserializeError::new(
self.offset,
"unexpected closing parenthesis",
)),
Some(_) => self.parse_atom(),
None => Err(DeserializeError::new(self.offset, "expected an expression")),
}
}
fn parse_fn_call(&mut self) -> Result<UntypedExpr, DeserializeError> {
let fn_call_offset = self.offset;
self.advance();
self.skip_whitespace();
if self.peek().is_none() {
return Err(DeserializeError::new(
fn_call_offset,
"unterminated function call",
));
}
if self.peek() == Some(')') {
return Err(DeserializeError::new(
self.offset,
"expected a function name",
));
}
let function_offset = self.offset;
let function_name = self.take_atom();
if function_name.is_empty() {
return Err(DeserializeError::new(
function_offset,
"expected an uppercase function name",
));
}
let function = parse_function(function_name, function_offset)?;
let mut args = Vec::new();
loop {
self.skip_whitespace();
match self.peek() {
Some(')') => {
self.advance();
return UntypedExpr::new_fn_call(function, args)
.map_err(|error| DeserializeError::new(fn_call_offset, error.to_string()));
}
Some(_) => args.push(self.parse_expr()?),
None => {
return Err(DeserializeError::new(
fn_call_offset,
"unterminated function call",
));
}
}
}
}
fn parse_atom(&mut self) -> Result<UntypedExpr, DeserializeError> {
let atom_offset = self.offset;
let atom = self.take_atom();
if let Some(literal) = parse_literal_atom(atom) {
if let Literal::F64(value) = &literal
&& !value.is_finite()
{
return Err(DeserializeError::new(
atom_offset,
format!("f64 literal `{atom}` must be finite"),
));
}
return Ok(UntypedExpr::Literal(literal));
}
Ok(UntypedExpr::Variable(Arc::from(atom)))
}
fn parse_quoted(&mut self, quote: char, kind: &str) -> Result<String, DeserializeError> {
let quoted_offset = self.offset;
self.advance();
let mut value = String::new();
loop {
let character_offset = self.offset;
let Some(character) = self.advance() else {
return Err(DeserializeError::new(
quoted_offset,
format!("unterminated {kind}"),
));
};
match character {
character if character == quote => return Ok(value),
'\\' => value.push(self.parse_escape(character_offset, quote)?),
character if character.is_control() => {
return Err(DeserializeError::new(
character_offset,
format!("unescaped control character in {kind}"),
));
}
character => value.push(character),
}
}
}
fn parse_escape(
&mut self,
escape_offset: usize,
quote: char,
) -> Result<char, DeserializeError> {
let Some(escaped) = self.advance() else {
return Err(DeserializeError::new(
escape_offset,
"unterminated string escape",
));
};
match escaped {
'`' if quote == '`' => Ok('`'),
'"' => Ok('"'),
'\\' => Ok('\\'),
'n' => Ok('\n'),
'r' => Ok('\r'),
't' => Ok('\t'),
'0' => Ok('\0'),
'u' => self.parse_unicode_escape(escape_offset),
_ => Err(DeserializeError::new(
escape_offset,
format!("unsupported string escape `\\{escaped}`"),
)),
}
}
fn parse_unicode_escape(&mut self, escape_offset: usize) -> Result<char, DeserializeError> {
if self.advance() != Some('{') {
return Err(DeserializeError::new(
escape_offset,
"Unicode escape must start with `\\u{`",
));
}
let digits_offset = self.offset;
while matches!(self.peek(), Some(character) if character.is_ascii_hexdigit()) {
self.advance();
}
let digits = &self.input[digits_offset..self.offset];
if digits.is_empty() || self.advance() != Some('}') {
return Err(DeserializeError::new(
escape_offset,
"invalid Unicode escape",
));
}
let codepoint = u32::from_str_radix(digits, 16).ok();
codepoint
.and_then(char::from_u32)
.ok_or_else(|| DeserializeError::new(escape_offset, "invalid Unicode scalar value"))
}
fn take_atom(&mut self) -> &'a str {
let start = self.offset;
while matches!(self.peek(), Some(character) if !is_delimiter(character)) {
self.advance();
}
&self.input[start..self.offset]
}
fn skip_whitespace(&mut self) {
while matches!(self.peek(), Some(character) if character.is_whitespace()) {
self.advance();
}
}
fn peek(&self) -> Option<char> {
self.input[self.offset..].chars().next()
}
fn advance(&mut self) -> Option<char> {
let character = self.peek()?;
self.offset += character.len_utf8();
Some(character)
}
}
fn is_delimiter(character: char) -> bool {
character.is_whitespace() || matches!(character, '(' | ')' | '"' | '`')
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_serialize_example() {
let expr = UntypedExpr::new_fn_call(
Function::Add,
vec![UntypedExpr::literal(1i64), UntypedExpr::variable("my_col")],
)
.unwrap();
assert_eq!(serialize(&expr), "(ADD 1i64 my_col)");
assert_eq!(format!("{expr}"), "(ADD 1i64 my_col)");
assert_eq!(format!("{expr:?}"), "(ADD 1i64 my_col)");
}
#[test]
fn test_serialize_literals() {
let cases = [
(UntypedExpr::Literal(Literal::None), "none"),
(UntypedExpr::literal(true), "true"),
(UntypedExpr::literal(false), "false"),
(UntypedExpr::literal(u64::MAX), "18446744073709551615u64"),
(UntypedExpr::literal(i64::MIN), "-9223372036854775808i64"),
(UntypedExpr::literal(1.5f64), "1.5f64"),
(UntypedExpr::literal(1.0f64), "1f64"),
];
for (expr, expected) in cases {
assert_eq!(serialize(&expr), expected);
assert_eq!(deserialize(expected).unwrap(), expr);
}
}
#[test]
fn test_nested_call_and_escaped_string_round_trip() {
let string = "quoted: \"hello\"\\world\n\t\0\u{7} café";
let regexp_extract = UntypedExpr::new_fn_call(
Function::RegexpExtract,
vec![
UntypedExpr::variable("message"),
UntypedExpr::literal(string),
UntypedExpr::literal(1u64),
],
)
.unwrap();
let expr = UntypedExpr::new_fn_call(
Function::Add,
vec![regexp_extract, UntypedExpr::literal(2i64)],
)
.unwrap();
let serialized = serialize(&expr);
assert_eq!(
serialized,
"(ADD (REGEXP_EXTRACT message \"quoted: \\\"hello\\\"\\\\world\\n\\t\\0\\u{7} café\" \
1u64) 2i64)"
);
assert_eq!(deserialize(&serialized).unwrap(), expr);
}
#[test]
fn test_deserialize_accepts_whitespace() {
let parsed = deserialize(" \n ( ADD\t1i64\nmy_col ) \r").unwrap();
let expected = UntypedExpr::new_fn_call(
Function::Add,
vec![UntypedExpr::literal(1i64), UntypedExpr::variable("my_col")],
)
.unwrap();
assert_eq!(parsed, expected);
}
#[test]
fn test_finite_float_edge_values_round_trip() {
for value in [
f64::MIN,
f64::MAX,
f64::MIN_POSITIVE,
f64::from_bits(1),
-0.0,
] {
let serialized = serialize(&UntypedExpr::literal(value));
let UntypedExpr::Literal(Literal::F64(parsed)) = deserialize(&serialized).unwrap()
else {
panic!("expected an f64 literal");
};
assert_eq!(parsed.to_bits(), value.to_bits());
}
}
#[test]
fn test_non_finite_float_literals_are_rejected() {
for atom in [
"NaNf64",
"nanf64",
"+NaNf64",
"-NaNf64",
"inff64",
"+inff64",
"-inff64",
"infinityf64",
"-INFINITYf64",
"1e999f64",
"-1e999f64",
] {
for (input, offset) in [(atom.to_string(), 0), (format!("(ADD 1u64 {atom})"), 10)] {
let error = deserialize(&input).unwrap_err();
assert_eq!(error.offset(), offset, "input: {input}");
assert_eq!(
error.message(),
format!("f64 literal `{atom}` must be finite")
);
}
// Rejected literal spellings are still usable as quoted field names.
let variable = UntypedExpr::variable(atom);
let serialized = serialize(&variable);
assert_eq!(serialized, format!("`{atom}`"));
assert_eq!(deserialize(&serialized).unwrap(), variable);
}
}
#[test]
fn test_from_str() {
let parsed: UntypedExpr = "(ADD 3u64 value)".parse().unwrap();
assert_eq!(serialize(&parsed), "(ADD 3u64 value)");
}
#[test]
fn test_bare_field_names_round_trip() {
for name in [
"HTTP.Status",
"@timestamp",
"_source",
"field-name",
"field/path:part",
"シャボン玉",
"café",
"🦀",
"ADD",
"IS_NULL",
"1i32",
"123",
"18446744073709551616u64",
] {
let expr = UntypedExpr::variable(name);
assert!(serialize(&expr).contains(name));
assert_eq!(deserialize(name).unwrap(), expr);
let call = UntypedExpr::new_fn_call(Function::Add, vec![expr]).unwrap();
let serialized = format!("(ADD {name})");
assert_eq!(deserialize(&serialized).unwrap(), call);
}
}
#[test]
fn test_literal_names_are_quoted() {
for name in [
"none", "true", "false", "1u64", "1i64", "1f64", "+1u64", "1e3f64", "-0f64",
] {
assert!(matches!(
deserialize(name).unwrap(),
UntypedExpr::Literal(_)
));
let expr = UntypedExpr::variable(name);
let serialized = format!("`{name}`");
assert_eq!(serialize(&expr), serialized);
assert_eq!(deserialize(&serialized).unwrap(), expr);
}
let expr = UntypedExpr::new_fn_call(
Function::Add,
vec![UntypedExpr::variable("1u64"), UntypedExpr::literal(1u64)],
)
.unwrap();
assert_eq!(serialize(&expr), "(ADD `1u64` 1u64)");
assert_eq!(deserialize("(ADD `1u64` 1u64)").unwrap(), expr);
}
#[test]
fn test_quoted_field_names_round_trip() {
let cases = [
("", "``"),
("two words", "`two words`"),
("(field)", "`(field)`"),
("a\"b", "`a\"b`"),
("a`b", "`a\\`b`"),
("a\\b", "`a\\\\b`"),
("a\n\r\t\0\u{7}", "`a\\n\\r\\t\\0\\u{7}`"),
("a\u{2003}b", "`a\u{2003}b`"),
];
for (name, serialized) in cases {
let expr = UntypedExpr::variable(name);
assert_eq!(serialize(&expr), serialized);
assert_eq!(format!("{expr:?}"), serialized);
assert_eq!(deserialize(serialized).unwrap(), expr);
let call = UntypedExpr::new_fn_call(Function::IsNull, vec![expr]).unwrap();
assert_eq!(deserialize(&serialize(&call)).unwrap(), call);
}
// Quoting does not force the canonical serializer to retain quotes.
assert_eq!(
serialize(&deserialize("`HTTP.Status`").unwrap()),
"HTTP.Status"
);
assert_eq!(
deserialize(r"`\u{30b7}\u{30e3}`").unwrap(),
UntypedExpr::variable("シャ")
);
// Backticks inside double quotes still belong to a string literal.
let string = UntypedExpr::literal("`field`");
assert_eq!(deserialize(&serialize(&string)).unwrap(), string);
}
#[test]
fn test_deserialize_rejects_invalid_fn_calls() {
let input = "(IS_NULL)";
let expected_message = "invalid number of arguments";
let error = deserialize(input).unwrap_err();
assert!(
error.message().contains(expected_message),
"input: {input}; error: {error}"
);
}
#[test]
fn test_deserialize_errors() {
let cases = [
("", 0, "expected an expression"),
("()", 1, "expected a function name"),
("(add 1i64)", 1, "must be uppercase"),
("(UNKNOWN 1i64)", 1, "unknown function"),
("(toto)", 1, "must be uppercase"),
("(`ADD` 1u64)", 1, "expected an uppercase function name"),
("`unterminated", 0, "unterminated quoted variable"),
("`bad\\x`", 4, "unsupported string escape"),
("`é\\x`", 3, "unsupported string escape"),
("`bad\\", 4, "unterminated string escape"),
(r"`\u{d800}`", 1, "invalid Unicode scalar value"),
(r"`\u{}`", 1, "invalid Unicode escape"),
(
"`bad\n`",
4,
"unescaped control character in quoted variable",
),
("\"unterminated", 0, "unterminated string literal"),
("\"bad\\x\"", 4, "unsupported string escape"),
("(ADD 1i64", 0, "unterminated function call"),
("value other", 6, "unexpected characters after expression"),
];
for (input, offset, expected_message) in cases {
let error = deserialize(input).unwrap_err();
assert_eq!(error.offset(), offset, "input: {input}");
assert!(
error.message().contains(expected_message),
"input: {input}; error: {error}"
);
}
}
}
+43
View File
@@ -0,0 +1,43 @@
use std::sync::Arc;
use crate::ast::{Function, Literal};
use crate::functions::InvalidFnCall;
/// An expression AST.
///
/// The expression at this point is untyped and not necessarily valid.
#[derive(Clone, PartialEq)]
pub enum UntypedExpr {
Literal(Literal),
Variable(Arc<str>),
FnCall {
function: Function,
args: Vec<UntypedExpr>,
},
}
impl UntypedExpr {
pub fn literal(val: impl Into<Literal>) -> UntypedExpr {
UntypedExpr::Literal(val.into())
}
pub fn variable(variable_name: impl ToString) -> UntypedExpr {
UntypedExpr::Variable(Arc::from(variable_name.to_string()))
}
/// Creates an untyped expression that is a function over different arguments.
///
/// This call will validate the arguments and
pub fn new_fn_call(
function: Function,
args: Vec<UntypedExpr>,
) -> Result<UntypedExpr, InvalidFnCall> {
function.call(args)
}
}
impl From<Literal> for UntypedExpr {
fn from(literal: Literal) -> Self {
UntypedExpr::Literal(literal)
}
}
+512
View File
@@ -0,0 +1,512 @@
use std::collections::HashMap;
use std::mem;
use std::sync::Arc;
use cranelift::codegen::Context as CodegenContext;
use cranelift::codegen::control::ControlPlane;
use cranelift::codegen::ir::{MemFlagsData, UserFuncName};
use cranelift::prelude::*;
use cranelift_jit::{JITBuilder, JITModule};
use cranelift_module::{FuncId, Module, ModuleError, default_libcall_names};
use super::compiled_fn::JitEntry;
use super::{
CompileError, CompiledFn, LoweringContext, TypedExpr, TypedExprAst, TypedLiteral, TypedVariable,
};
use crate::ast::{InferredTypeSet, Literal, UntypedExpr};
use crate::functions::{declare_native_functions, register_jit_symbols};
use crate::types::VarType;
pub(crate) struct CompileFnBuilder<'types, 'names> {
variable_types: &'types HashMap<&'names str, VarType>,
input_vars: Vec<TypedVariable>,
}
struct LoweredFunction {
module: JITModule,
context: CodegenContext,
function_id: FuncId,
input_vars: Vec<TypedVariable>,
expression: Box<TypedExpr>,
}
fn make_jit_builder() -> Result<JITBuilder, CompileError> {
let mut shared_flags = settings::builder();
shared_flags
.set("opt_level", "speed")
.map_err(ModuleError::from)?;
shared_flags
.set("use_colocated_libcalls", "false")
.map_err(ModuleError::from)?;
shared_flags
.set("is_pic", "false")
.map_err(ModuleError::from)?;
let isa_builder = cranelift_native::builder().unwrap_or_else(|message| {
panic!("host machine is not supported: {message}");
});
#[cfg(target_arch = "aarch64")]
let isa_builder = {
let mut isa_builder = isa_builder;
isa_builder
.set("sign_return_address", "false")
.map_err(ModuleError::from)?;
isa_builder
.set("sign_return_address_all", "false")
.map_err(ModuleError::from)?;
isa_builder
};
let isa = isa_builder
.finish(settings::Flags::new(shared_flags))
.map_err(ModuleError::from)?;
Ok(JITBuilder::with_isa(isa, default_libcall_names()))
}
impl<'types, 'names> CompileFnBuilder<'types, 'names> {
pub(crate) fn new(variable_types: &'types HashMap<&'names str, VarType>) -> Self {
CompileFnBuilder {
variable_types,
input_vars: Vec::new(),
}
}
pub(crate) fn variable_types(&self) -> &HashMap<&'names str, VarType> {
self.variable_types
}
/// If a variable is missing from `variable_types`, it is treated as `None`.
pub(crate) fn build_typed_expr(
&mut self,
untyped_expr: &UntypedExpr,
) -> Result<TypedExpr, CompileError> {
let mut typed_expr = self.apply_types(untyped_expr, InferredTypeSet::ALL)?;
self.assign_variable_ids(&mut typed_expr);
Ok(typed_expr)
}
pub(crate) fn assign_variable_ids(&mut self, typed_expr: &mut TypedExpr) {
self.input_vars = assign_variable_ids(typed_expr);
}
fn apply_literal_type(
&mut self,
literal: &Literal,
target_type_set: InferredTypeSet,
) -> TypedLiteral {
if literal.is_none() {
return TypedLiteral::None;
}
let inferred_type_set = literal.types();
let intersection = inferred_type_set.intersect(target_type_set);
if intersection.contains(VarType::Bool) {
match literal {
Literal::Bool(value) => TypedLiteral::Bool(*value),
_ => panic!("cannot coerce literal {literal:?} to bool"),
}
} else if intersection.contains(VarType::I64) {
match literal {
Literal::U64(value) => TypedLiteral::I64(*value as i64),
Literal::I64(value) => TypedLiteral::I64(*value),
Literal::F64(value) if f64_to_i64_lossless(*value).is_some() => {
TypedLiteral::I64(f64_to_i64_lossless(*value).unwrap())
}
_ => panic!("cannot coerce literal {literal:?} to i64"),
}
} else if intersection.contains(VarType::U64) {
match literal {
Literal::U64(value) => TypedLiteral::U64(*value),
Literal::I64(value) => TypedLiteral::U64(*value as u64),
Literal::F64(value) if f64_to_u64_lossless(*value).is_some() => {
TypedLiteral::U64(f64_to_u64_lossless(*value).unwrap())
}
_ => panic!("cannot coerce literal {literal:?} to u64"),
}
} else if intersection.contains(VarType::F64) {
match literal {
Literal::U64(value) => TypedLiteral::F64(*value as f64),
Literal::I64(value) => TypedLiteral::F64(*value as f64),
Literal::F64(value) => TypedLiteral::F64(*value),
_ => panic!("cannot coerce literal {literal:?} to f64"),
}
} else if intersection.contains(VarType::Str) {
match literal {
Literal::String(value) => TypedLiteral::String(value.clone()),
_ => panic!("cannot coerce literal {literal:?} to string"),
}
} else if intersection.contains(VarType::None) {
match literal {
Literal::None => TypedLiteral::None,
_ => panic!("cannot coerce literal {literal:?} to none"),
}
} else {
panic!(
"no compatible type for literal {literal:?} with target type set \
{target_type_set:?}"
)
}
}
pub(crate) fn apply_types(
&mut self,
untyped_expr: &UntypedExpr,
target_type_set: InferredTypeSet,
) -> Result<TypedExpr, CompileError> {
match untyped_expr {
UntypedExpr::Literal(literal) => {
let typed_literal = self.apply_literal_type(literal, target_type_set);
let return_type = typed_literal.r#type();
Ok(TypedExpr {
return_type,
ast: TypedExprAst::Literal(typed_literal),
})
}
UntypedExpr::Variable(variable_name) => {
let variable_type = self
.variable_types
.get(variable_name.as_ref())
.copied()
.unwrap_or(VarType::None);
if variable_type == VarType::None {
Ok(TypedExpr {
return_type: VarType::None,
ast: TypedExprAst::Literal(TypedLiteral::None),
})
} else {
let typed_expr = TypedExpr {
return_type: variable_type,
ast: TypedExprAst::variable(variable_name, variable_type),
};
if target_type_set.contains(variable_type) {
Ok(typed_expr)
} else if let Some(target_type) = preferred_numerical_type(target_type_set)
&& is_numerical(variable_type)
{
Ok(typed_expr.coerce(target_type))
} else {
Ok(typed_expr)
}
}
}
UntypedExpr::FnCall { function, args } => {
function.call_with_types(args, target_type_set, self)
}
}
}
pub(super) fn compile_typed_expr(
self,
expression: TypedExpr,
) -> Result<CompiledFn, CompileError> {
self.lower_typed_expr(expression)?.into_compiled_fn()
}
pub(super) fn compile_typed_expr_to_assembly(
self,
expression: TypedExpr,
) -> Result<String, CompileError> {
self.lower_typed_expr(expression)?.into_assembly()
}
fn lower_typed_expr(self, expression: TypedExpr) -> Result<LoweredFunction, CompileError> {
let CompileFnBuilder { input_vars, .. } = self;
let expression = Box::new(expression);
let mut jit_builder = make_jit_builder()?;
register_jit_symbols(&mut jit_builder);
let mut module = JITModule::new(jit_builder);
let target_config = module.target_config();
let pointer_type = target_config.pointer_type();
// The native entry point mirrors JitEntry: its arguments point to the
// input slots and call-scoped string arena. 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)?;
let mut context = module.make_context();
context.func.signature = signature;
context.func.name = UserFuncName::user(0, function_id.as_u32());
let native_functions =
declare_native_functions(&mut module, &mut context.func, pointer_type)?;
let mut function_builder_context = FunctionBuilderContext::new();
{
let mut builder =
FunctionBuilder::new(&mut context.func, &mut function_builder_context);
let entry_block = builder.create_block();
builder.append_block_params_for_function_params(entry_block);
builder.switch_to_block(entry_block);
builder.seal_block(entry_block);
let args_ptr = builder.block_params(entry_block)[0];
let string_arena_ptr = builder.block_params(entry_block)[1];
let mut lowering_context = LoweringContext {
args_ptr,
string_arena_ptr,
string_arena_was_reset: false,
pointer_type,
native_functions: &native_functions,
};
let lowered = lowering_context.compile_expr(&expression, &mut builder)?;
let value_bits = match expression.return_type {
VarType::Bool => builder.ins().uextend(types::I64, lowered.value),
VarType::F64 => {
builder
.ins()
.bitcast(types::I64, MemFlagsData::new(), lowered.value)
}
VarType::U64 | VarType::I64 | VarType::Str | VarType::None => lowered.value,
};
let second_word = if expression.return_type == VarType::Str {
lowered.string_len
} else {
builder.ins().uextend(types::I64, lowered.is_present)
};
builder.ins().return_(&[value_bits, second_word]);
builder.finalize(target_config);
}
Ok(LoweredFunction {
module,
context,
function_id,
input_vars,
expression,
})
}
}
fn preferred_numerical_type(inferred_types: InferredTypeSet) -> Option<VarType> {
if inferred_types.i64 {
Some(VarType::I64)
} else if inferred_types.u64 {
Some(VarType::U64)
} else if inferred_types.f64 {
Some(VarType::F64)
} else {
None
}
}
fn is_numerical(var_type: VarType) -> bool {
matches!(var_type, VarType::I64 | VarType::U64 | VarType::F64)
}
impl LoweredFunction {
fn into_compiled_fn(self) -> Result<CompiledFn, CompileError> {
let LoweredFunction {
mut module,
mut context,
function_id,
input_vars,
expression,
} = self;
module.define_function(function_id, &mut context)?;
module.finalize_definitions()?;
let code = module.get_finalized_function(function_id);
// SAFETY: `code` is the finalized entry point for the function whose ABI
// was built above to exactly match `JitEntry`. The module is retained by
// `CompiledFn`, so its executable allocation outlives `entry`.
let entry = unsafe { mem::transmute::<*const u8, JitEntry>(code) };
Ok(CompiledFn {
entry,
_module: module,
inputs: input_vars,
_typed_expr: expression,
})
}
fn into_assembly(mut self) -> Result<String, CompileError> {
self.context.set_disasm(true);
let compiled_code = self
.context
.compile(self.module.isa(), &mut ControlPlane::default())
.map_err(ModuleError::from)?;
Ok(compiled_code
.vcode
.clone()
.expect("Cranelift assembly was requested before compilation"))
}
}
/// Converts an `f64` to an `i64` only when the value can be represented exactly.
fn f64_to_i64_lossless(value: f64) -> Option<i64> {
let is_integral = value.is_finite() && value.fract() == 0.0;
if is_integral && value >= i64::MIN as f64 && value < -(i64::MIN as f64) {
Some(value as i64)
} else {
None
}
}
/// Converts an `f64` to a `u64` only when the value can be represented exactly.
fn f64_to_u64_lossless(value: f64) -> Option<u64> {
let is_integral = value.is_finite() && value.fract() == 0.0;
if is_integral && value >= 0.0 && value < u64::MAX as f64 {
Some(value as u64)
} else {
None
}
}
fn assign_variable_ids(expr: &mut TypedExpr) -> Vec<TypedVariable> {
let mut name_to_vars: HashMap<Arc<str>, TypedVariable> = HashMap::new();
assign_variable_ids_aux(&mut expr.ast, &mut name_to_vars);
let mut input_vars: Vec<TypedVariable> = name_to_vars.into_values().collect();
input_vars.sort_by_key(|var| var.variable_id);
input_vars
}
fn assign_variable_ids_aux(
ast: &mut TypedExprAst,
name_to_vars: &mut HashMap<Arc<str>, TypedVariable>,
) {
match ast {
TypedExprAst::Literal(_) => {}
TypedExprAst::Variable(var) => {
if let Some(typed_var) = name_to_vars.get(&var.variable_name) {
assert_eq!(
typed_var.r#type, var.r#type,
"variable `{}` appears with two different types (`{:?}` and `{:?}`); a typed \
expr AST must be built with a single explicit type per variable",
var.variable_name, typed_var.r#type, var.r#type,
);
var.variable_id = typed_var.variable_id;
} else {
var.variable_id = name_to_vars.len();
name_to_vars.insert(var.variable_name.clone(), var.clone());
};
}
TypedExprAst::Coerce { expr, .. } => {
assign_variable_ids_aux(&mut expr.ast, name_to_vars);
}
TypedExprAst::FnCall(fn_call) => {
for arg in fn_call.args_mut() {
assign_variable_ids_aux(&mut arg.ast, name_to_vars);
}
}
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use super::*;
use crate::ast::Function;
use crate::functions::{AddFnCall, FnCallEnum};
#[test]
fn test_apply_types_to_literal() {
let untyped_expr = UntypedExpr::literal("hello");
let variable_types = HashMap::new();
let mut builder = CompileFnBuilder::new(&variable_types);
let typed_expr = builder.build_typed_expr(&untyped_expr).unwrap();
assert_eq!(typed_expr.return_type, VarType::Str);
let TypedExprAst::Literal(TypedLiteral::String(value)) = typed_expr.ast else {
panic!("expected a typed string literal");
};
assert_eq!(value.as_ref(), "hello");
}
#[test]
fn test_assign_variable_ids_two_variables_different_types() {
let untyped_expr = UntypedExpr::new_fn_call(
Function::Add,
vec![UntypedExpr::variable("x"), UntypedExpr::variable("y")],
)
.unwrap();
let variable_types = HashMap::from([("x", VarType::U64), ("y", VarType::F64)]);
let mut builder = CompileFnBuilder::new(&variable_types);
let _typed_expr = builder.build_typed_expr(&untyped_expr).unwrap();
let var_args = &builder.input_vars;
assert_eq!(var_args.len(), 2);
assert_eq!(var_args[0].variable_name.as_ref(), "x");
assert_eq!(var_args[0].r#type, VarType::U64);
assert_eq!(var_args[0].variable_id, 0);
assert_eq!(var_args[1].variable_name.as_ref(), "y");
assert_eq!(var_args[1].r#type, VarType::F64);
assert_eq!(var_args[1].variable_id, 1);
}
#[test]
#[should_panic(expected = "appears with two different types")]
fn test_assign_variable_ids_panics_on_inconsistent_types() {
let mut typed_expr = TypedExpr {
return_type: VarType::F64,
ast: TypedExprAst::FnCall(FnCallEnum::Add(AddFnCall {
args: vec![
TypedExprAst::variable("x", VarType::U64).with_type(VarType::U64),
TypedExprAst::variable("x", VarType::F64).with_type(VarType::F64),
]
.into_boxed_slice(),
})),
};
assign_variable_ids(&mut typed_expr);
}
#[test]
fn test_assign_variable_ids_dedups_repeated_variable() {
let untyped_expr = UntypedExpr::new_fn_call(
Function::Add,
vec![
UntypedExpr::variable("x"),
UntypedExpr::new_fn_call(
Function::Add,
vec![UntypedExpr::variable("y"), UntypedExpr::variable("x")],
)
.unwrap(),
],
)
.unwrap();
let variable_types: HashMap<&str, VarType> =
HashMap::from([("x", VarType::U64), ("y", VarType::U64)]);
let mut builder = CompileFnBuilder::new(&variable_types);
let _typed_expr = builder.build_typed_expr(&untyped_expr).unwrap();
let var_args = &builder.input_vars;
assert_eq!(var_args.len(), 2);
assert_eq!(var_args[0].variable_name.as_ref(), "x");
assert_eq!(var_args[0].r#type, VarType::U64);
assert_eq!(var_args[0].variable_id, 0);
assert_eq!(var_args[1].variable_name.as_ref(), "y");
assert_eq!(var_args[1].r#type, VarType::U64);
assert_eq!(var_args[1].variable_id, 1);
}
#[test]
fn test_jit_entry_returns_variable_value_as_two_abi_words() {
let variable_types = HashMap::new();
let mut builder = CompileFnBuilder::new(&variable_types);
let expression = builder
.build_typed_expr(&UntypedExpr::literal(1u64))
.unwrap();
let lowered = builder.lower_typed_expr(expression).unwrap();
let signature = &lowered.context.func.signature;
assert_eq!(signature.params.len(), 2);
assert!(
signature
.params
.iter()
.all(|param| param.value_type == types::I64)
);
assert_eq!(signature.returns.len(), 2);
assert_eq!(signature.returns[0].value_type, types::I64);
assert_eq!(signature.returns[1].value_type, types::I64);
}
}
+158
View File
@@ -0,0 +1,158 @@
use std::ops::Deref;
use std::sync::Arc;
use cranelift_jit::JITModule;
use super::{StringArena, TypedExpr, TypedVariable};
use crate::types::{VarType, VariableValue};
#[cfg(not(any(
all(target_arch = "x86_64", not(target_os = "windows")),
target_arch = "aarch64"
)))]
// Windows is not supported because apparently returning more than one 64 bits word throught
// registers is not supported by its ABI.
compile_error!(
"the direct VariableValue JIT return ABI is only implemented for x86-64 System V and AArch64"
);
// On the supported targets, VariableValue's two eightbytes are returned in two
// integer registers by the platform C ABI. The lifetime is selected by
// CompiledFn::call so that it is bounded by all possible sources of strings.
// This is a Rust-to-JIT boundary whose VariableValue layout is asserted in
// 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>, *mut StringArena) -> VariableValue<'a>;
/// An expression compiled to native machine code.
///
/// This object owns the JIT module containing its executable memory and every
/// resource referenced by the generated code.
pub struct CompiledFn {
pub(crate) entry: JitEntry,
pub(crate) _module: JITModule,
pub(super) inputs: Vec<TypedVariable>,
// This AST owns the Arc-backed literals and regexes embedded in generated code.
pub(crate) _typed_expr: Box<TypedExpr>,
}
// `JITModule` is not `Sync` because it supports lazily looking up symbols through
// interior mutability. A `CompiledFn` only retains a finalized module to keep its
// executable allocation alive and never invokes those mutable APIs. Its entry
// point and the immutable resources referenced by the generated code can be
// called concurrently when each caller supplies a distinct `StringArena`.
unsafe impl Sync for CompiledFn {}
impl CompiledFn {
/// Returns the input slots in the exact order expected by [`CompiledFn::call`].
pub fn inputs(&self) -> &[TypedVariable] {
&self.inputs
}
/// Returns the concrete result type selected during compilation.
pub fn result_type(&self) -> VarType {
self._typed_expr.return_type
}
/// Creates an evaluation context with a private string arena.
pub fn context(self: &Arc<Self>) -> CompiledFnCtx {
CompiledFnCtx::new(Arc::clone(self))
}
/// Evaluates the compiled expression using the supplied string arena.
///
/// The mutable arena borrow prevents another evaluation from clearing the
/// arena while an arena-backed result from this call is still live.
///
/// # Safety
///
/// `args` must follow [`CompiledFn::inputs`] exactly: every present slot must
/// contain the union member corresponding to that variable's type. Absent
/// slots must use [`VariableValue::none`], and any borrowed strings must
/// remain alive for the duration of this call.
///
/// The result cannot outlive the compiled function, the string arena, or
/// the passed arguments' lifetime.
#[inline(always)]
pub unsafe fn call<'args, 'compiled, 'arena, 'output>(
&'compiled self,
args: &[VariableValue<'args>],
string_arena: &'arena mut StringArena,
) -> VariableValue<'output>
where
'args: 'output,
'compiled: 'output,
'arena: 'output,
{
debug_assert_eq!(args.len(), self.inputs.len());
let args: &[VariableValue<'output>] = args;
let string_arena = &raw mut *string_arena;
// SAFETY: Guaranteed by the caller. The input, compiled-function, and
// arena lifetimes outlive the lifetime selected for the returned value.
unsafe { (self.entry)(args.as_ptr(), string_arena) }
}
}
/// Per-caller mutable state used to evaluate a shared [`CompiledFn`].
pub struct CompiledFnCtx {
compiled_fn: Arc<CompiledFn>,
pub(crate) string_arena: StringArena,
}
impl CompiledFnCtx {
/// Creates an evaluation context for `compiled_fn`.
pub fn new(compiled_fn: Arc<CompiledFn>) -> Self {
Self {
compiled_fn,
string_arena: StringArena::new(),
}
}
/// Evaluates the compiled expression using this context's string arena.
///
/// The mutable borrow prevents another evaluation from clearing the string
/// arena while an arena-backed result from this call is still live.
///
/// # Safety
///
/// `args` must follow [`CompiledFn::inputs`] exactly: every present slot must
/// contain the union member corresponding to that variable's type. Absent
/// slots must use [`VariableValue::none`], and any borrowed strings must
/// remain alive for the duration of this call.
///
/// The result cannot outlive this context nor the passed arguments'
/// lifetime.
#[inline(always)]
pub unsafe fn call<'args, 'ctx, 'output>(
&'ctx mut self,
args: &[VariableValue<'args>],
) -> VariableValue<'output>
where
'args: 'output,
'ctx: 'output,
{
// SAFETY: Guaranteed by the caller. The context owns both the compiled
// function and arena for the lifetime selected for the returned value.
unsafe { self.compiled_fn.call(args, &mut self.string_arena) }
}
/// Returns the shared compiled expression owned by this context.
pub fn compiled_fn(&self) -> &Arc<CompiledFn> {
&self.compiled_fn
}
}
impl From<Arc<CompiledFn>> for CompiledFnCtx {
fn from(compiled_fn: Arc<CompiledFn>) -> Self {
Self::new(compiled_fn)
}
}
impl Deref for CompiledFnCtx {
type Target = CompiledFn;
fn deref(&self) -> &Self::Target {
&self.compiled_fn
}
}
+31
View File
@@ -0,0 +1,31 @@
use crate::ast::{Function, InvalidFnCall, TypeError};
use crate::types::VarType;
#[derive(Debug, thiserror::Error)]
pub enum CompileError {
#[error("type inference failed: {0}")]
TypeInference(#[from] TypeError),
#[error("JIT compilation failed: {0}")]
Module(#[source] Box<cranelift_module::ModuleError>),
#[error("cannot coerce an expression from {from_type:?} to {target:?}")]
UnsupportedCoercion { from_type: VarType, target: VarType },
#[error("cannot compile {function:?} with result type {return_type:?}")]
UnsupportedFunctionType {
function: Function,
return_type: VarType,
},
#[error("invalid regular expression `{pattern}`: {source}")]
InvalidRegex {
pattern: String,
#[source]
source: regex::Error,
},
#[error("arguments do not match the function {0}")]
InvalidArguments(#[from] InvalidFnCall),
}
impl From<cranelift_module::ModuleError> for CompileError {
fn from(error: cranelift_module::ModuleError) -> Self {
CompileError::Module(Box::new(error))
}
}
+353
View File
@@ -0,0 +1,353 @@
mod compile_fn_builder;
mod compiled_fn;
mod error;
mod string_arena;
mod typed_expr;
mod typed_expr_serialize;
use std::collections::HashMap;
use std::sync::Arc;
pub(crate) use compile_fn_builder::CompileFnBuilder;
pub use compiled_fn::{CompiledFn, CompiledFnCtx};
use cranelift::codegen::ir::{
InstBuilder as _, MemFlagsData, Type, Value as CraneliftValue, types as cranelift_types,
};
use cranelift::frontend::FunctionBuilder;
pub use error::CompileError;
pub use string_arena::StringArena;
pub use typed_expr::TypedVariable;
pub(crate) use typed_expr::{TypedExpr, TypedExprAst, TypedLiteral};
pub(crate) use typed_expr_serialize::{format_fn_call, format_string_literal};
use crate::ast::UntypedExpr;
use crate::functions::NativeFunctions;
use crate::types::{VarType, VariablePrimitiveOpt, VariableValue};
/// Compiles an expression into an immutable, shareable native function.
pub fn compile(
untyped_expr: &UntypedExpr,
var_types: &HashMap<&str, VarType>,
) -> Result<Arc<CompiledFn>, CompileError> {
let mut builder = CompileFnBuilder::new(var_types);
let typed_expr = builder.build_typed_expr(untyped_expr)?;
builder.compile_typed_expr(typed_expr).map(Arc::new)
}
/// Applies concrete variable types and serializes the resulting typed expression.
pub fn serialize(
untyped_expr: &UntypedExpr,
var_types: &HashMap<&str, VarType>,
) -> Result<String, CompileError> {
let mut builder = CompileFnBuilder::new(var_types);
let typed_expr = builder.build_typed_expr(untyped_expr)?;
Ok(typed_expr_serialize::serialize(&typed_expr))
}
/// Compiles an expression and returns Cranelift's assembly listing for the host target.
pub fn compile_to_assembly(
untyped_expr: &UntypedExpr,
var_types: &HashMap<&str, VarType>,
) -> Result<String, CompileError> {
let mut builder = CompileFnBuilder::new(var_types);
let typed_expr = builder.build_typed_expr(untyped_expr)?;
builder.compile_typed_expr_to_assembly(typed_expr)
}
pub(crate) struct LoweringContext<'a> {
args_ptr: CraneliftValue,
string_arena_ptr: CraneliftValue,
string_arena_was_reset: bool,
pointer_type: Type,
native_functions: &'a NativeFunctions,
}
/// The two SSA values used to represent a nullable expression result.
#[derive(Clone, Copy)]
pub(crate) struct LoweredValue {
pub(crate) value: CraneliftValue,
pub(crate) is_present: CraneliftValue,
pub(crate) string_len: CraneliftValue,
}
impl LoweringContext<'_> {
pub(crate) fn compile_expr(
&mut self,
expression: &TypedExpr,
builder: &mut FunctionBuilder<'_>,
) -> Result<LoweredValue, CompileError> {
match &expression.ast {
TypedExprAst::Literal(literal) => Ok(lower_literal(literal, self, builder)),
TypedExprAst::Variable(variable) => {
let slot_offset = variable.variable_id * std::mem::size_of::<VariableValue>();
let value_offset = slot_offset as i32;
let second_word_offset =
(slot_offset + std::mem::offset_of!(VariablePrimitiveOpt, is_present)) as i32;
let value = builder.ins().load(
cranelift_type(variable.r#type, self.pointer_type),
MemFlagsData::trusted(),
self.args_ptr,
value_offset,
);
let (is_present, string_len) = if variable.r#type == VarType::Str {
let string_len = builder.ins().load(
cranelift_types::I64,
MemFlagsData::trusted(),
self.args_ptr,
second_word_offset,
);
let is_present =
builder
.ins()
.icmp_imm_u(cranelift::prelude::IntCC::NotEqual, value, 0);
(is_present, string_len)
} else {
let is_present = builder.ins().load(
cranelift_types::I8,
MemFlagsData::trusted(),
self.args_ptr,
second_word_offset,
);
let string_len = builder.ins().iconst(cranelift_types::I64, 0);
(is_present, string_len)
};
Ok(LoweredValue {
value,
is_present,
string_len,
})
}
TypedExprAst::Coerce { target_type, expr } => {
let source_type = expr.return_type;
let lowered = self.compile_expr(expr, builder)?;
let value = lower_coercion(lowered.value, source_type, *target_type, builder)?;
Ok(LoweredValue {
value,
is_present: lowered.is_present,
string_len: lowered.string_len,
})
}
TypedExprAst::FnCall(fn_call) => fn_call.lower(expression.return_type, self, builder),
}
}
pub(crate) fn pointer_type(&self) -> Type {
self.pointer_type
}
pub(crate) fn string_arena_ptr(&mut self, builder: &mut FunctionBuilder<'_>) -> CraneliftValue {
if !self.string_arena_was_reset {
let zero = builder.ins().iconst(cranelift_types::I64, 0);
builder.ins().store(
MemFlagsData::trusted(),
zero,
self.string_arena_ptr,
StringArena::CURSOR_OFFSET,
);
self.string_arena_was_reset = true;
}
self.string_arena_ptr
}
pub(crate) fn native_functions(&self) -> &NativeFunctions {
self.native_functions
}
}
fn lower_literal(
literal: &TypedLiteral,
context: &mut LoweringContext<'_>,
builder: &mut FunctionBuilder<'_>,
) -> LoweredValue {
let value = match literal {
TypedLiteral::None => builder.ins().iconst(cranelift_types::I64, 0),
TypedLiteral::Bool(value) => builder.ins().iconst(cranelift_types::I8, i64::from(*value)),
TypedLiteral::U64(value) => builder.ins().iconst(cranelift_types::I64, *value as i64),
TypedLiteral::I64(value) => builder.ins().iconst(cranelift_types::I64, *value),
TypedLiteral::F64(value) => {
builder
.ins()
.f64const(cranelift::codegen::ir::immediates::Ieee64::with_bits(
value.to_bits(),
))
}
TypedLiteral::String(value) => {
let string_ptr = value.as_ptr() as usize;
builder
.ins()
.iconst(context.pointer_type, string_ptr as i64)
}
};
let is_present = builder.ins().iconst(
cranelift_types::I8,
i64::from(!matches!(literal, TypedLiteral::None)),
);
let string_len = match literal {
TypedLiteral::String(value) => builder
.ins()
.iconst(cranelift_types::I64, value.len() as i64),
_ => builder.ins().iconst(cranelift_types::I64, 0),
};
LoweredValue {
value,
is_present,
string_len,
}
}
fn lower_coercion(
value: CraneliftValue,
source: VarType,
target: VarType,
builder: &mut FunctionBuilder<'_>,
) -> Result<CraneliftValue, CompileError> {
let coerced = match (source, target) {
(source, target) if source == target => value,
(VarType::U64, VarType::F64) => builder.ins().fcvt_from_uint(cranelift_types::F64, value),
(VarType::I64, VarType::F64) => builder.ins().fcvt_from_sint(cranelift_types::F64, value),
// Cranelift integers do not carry signedness. These two coercions have
// the same machine representation and therefore need no instruction.
(VarType::U64, VarType::I64) | (VarType::I64, VarType::U64) => value,
_ => {
return Err(CompileError::UnsupportedCoercion {
from_type: source,
target,
});
}
};
Ok(coerced)
}
fn cranelift_type(var_type: VarType, pointer_type: Type) -> Type {
match var_type {
VarType::Bool => cranelift_types::I8,
VarType::F64 => cranelift_types::F64,
VarType::U64 | VarType::I64 | VarType::None => cranelift_types::I64,
VarType::Str => pointer_type,
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use super::*;
use crate::ast::{Function, UntypedExpr};
use crate::types::{VarType, VariableValue};
#[test]
fn test_compile_bool_variable() {
let untyped_expr = UntypedExpr::variable("flag");
let variable_types = HashMap::from([("flag", VarType::Bool)]);
let compiled_fn = compile(&untyped_expr, &variable_types).unwrap();
let mut string_arena = StringArena::new();
assert!(string_arena.allocate(1).is_some());
let input = [VariableValue::some(true)];
let output = unsafe { compiled_fn.call(&input, &mut string_arena) };
assert_eq!(unsafe { output.as_bool() }, Some(true));
assert_eq!(string_arena.used_bytes(), 1);
}
#[test]
fn test_compiled_fn_is_send_and_sync() {
fn assert_send_and_sync<T: Send + Sync>() {}
assert_send_and_sync::<CompiledFn>();
}
#[test]
fn test_compile_none_variable() {
let untyped_expr = UntypedExpr::variable("value");
let variable_types = HashMap::from([("value", VarType::U64)]);
let mut compiled_fn = compile(&untyped_expr, &variable_types).unwrap().context();
let input = [VariableValue::none()];
assert_eq!(compiled_fn.result_type(), VarType::U64);
let output = unsafe { compiled_fn.call(&input) };
assert_eq!(unsafe { output.as_u64() }, None);
}
#[test]
fn test_compile_string_literal_keeps_backing_data_alive() {
let untyped_expr = UntypedExpr::literal("hello");
let mut compiled_fn = compile(&untyped_expr, &HashMap::new()).unwrap().context();
drop(untyped_expr);
let output = unsafe { compiled_fn.call(&[]) };
assert_eq!(unsafe { output.as_str() }, Some("hello"));
}
#[test]
fn test_compile_returns_borrowed_string_variable_directly() {
let untyped_expr = UntypedExpr::variable("value");
let variable_types = HashMap::from([("value", VarType::Str)]);
let mut compiled_fn = compile(&untyped_expr, &variable_types).unwrap().context();
let value = String::from("hello from an input");
let input = [VariableValue::some(value.as_str())];
let output = unsafe { compiled_fn.call(&input) };
let output = unsafe { output.as_str() }.unwrap();
assert_eq!(output, value);
assert_eq!(output.as_ptr(), value.as_ptr());
}
#[test]
fn test_compile_none_literal_returns_absent_value() {
let untyped_expr = UntypedExpr::literal(crate::ast::Literal::None);
let mut compiled_fn = compile(&untyped_expr, &HashMap::new()).unwrap().context();
assert_eq!(compiled_fn.result_type(), VarType::None);
let output = unsafe { compiled_fn.call(&[]) };
assert_eq!(unsafe { output.as_u64() }, None);
}
#[test]
fn test_compile_to_assembly() {
let untyped_expr = UntypedExpr::variable("value");
let variable_types = HashMap::from([("value", VarType::F64)]);
let assembly = compile_to_assembly(&untyped_expr, &variable_types).unwrap();
assert!(assembly.contains("block0:"));
assert!(!assembly.trim().is_empty());
}
#[test]
fn test_compile_native_call_to_assembly() {
let untyped_expr = UntypedExpr::new_fn_call(
Function::RegexpExtract,
vec![
UntypedExpr::variable("message"),
UntypedExpr::literal("([a-z]+)"),
UntypedExpr::literal(0u64),
],
)
.unwrap();
let variable_types = HashMap::from([("message", VarType::Str)]);
let assembly = compile_to_assembly(&untyped_expr, &variable_types).unwrap();
assert!(assembly.contains("block0:"));
assert!(!assembly.trim().is_empty());
}
#[cfg(target_arch = "aarch64")]
#[test]
fn test_compile_to_assembly_does_not_sign_return_address() {
let untyped_expr = UntypedExpr::new_fn_call(
Function::RegexpExtract,
vec![
UntypedExpr::variable("value"),
UntypedExpr::literal("([a-z]+)"),
],
)
.unwrap();
let variable_types = HashMap::from([("value", VarType::Str)]);
let assembly = compile_to_assembly(&untyped_expr, &variable_types).unwrap();
assert!(!assembly.contains("pacibsp"));
assert!(!assembly.contains("retabsp"));
}
}
+43
View File
@@ -0,0 +1,43 @@
pub(crate) const STRING_ARENA_CAPACITY: usize = 262_144;
/// Fixed-capacity storage for strings constructed while evaluating an expression.
pub struct StringArena {
buffer: Box<[u8; STRING_ARENA_CAPACITY]>,
cursor: usize,
}
impl StringArena {
pub(crate) const CURSOR_OFFSET: i32 = std::mem::offset_of!(StringArena, cursor) as i32;
/// Creates an empty string arena.
pub fn new() -> Self {
let buffer = vec![0; STRING_ARENA_CAPACITY].into_boxed_slice();
let buffer = buffer
.try_into()
.unwrap_or_else(|_| unreachable!("the arena buffer has the requested capacity"));
Self { buffer, cursor: 0 }
}
/// Reserves `len` contiguous bytes without growing the backing allocation.
pub(crate) fn allocate(&mut self, len: usize) -> Option<*mut u8> {
let end = self.cursor.checked_add(len)?;
if end > STRING_ARENA_CAPACITY {
return None;
}
// SAFETY: `cursor <= end <= STRING_ARENA_CAPACITY`.
let allocation = unsafe { self.buffer.as_mut_ptr().add(self.cursor) };
self.cursor = end;
Some(allocation)
}
#[cfg(test)]
pub(crate) fn used_bytes(&self) -> usize {
self.cursor
}
}
impl Default for StringArena {
fn default() -> Self {
Self::new()
}
}
+177
View File
@@ -0,0 +1,177 @@
use std::sync::Arc;
#[cfg(test)]
use crate::ast::Literal;
use crate::functions::FnCallEnum;
use crate::types::VarType;
#[derive(Clone, PartialEq)]
pub struct TypedVariable {
/// The source-level variable name.
pub variable_name: Arc<str>,
/// The concrete type expected in this input slot.
pub r#type: VarType,
/// The position in the compiled input array.
pub variable_id: usize,
}
#[derive(Clone, PartialEq)]
pub(crate) struct TypedExpr {
pub(crate) return_type: VarType,
pub(crate) ast: TypedExprAst,
}
impl TypedExpr {
pub(crate) fn coerce(self, target_type: VarType) -> TypedExpr {
if target_type == self.return_type {
self
} else {
let TypedExpr { return_type, ast } = self;
let ast = match (ast, target_type) {
(TypedExprAst::Literal(TypedLiteral::U64(value)), VarType::I64)
if value <= i64::MAX as u64 =>
{
TypedExprAst::Literal(TypedLiteral::I64(value as i64))
}
(TypedExprAst::Literal(TypedLiteral::U64(value)), VarType::F64) => {
TypedExprAst::Literal(TypedLiteral::F64(value as f64))
}
(TypedExprAst::Literal(TypedLiteral::I64(value)), VarType::U64) if value >= 0 => {
TypedExprAst::Literal(TypedLiteral::U64(value as u64))
}
(TypedExprAst::Literal(TypedLiteral::I64(value)), VarType::F64) => {
TypedExprAst::Literal(TypedLiteral::F64(value as f64))
}
(TypedExprAst::Literal(TypedLiteral::F64(value)), VarType::U64)
if value.is_finite()
&& value.fract() == 0.0
&& value >= 0.0
&& value < u64::MAX as f64 =>
{
TypedExprAst::Literal(TypedLiteral::U64(value as u64))
}
(TypedExprAst::Literal(TypedLiteral::F64(value)), VarType::I64)
if value.is_finite()
&& value.fract() == 0.0
&& value >= i64::MIN as f64
&& value < -(i64::MIN as f64) =>
{
TypedExprAst::Literal(TypedLiteral::I64(value as i64))
}
(ast, target_type) => TypedExprAst::Coerce {
target_type,
expr: Box::new(TypedExpr { return_type, ast }),
},
};
TypedExpr {
return_type: target_type,
ast,
}
}
}
pub(crate) fn none() -> TypedExpr {
TypedExpr {
return_type: VarType::None,
ast: TypedExprAst::Literal(TypedLiteral::None),
}
}
#[cfg(test)]
pub(crate) fn literal(val: impl Into<Literal>) -> TypedExpr {
let literal: Literal = val.into();
let r#type = literal.r#type();
let literal = match literal {
Literal::None => TypedLiteral::None,
Literal::Bool(value) => TypedLiteral::Bool(value),
Literal::U64(value) => TypedLiteral::U64(value),
Literal::I64(value) => TypedLiteral::I64(value),
Literal::F64(value) => TypedLiteral::F64(value),
Literal::String(_) => panic!("typed string literals require registered backing data"),
};
TypedExprAst::Literal(literal).with_type(r#type)
}
}
#[derive(Clone, Debug, PartialEq)]
pub(crate) enum TypedLiteral {
None,
Bool(bool),
U64(u64),
I64(i64),
F64(f64),
String(Arc<str>),
}
impl TypedLiteral {
pub(crate) fn r#type(&self) -> VarType {
match self {
TypedLiteral::None => VarType::None,
TypedLiteral::Bool(_) => VarType::Bool,
TypedLiteral::U64(_) => VarType::U64,
TypedLiteral::I64(_) => VarType::I64,
TypedLiteral::F64(_) => VarType::F64,
TypedLiteral::String(_) => VarType::Str,
}
}
}
#[derive(Clone, PartialEq)]
pub(crate) enum TypedExprAst {
Literal(TypedLiteral),
Variable(TypedVariable),
Coerce {
target_type: VarType,
expr: Box<TypedExpr>,
},
FnCall(FnCallEnum),
}
impl TypedExprAst {
#[cfg(test)]
pub(crate) fn with_type(self, return_type: VarType) -> TypedExpr {
TypedExpr {
return_type,
ast: self,
}
}
pub(crate) fn variable(variable_name: impl ToString, r#type: VarType) -> TypedExprAst {
TypedExprAst::Variable(TypedVariable {
variable_name: Arc::from(variable_name.to_string()),
r#type,
variable_id: 0,
})
}
pub(crate) fn from_fn_call(fn_call: impl Into<FnCallEnum>) -> TypedExprAst {
TypedExprAst::FnCall(fn_call.into())
}
}
impl std::fmt::Debug for TypedExpr {
fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
write!(f, "({:?} : {:?})", self.ast, self.return_type)
}
}
impl std::fmt::Debug for TypedExprAst {
fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
match self {
TypedExprAst::Literal(literal) => write!(f, "{:?}", literal),
TypedExprAst::Variable(variable) => write!(f, "{:?}", variable),
TypedExprAst::Coerce { target_type, expr } => {
write!(f, "coerce({:?} as {:?})", expr, target_type)
}
TypedExprAst::FnCall(fn_call) => {
write!(f, "{fn_call:?}")
}
}
}
}
impl std::fmt::Debug for TypedVariable {
fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
write!(f, "{{{}:{:?}}}", self.variable_name, self.r#type)
}
}
+201
View File
@@ -0,0 +1,201 @@
//! Serialization for the normalized typed expression tree.
//!
//! Calls, variables, and coercions use `[type: expression]`, while literals retain the canonical
//! untyped literal syntax because numerical suffixes and literal spellings already identify their
//! types. For example:
//!
//! ```text
//! [int64: ADD 1i64 [int64: my_col]]
//! ```
use std::fmt;
use super::{TypedExpr, TypedExprAst, TypedLiteral};
use crate::types::VarType;
pub(super) fn serialize(expr: &TypedExpr) -> String {
expr.to_string()
}
impl fmt::Display for TypedExpr {
fn fmt(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
format_expr(self, formatter)
}
}
fn format_expr(expr: &TypedExpr, formatter: &mut fmt::Formatter) -> fmt::Result {
if let TypedExprAst::Literal(literal) = &expr.ast {
return format_literal(literal, formatter);
}
write!(formatter, "[{}: ", type_name(expr.return_type))?;
match &expr.ast {
TypedExprAst::Literal(_) => unreachable!(),
TypedExprAst::Variable(variable) => {
crate::ast::format_variable_name(&variable.variable_name, formatter)?;
}
TypedExprAst::Coerce { expr, .. } => write!(formatter, "COERCE {expr}")?,
TypedExprAst::FnCall(fn_call) => fn_call.serialize(formatter)?,
}
formatter.write_str("]")
}
fn format_literal(literal: &TypedLiteral, formatter: &mut fmt::Formatter) -> fmt::Result {
match literal {
TypedLiteral::None => formatter.write_str("none"),
TypedLiteral::Bool(value) => write!(formatter, "{value}"),
TypedLiteral::U64(value) => write!(formatter, "{value}u64"),
TypedLiteral::I64(value) => write!(formatter, "{value}i64"),
TypedLiteral::F64(value) => write!(formatter, "{value}f64"),
TypedLiteral::String(value) => format_string_literal(value, formatter),
}
}
pub(crate) fn format_fn_call<'a>(
name: &str,
args: impl IntoIterator<Item = &'a TypedExpr>,
formatter: &mut fmt::Formatter,
) -> fmt::Result {
formatter.write_str(name)?;
for arg in args {
write!(formatter, " {arg}")?;
}
Ok(())
}
pub(crate) fn format_string_literal(value: &str, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("\"")?;
for character in value.chars() {
match character {
'"' => formatter.write_str("\\\""),
'\\' => formatter.write_str("\\\\"),
'\n' => formatter.write_str("\\n"),
'\r' => formatter.write_str("\\r"),
'\t' => formatter.write_str("\\t"),
'\0' => formatter.write_str("\\0"),
character if character.is_control() => {
write!(formatter, "{}", character.escape_unicode())
}
character => write!(formatter, "{character}"),
}?;
}
formatter.write_str("\"")
}
fn type_name(var_type: VarType) -> &'static str {
match var_type {
VarType::Bool => "boolean",
VarType::F64 => "float64",
VarType::U64 => "uint64",
VarType::I64 => "int64",
VarType::Str => "string",
VarType::None => "none",
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use crate::types::VarType;
use crate::{ast, compile};
fn serialize(expression: &str, variable_types: &HashMap<&str, VarType>) -> String {
let expression = ast::deserialize(expression).unwrap();
compile::serialize(&expression, variable_types).unwrap()
}
#[test]
fn test_variable_names_use_untyped_quoting() {
let cases = [
("my_col", "my_col"),
("field.name", "field.name"),
("", "``"),
("two words", "`two words`"),
("true", "`true`"),
("none", "`none`"),
("1u64", "`1u64`"),
("field1", "`field1`"),
("a]b", "`a]b`"),
("café", "`café`"),
("a`b", r"`a\`b`"),
("a\\b", r"`a\\b`"),
("a\n\r\t\0\u{7}", r"`a\n\r\t\0\u{7}`"),
];
for (variable_name, expected) in cases {
let expression = ast::UntypedExpr::variable(variable_name);
let variable_types = HashMap::from([(variable_name, VarType::I64)]);
assert_eq!(expression.to_string(), expected);
assert_eq!(
compile::serialize(&expression, &variable_types).unwrap(),
format!("[int64: {expected}]")
);
}
}
#[test]
fn test_serializes_typed_add_expression() {
let variable_types = HashMap::from([("my_col", VarType::I64)]);
assert_eq!(
serialize("(ADD 1i64 my_col)", &variable_types),
"[int64: ADD 1i64 [int64: my_col]]"
);
}
#[test]
fn test_serializes_explicit_coercion() {
let variable_types = HashMap::from([("my_col", VarType::I64)]);
assert_eq!(
serialize("(ADD 1.5f64 my_col)", &variable_types),
"[float64: ADD 1.5f64 [float64: COERCE [int64: my_col]]]"
);
}
#[test]
fn test_serializes_all_variable_types_and_literals() {
let cases = [
(VarType::Bool, "[boolean: value]"),
(VarType::F64, "[float64: value]"),
(VarType::U64, "[uint64: value]"),
(VarType::I64, "[int64: value]"),
(VarType::Str, "[string: value]"),
];
for (var_type, expected) in cases {
assert_eq!(
serialize("value", &HashMap::from([("value", var_type)])),
expected
);
}
let literals = [
("none", "none"),
("true", "true"),
("1u64", "1i64"),
("18446744073709551615u64", "18446744073709551615u64"),
("-2i64", "-2i64"),
("1.5f64", "1.5f64"),
(
r#""quoted: \"hello\"\\world\n\t\0\u{7} café""#,
r#""quoted: \"hello\"\\world\n\t\0\u{7} café""#,
),
];
for (literal, expected) in literals {
assert_eq!(serialize(literal, &HashMap::new()), expected);
}
}
#[test]
fn test_serializes_normalized_compile_time_arguments() {
let variable_types = HashMap::from([("message", VarType::Str)]);
let cases = [(
r#"(REGEXP_EXTRACT message "([a-z]+)")"#,
r#"[string: REGEXP_EXTRACT [string: message] "([a-z]+)" 0u64]"#,
)];
for (expression, expected) in cases {
assert_eq!(serialize(expression, &variable_types), expected);
}
}
}
+509
View File
@@ -0,0 +1,509 @@
// Adds takes an arbitrary number of arguments and adds them.
//
// The type of the addition is rather complex.
// We consider the possible types of all arguments, make an intersection of those, and
// pick the first available type with the order of priority i64, u64, f64.
//
// For instance (ADD mycol 1f64) where mycol is i64 will coerce
// 1f64 to 1i64 at compile time (because we have detected that the conversion was lossless), and the
// operation will run over integer.
//
// On the other hand, (ADD mycol 1.2f64) where mycol is i64 will coerce
// mycol to float dynamically (because 1.2f64 cannot be converted to u64 with loss).
//
// If any of the values of the arguments is NaN, the function return NaN.
use std::collections::HashMap;
use cranelift::prelude::{FunctionBuilder, InstBuilder, types};
use crate::ast::{Function, InferredTypeSet, TypeError, UntypedExpr};
use crate::compile::{
CompileError, CompileFnBuilder, LoweredValue, LoweringContext, TypedExpr, TypedExprAst,
};
use crate::functions::{FnCall, FnCallEnum};
use crate::types::VarType;
#[derive(Clone, Debug, PartialEq)]
pub(crate) struct AddFnCall {
pub(crate) args: Box<[TypedExpr]>,
}
impl FnCall for AddFnCall {
const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Any;
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::NUMERICAL).is_none() {
return Err(TypeError::WrongFunctionReturnType {
function: Function::Add,
expected: target_type,
got: InferredTypeSet::NUMERICAL,
});
}
let mut return_types = InferredTypeSet::NUMERICAL;
for arg in args {
let arg_types =
crate::ast::infer_types_aux(arg, InferredTypeSet::NUMERICAL, inferred_types)?;
return_types = return_types.intersect(arg_types);
}
return_types = with_float_fallback(return_types);
let constrained_return_types = return_types.intersect(target_type);
if constrained_return_types.is_none() {
return Err(TypeError::WrongFunctionReturnType {
function: Function::Add,
expected: target_type,
got: return_types,
});
}
Ok(constrained_return_types)
}
fn call_with_types(
args: &[UntypedExpr],
target_type_set: InferredTypeSet,
context: &mut CompileFnBuilder<'_, '_>,
) -> Result<TypedExpr, CompileError> {
Self::ARG_COUNT.validate(args)?;
let mut return_types = InferredTypeSet::NUMERICAL.intersect(target_type_set);
for arg in args {
let arg_types = crate::ast::infer_type_with_variable_types(
arg,
InferredTypeSet::NUMERICAL,
context.variable_types(),
)?;
return_types = return_types.intersect(arg_types);
}
let return_type = select_return_type(with_float_fallback(return_types));
let typed_args: Vec<TypedExpr> = args
.iter()
.map(|arg| context.apply_types(arg, InferredTypeSet::singleton(return_type)))
.collect::<Result<_, _>>()?;
if typed_args
.iter()
.any(|typed_arg| !is_numerical(typed_arg.return_type))
{
return Ok(TypedExpr::none());
}
Ok(TypedExpr {
return_type,
ast: TypedExprAst::from_fn_call(AddFnCall {
args: typed_args.into_boxed_slice(),
}),
})
}
fn args_mut(&mut self) -> &mut [TypedExpr] {
&mut self.args
}
fn serialize(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
crate::compile::format_fn_call("ADD", self.args.iter(), formatter)
}
fn emit_cranelift_ir(
&self,
return_type: VarType,
context: &mut LoweringContext<'_>,
builder: &mut FunctionBuilder<'_>,
) -> Result<LoweredValue, CompileError> {
let mut sum = match return_type {
VarType::U64 | VarType::I64 => builder.ins().iconst(types::I64, 0),
VarType::F64 => builder.ins().f64const(0.0),
_ => {
return Err(CompileError::UnsupportedFunctionType {
function: Function::Add,
return_type,
});
}
};
let mut is_present = builder.ins().iconst(types::I8, 1);
for arg in &self.args {
let lowered = context.compile_expr(arg, builder)?;
sum = match return_type {
VarType::U64 | VarType::I64 => builder.ins().iadd(sum, lowered.value),
VarType::F64 => builder.ins().fadd(sum, lowered.value),
_ => unreachable!("the return type was checked above"),
};
is_present = builder.ins().band(is_present, lowered.is_present);
}
Ok(LoweredValue {
value: sum,
is_present,
string_len: builder.ins().iconst(types::I64, 0),
})
}
}
pub(super) fn with_float_fallback(inferred_types: InferredTypeSet) -> InferredTypeSet {
if inferred_types.is_none() {
InferredTypeSet::F64
} else {
inferred_types
}
}
pub(super) fn select_return_type(inferred_types: InferredTypeSet) -> VarType {
if inferred_types.i64 {
VarType::I64
} else if inferred_types.u64 {
VarType::U64
} else {
debug_assert!(inferred_types.f64);
VarType::F64
}
}
pub(super) fn is_numerical(var_type: VarType) -> bool {
matches!(var_type, VarType::I64 | VarType::U64 | VarType::F64)
}
impl From<AddFnCall> for FnCallEnum {
fn from(call: AddFnCall) -> Self {
FnCallEnum::Add(call)
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use super::*;
use crate::ast::{Literal, infer_types};
use crate::compile::{TypedExprAst, compile};
use crate::types::VariableValue;
#[test]
fn test_infer_types_rejects_string_argument() {
let expr = UntypedExpr::new_fn_call(
Function::Add,
vec![UntypedExpr::literal(1.0), UntypedExpr::literal("hello")],
)
.unwrap();
let error = infer_types(&expr).unwrap_err();
assert!(matches!(
error,
TypeError::InvalidLiteralType {
literal: Literal::String(_),
expected: InferredTypeSet::NUMERICAL,
}
));
}
#[test]
fn test_infer_types_constrains_variables_to_numerical() {
let expr = UntypedExpr::new_fn_call(
Function::Add,
vec![UntypedExpr::variable("a"), UntypedExpr::variable("b")],
)
.unwrap();
let inferred_types = infer_types(&expr).unwrap();
assert_eq!(inferred_types.get("a"), Some(&InferredTypeSet::NUMERICAL));
assert_eq!(inferred_types.get("b"), Some(&InferredTypeSet::NUMERICAL));
}
#[test]
fn test_call_with_types_preserves_u64() {
let variable_types = HashMap::from([("present", VarType::U64)]);
let typed_expr = crate::typed_expr_from_str("(ADD present 1u64)", &variable_types);
assert_eq!(typed_expr.return_type, VarType::U64);
assert_eq!(
typed_expr,
TypedExpr {
return_type: VarType::U64,
ast: TypedExprAst::from_fn_call(AddFnCall {
args: vec![
TypedExprAst::variable("present", VarType::U64).with_type(VarType::U64),
TypedExpr::literal(1u64),
]
.into_boxed_slice()
}),
}
);
}
#[test]
fn test_call_with_types_rematerializes_compatible_literal_as_u64() {
let variable_types = HashMap::from([("present", VarType::U64)]);
let typed_expr = crate::typed_expr_from_str("(ADD present 1i64)", &variable_types);
assert_eq!(typed_expr.return_type, VarType::U64);
assert_eq!(
typed_expr,
TypedExpr {
return_type: VarType::U64,
ast: TypedExprAst::from_fn_call(AddFnCall {
args: vec![
TypedExprAst::variable("present", VarType::U64).with_type(VarType::U64),
TypedExpr::literal(1u64),
]
.into_boxed_slice()
}),
}
);
}
#[test]
fn test_call_with_types_rematerializes_integral_f64_literal_as_u64() {
let variable_types = HashMap::from([("present", VarType::U64)]);
let typed_expr = crate::typed_expr_from_str("(ADD 1.0f64 present)", &variable_types);
assert_eq!(typed_expr.return_type, VarType::U64);
let TypedExprAst::FnCall(FnCallEnum::Add(call)) = typed_expr.ast else {
panic!("expected an ADD call");
};
assert_eq!(call.args[0], TypedExpr::literal(1u64));
}
#[test]
fn test_call_with_types_rematerializes_compatible_literal_as_i64() {
let variable_types = HashMap::from([("present", VarType::I64)]);
let typed_expr = crate::typed_expr_from_str("(ADD present 1u64)", &variable_types);
assert_eq!(typed_expr.return_type, VarType::I64);
let TypedExprAst::FnCall(FnCallEnum::Add(call)) = typed_expr.ast else {
panic!("expected an ADD call");
};
assert_eq!(call.args[1], TypedExpr::literal(1i64));
}
#[test]
fn test_call_with_types_prefers_i64_for_compatible_literals() {
let variable_types = HashMap::new();
let typed_expr = crate::typed_expr_from_str("(ADD 1u64 2.0f64)", &variable_types);
assert_eq!(typed_expr.return_type, VarType::I64);
let TypedExprAst::FnCall(FnCallEnum::Add(call)) = typed_expr.ast else {
panic!("expected an ADD call");
};
assert_eq!(call.args[0], TypedExpr::literal(1i64));
assert_eq!(call.args[1], TypedExpr::literal(2i64));
}
#[test]
fn test_call_with_types_uses_u64_when_i64_is_not_possible() {
let variable_types = HashMap::new();
let expression = crate::ast::deserialize("(ADD 9223372036854775808u64)").unwrap();
let typed_expr =
crate::typed_expr_from_str("(ADD 9223372036854775808u64)", &variable_types);
assert_eq!(typed_expr.return_type, VarType::U64);
assert_eq!(
typed_expr,
TypedExpr {
return_type: VarType::U64,
ast: TypedExprAst::from_fn_call(AddFnCall {
args: vec![TypedExpr::literal(9223372036854775808u64)].into_boxed_slice(),
}),
}
);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let output = unsafe { compiled.call(&[]) };
assert_eq!(unsafe { output.as_u64() }, Some(9223372036854775808u64));
}
#[test]
fn test_call_with_types_coerces_mixed_numbers_to_f64() {
let variable_types = HashMap::from([("present", VarType::U64)]);
let typed_expr = crate::typed_expr_from_str("(ADD present 1.2f64)", &variable_types);
assert_eq!(typed_expr.return_type, VarType::F64);
assert!(matches!(typed_expr.ast, TypedExprAst::FnCall(_)));
}
#[test]
fn test_call_with_types_propagates_missing_variable() {
let variable_types = HashMap::from([("present", VarType::U64)]);
let typed_expr =
crate::typed_expr_from_str("(ADD present (ADD 1u64 missing))", &variable_types);
assert_eq!(typed_expr, TypedExpr::none());
}
#[test]
fn test_compile_signed_add() {
let expression = UntypedExpr::new_fn_call(
Function::Add,
vec![
UntypedExpr::literal(-4i64),
UntypedExpr::variable("myfield"),
],
)
.unwrap();
let variable_types = HashMap::from([("myfield", VarType::I64)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let input = [VariableValue::some(-8i64)];
let output = unsafe { compiled.call(&input) };
assert_eq!(unsafe { output.as_i64() }, Some(-12));
}
#[test]
fn test_compile_adds_i64_literal_to_u64_variable_without_float_coercion() {
let variable_types = HashMap::from([("myfield", VarType::U64)]);
let argument_orders = [
vec![UntypedExpr::variable("myfield"), UntypedExpr::literal(1i64)],
vec![UntypedExpr::literal(1i64), UntypedExpr::variable("myfield")],
];
for args in argument_orders {
let expression = UntypedExpr::new_fn_call(Function::Add, args).unwrap();
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let input = [VariableValue::some(41u64)];
let output = unsafe { compiled.call(&input) };
assert_eq!(unsafe { output.as_u64() }, Some(42));
}
}
#[test]
fn test_compile_coerces_compatible_nested_add_to_u64() {
let nested_literals = UntypedExpr::new_fn_call(
Function::Add,
vec![UntypedExpr::literal(1i64), UntypedExpr::literal(2u64)],
)
.unwrap();
let expression = UntypedExpr::new_fn_call(
Function::Add,
vec![UntypedExpr::variable("myfield"), nested_literals],
)
.unwrap();
let variable_types = HashMap::from([("myfield", VarType::U64)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let input = [VariableValue::some(39u64)];
let output = unsafe { compiled.call(&input) };
assert_eq!(unsafe { output.as_u64() }, Some(42));
}
#[test]
fn test_compile_coerces_integers_to_float() {
let expression = UntypedExpr::new_fn_call(
Function::Add,
vec![
UntypedExpr::variable("myfield"),
UntypedExpr::literal(-2i64),
UntypedExpr::literal(0.5f64),
],
)
.unwrap();
let variable_types = HashMap::from([("myfield", VarType::U64)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let input = [VariableValue::some(10u64)];
let output = unsafe { compiled.call(&input) };
assert_eq!(unsafe { output.as_f64() }, Some(8.5));
}
#[test]
fn test_compile_loads_multiple_variable_slots() {
let expression = UntypedExpr::new_fn_call(
Function::Add,
vec![UntypedExpr::variable("x"), UntypedExpr::variable("y")],
)
.unwrap();
let variable_types = HashMap::from([("x", VarType::U64), ("y", VarType::F64)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let input = [VariableValue::some(10u64), VariableValue::some(0.5f64)];
let output = unsafe { compiled.call(&input) };
assert_eq!(unsafe { output.as_f64() }, Some(10.5));
}
#[test]
fn test_compile_add_propagates_none_input() {
let expression = crate::ast::deserialize(r#"(ADD x 0.5f64)"#).unwrap();
let variable_types = HashMap::from([("x", VarType::U64)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let input = [VariableValue::none()];
let output = unsafe { compiled.call(&input) };
assert_eq!(unsafe { output.as_f64() }, None);
}
#[test]
fn test_compile_u64_to_float_coercion_is_unsigned() {
let expression = UntypedExpr::new_fn_call(
Function::Add,
vec![UntypedExpr::variable("x"), UntypedExpr::literal(0.5f64)],
)
.unwrap();
let variable_types = HashMap::from([("x", VarType::U64)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let input = [VariableValue::some(u64::MAX)];
let output = unsafe { compiled.call(&input) };
assert_eq!(unsafe { output.as_f64() }, Some(u64::MAX as f64 + 0.5));
}
#[test]
fn test_compile_reuses_repeated_variable_slot() {
let expression = UntypedExpr::new_fn_call(
Function::Add,
vec![
UntypedExpr::variable("x"),
UntypedExpr::variable("x"),
UntypedExpr::literal(1u64),
],
)
.unwrap();
let variable_types = HashMap::from([("x", VarType::U64)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let input = [VariableValue::some(4i64)];
assert_eq!(compiled.inputs().len(), 1);
assert_eq!(compiled.inputs()[0].variable_name.as_ref(), "x");
assert_eq!(compiled.inputs()[0].r#type, VarType::U64);
assert_eq!(compiled.inputs()[0].variable_id, 0);
assert_eq!(compiled.result_type(), VarType::U64);
let output = unsafe { compiled.call(&input) };
assert_eq!(unsafe { output.as_u64() }, Some(9));
}
#[test]
fn test_compile_can_coerce_variable_when_necessary() {
let expression = UntypedExpr::new_fn_call(
Function::Add,
vec![UntypedExpr::variable("x"), UntypedExpr::literal(1.2f64)],
)
.unwrap();
let variable_types = HashMap::from([("x", VarType::U64)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let input = [VariableValue::some(4u64)];
assert_eq!(compiled.inputs().len(), 1);
assert_eq!(compiled.inputs()[0].variable_name.as_ref(), "x");
assert_eq!(compiled.inputs()[0].r#type, VarType::U64);
assert_eq!(compiled.inputs()[0].variable_id, 0);
assert_eq!(compiled.result_type(), VarType::F64);
let output = unsafe { compiled.call(&input) };
assert_eq!(unsafe { output.as_f64() }, Some(5.2f64));
}
#[test]
fn test_compile_empty_add_uses_zero_identity() {
let variable_types = HashMap::new();
let typed_expr = crate::typed_expr_from_str("(ADD)", &variable_types);
assert_eq!(typed_expr.return_type, VarType::I64);
let expression = UntypedExpr::new_fn_call(Function::Add, Vec::new()).unwrap();
let mut compiled = compile(&expression, &HashMap::new()).unwrap().context();
let output = unsafe { compiled.call(&[]) };
assert_eq!(unsafe { output.as_i64() }, Some(0));
}
#[test]
fn test_no_variable_works() {
let args = vec![UntypedExpr::literal(1.2f64), UntypedExpr::literal(1u64)];
let variable_types = HashMap::new();
let typed_expr = crate::typed_expr_from_str("(ADD 1.2f64 1u64)", &variable_types);
assert_eq!(typed_expr.return_type, VarType::F64);
let expression = UntypedExpr::new_fn_call(Function::Add, args).unwrap();
let mut compiled = compile(&expression, &HashMap::new()).unwrap().context();
let output = unsafe { compiled.call(&[]) };
assert_eq!(unsafe { output.as_f64() }, Some(2.2f64));
}
}
+196
View File
@@ -0,0 +1,196 @@
//! `IS_NULL` observes absence without propagating it.
//!
//! It accepts exactly one expression of any supported type and always returns a present boolean:
//! `true` when its argument is absent and `false` when it is present. Payload values such as
//! `false`, zero, `NaN`, and the empty string are present and therefore return `false`.
use std::collections::HashMap;
use cranelift::prelude::{FunctionBuilder, InstBuilder, types};
use crate::ast::{Function, InferredTypeSet, TypeError, UntypedExpr};
use crate::compile::{
CompileError, CompileFnBuilder, LoweredValue, LoweringContext, TypedExpr, TypedExprAst,
};
use crate::functions::{FnCall, FnCallEnum};
use crate::types::VarType;
#[derive(Clone, Debug, PartialEq)]
pub(crate) struct IsNullFnCall {
pub(crate) args: Box<[TypedExpr]>,
}
impl FnCall for IsNullFnCall {
const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(1);
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::BOOLEAN).is_none() {
return Err(TypeError::WrongFunctionReturnType {
function: Function::IsNull,
expected: target_type,
got: InferredTypeSet::BOOLEAN,
});
}
if args.len() != 1 {
return Err(TypeError::InvalidNumberOfArguments {
function: Function::IsNull,
expected: 1,
got: args.len(),
});
}
crate::ast::infer_types_aux(&args[0], InferredTypeSet::ALL, inferred_types)?;
Ok(InferredTypeSet::BOOLEAN)
}
fn call_with_types(
args: &[UntypedExpr],
target_type_set: InferredTypeSet,
context: &mut CompileFnBuilder<'_, '_>,
) -> Result<TypedExpr, CompileError> {
Self::ARG_COUNT.validate(args)?;
debug_assert!(target_type_set.contains(VarType::Bool));
let arg = context.apply_types(&args[0], InferredTypeSet::ALL)?;
Ok(TypedExpr {
return_type: VarType::Bool,
ast: TypedExprAst::from_fn_call(IsNullFnCall {
args: vec![arg].into_boxed_slice(),
}),
})
}
fn args_mut(&mut self) -> &mut [TypedExpr] {
&mut self.args
}
fn serialize(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
crate::compile::format_fn_call("IS_NULL", self.args.iter(), formatter)
}
fn emit_cranelift_ir(
&self,
return_type: VarType,
context: &mut LoweringContext<'_>,
builder: &mut FunctionBuilder<'_>,
) -> Result<LoweredValue, CompileError> {
debug_assert_eq!(return_type, VarType::Bool);
let arg = context.compile_expr(&self.args[0], builder)?;
let value = builder.ins().bxor_imm_u(arg.is_present, 1);
Ok(LoweredValue {
value,
is_present: builder.ins().iconst(types::I8, 1),
string_len: builder.ins().iconst(types::I64, 0),
})
}
}
impl From<IsNullFnCall> for FnCallEnum {
fn from(call: IsNullFnCall) -> Self {
FnCallEnum::IsNull(call)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ast::{InvalidFnCall, deserialize, infer_types};
use crate::compile::compile;
use crate::functions::ArgumentCount;
use crate::types::VariableValue;
fn eval(expression: &str) -> Option<bool> {
let expression = deserialize(expression).unwrap();
let mut compiled = compile(&expression, &HashMap::new()).unwrap().context();
// SAFETY: These expressions have no runtime inputs and return booleans.
unsafe { compiled.call(&[]).as_bool() }
}
#[test]
fn test_infer_types_accepts_any_single_argument() {
let expression = deserialize("(IS_NULL value)").unwrap();
let inferred_types = infer_types(&expression).unwrap();
assert_eq!(inferred_types.get("value"), Some(&InferredTypeSet::ALL));
}
#[test]
fn test_rejects_wrong_arity() {
for args in [
vec![],
vec![
UntypedExpr::variable("value"),
UntypedExpr::variable("other"),
],
] {
let invalid_fn_call: InvalidFnCall =
UntypedExpr::new_fn_call(Function::IsNull, args).unwrap_err();
assert!(matches!(
invalid_fn_call,
InvalidFnCall::InvalidNumberOfArguments {
expected: ArgumentCount::Exactly(1),
..
}
));
}
}
#[test]
fn test_absent_values_are_true() {
assert_eq!(eval("(IS_NULL none)"), Some(true));
assert_eq!(eval("(IS_NULL missing)"), Some(true));
assert_eq!(
eval(r#"(IS_NULL (REGEXP_EXTRACT "b" "(a+)" 1u64))"#),
Some(true)
);
}
#[test]
fn test_present_edge_values_are_false() {
for expression in [
"(IS_NULL false)",
"(IS_NULL 0i64)",
"(IS_NULL -0f64)",
r#"(IS_NULL "")"#,
r#"(IS_NULL (REGEXP_EXTRACT "b" "(a*)b" 1u64))"#,
] {
assert_eq!(eval(expression), Some(false), "expression: {expression}");
}
}
#[test]
fn test_non_finite_values_are_present() {
// The textual parser rejects non-finite literals, but programmatically
// constructed expressions can still contain them. They are not null.
for value in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
let expression =
UntypedExpr::new_fn_call(Function::IsNull, vec![UntypedExpr::literal(value)])
.unwrap();
let mut compiled = compile(&expression, &HashMap::new()).unwrap().context();
// SAFETY: The expression has no runtime inputs and returns a boolean.
assert_eq!(unsafe { compiled.call(&[]).as_bool() }, Some(false));
}
}
#[test]
fn test_runtime_absent_variable_is_true_and_present_variable_is_false() {
let expression = deserialize("(IS_NULL value)").unwrap();
let variable_types = HashMap::from([("value", VarType::I64)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
// SAFETY: The input and output types match the compiled signature.
assert_eq!(
unsafe { compiled.call(&[VariableValue::none()]).as_bool() },
Some(true)
);
// SAFETY: The input and output types match the compiled signature.
assert_eq!(
unsafe { compiled.call(&[VariableValue::some(0i64)]).as_bool() },
Some(false)
);
}
}
+380
View File
@@ -0,0 +1,380 @@
mod add;
mod is_null;
mod native_function;
mod regexp_extract;
use std::collections::HashMap;
use cranelift::frontend::FunctionBuilder;
pub(crate) use self::add::AddFnCall;
pub(crate) use self::is_null::IsNullFnCall;
pub(crate) use self::native_function::{
NativeFunctions, declare_native_functions, register_jit_symbols,
};
pub(crate) use self::regexp_extract::RegexpExtractFnCall;
use crate::ast::{InferredTypeSet, Literal, TypeError, UntypedExpr};
use crate::compile::{CompileError, CompileFnBuilder, LoweredValue, LoweringContext, TypedExpr};
use crate::types::VarType;
/// A function supported by the first expression-language milestone.
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
pub enum Function {
/// Adds zero or more numerical expressions.
Add,
/// Tests whether an expression produced an absent value.
IsNull,
/// Extracts a capture group from a string using a constant regular expression.
RegexpExtract,
}
impl Function {
pub(crate) fn call(self, args: Vec<UntypedExpr>) -> Result<UntypedExpr, InvalidFnCall> {
match self {
Function::Add => <AddFnCall as FnCall>::validate_args(&args)?,
Function::IsNull => <IsNullFnCall as FnCall>::validate_args(&args)?,
Function::RegexpExtract => <RegexpExtractFnCall as FnCall>::validate_args(&args)?,
}
Ok(UntypedExpr::FnCall {
function: self,
args,
})
}
pub(crate) fn call_with_types(
self,
args: &[UntypedExpr],
target_type_set: InferredTypeSet,
context: &mut CompileFnBuilder<'_, '_>,
) -> Result<TypedExpr, CompileError> {
match self {
Function::Add => <AddFnCall as FnCall>::call_with_types(args, target_type_set, context),
Function::IsNull => {
<IsNullFnCall as FnCall>::call_with_types(args, target_type_set, context)
}
Function::RegexpExtract => {
<RegexpExtractFnCall as FnCall>::call_with_types(args, target_type_set, context)
}
}
}
pub(crate) fn infer_types<'a>(
self,
args: &'a [UntypedExpr],
target_type: InferredTypeSet,
inferred_types: &mut HashMap<&'a str, InferredTypeSet>,
) -> Result<InferredTypeSet, TypeError> {
match self {
Function::Add => <AddFnCall as FnCall>::infer_types(args, target_type, inferred_types),
Function::IsNull => {
<IsNullFnCall as FnCall>::infer_types(args, target_type, inferred_types)
}
Function::RegexpExtract => {
<RegexpExtractFnCall as FnCall>::infer_types(args, target_type, inferred_types)
}
}
}
}
#[derive(Clone, Debug, PartialEq)]
pub(crate) enum FnCallEnum {
Add(AddFnCall),
IsNull(IsNullFnCall),
RegexpExtract(RegexpExtractFnCall),
}
impl FnCallEnum {
pub(crate) fn serialize(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
match self {
FnCallEnum::Add(call) => call.serialize(formatter),
FnCallEnum::IsNull(call) => call.serialize(formatter),
FnCallEnum::RegexpExtract(call) => call.serialize(formatter),
}
}
pub(crate) fn args_mut(&mut self) -> &mut [TypedExpr] {
match self {
FnCallEnum::Add(call) => call.args_mut(),
FnCallEnum::IsNull(call) => call.args_mut(),
FnCallEnum::RegexpExtract(call) => call.args_mut(),
}
}
/// Produce CraneLift IR for the given function call.
pub(crate) fn lower(
&self,
return_type: VarType,
context: &mut LoweringContext<'_>,
builder: &mut FunctionBuilder<'_>,
) -> Result<LoweredValue, CompileError> {
match self {
FnCallEnum::Add(call) => call.emit_cranelift_ir(return_type, context, builder),
FnCallEnum::IsNull(call) => call.emit_cranelift_ir(return_type, context, builder),
FnCallEnum::RegexpExtract(call) => {
call.emit_cranelift_ir(return_type, context, builder)
}
}
}
}
/// Error representing an invalid function call.
#[derive(Debug, Eq, PartialEq, thiserror::Error)]
pub enum InvalidFnCall {
#[error("invalid number of arguments: expected {expected}, got {provided}")]
InvalidNumberOfArguments {
expected: ArgumentCount,
provided: usize,
},
#[error("argument {argument} must be a {expected:?} literal")]
ExpectedLiteral { argument: usize, expected: VarType },
#[error("invalid value for argument {argument}: expected {expected}")]
InvalidLiteralValue {
argument: usize,
expected: &'static str,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ArgumentCount {
Any,
Exactly(usize),
AtLeast(usize),
Between { min: usize, max: usize },
}
impl std::fmt::Display for ArgumentCount {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match *self {
ArgumentCount::Any => formatter.write_str("any number of arguments"),
ArgumentCount::Exactly(1) => formatter.write_str("exactly 1 argument"),
ArgumentCount::Exactly(count) => {
write!(formatter, "exactly {count} arguments")
}
ArgumentCount::AtLeast(1) => formatter.write_str("at least 1 argument"),
ArgumentCount::AtLeast(count) => {
write!(formatter, "at least {count} arguments")
}
ArgumentCount::Between { min: 1, max: 1 } => formatter.write_str("exactly 1 argument"),
ArgumentCount::Between { min, max } if min == max => {
write!(formatter, "exactly {min} arguments")
}
ArgumentCount::Between { min, max } => {
write!(formatter, "between {min} and {max} arguments")
}
}
}
}
impl ArgumentCount {
fn validate(self, args: &[UntypedExpr]) -> Result<(), InvalidFnCall> {
let provided = args.len();
let is_valid = match self {
ArgumentCount::Any => true,
ArgumentCount::Exactly(expected) => provided == expected,
ArgumentCount::AtLeast(expected) => provided >= expected,
ArgumentCount::Between { min, max } => (min..=max).contains(&provided),
};
if is_valid {
Ok(())
} else {
Err(InvalidFnCall::InvalidNumberOfArguments {
expected: self,
provided,
})
}
}
}
pub(crate) fn validate_literal(
args: &[UntypedExpr],
index: usize,
expected: VarType,
is_valid: impl FnOnce(&Literal) -> bool,
) -> Result<(), InvalidFnCall> {
let Some(UntypedExpr::Literal(literal)) = args.get(index) else {
return Err(InvalidFnCall::ExpectedLiteral {
argument: index + 1,
expected,
});
};
if is_valid(literal) {
Ok(())
} else {
Err(InvalidFnCall::ExpectedLiteral {
argument: index + 1,
expected,
})
}
}
/// Implements the type-inference, typed-AST, and lowering phases of a function call.
///
/// The static methods operate on an [`UntypedExpr`] call before a concrete call node exists.
/// Once [`FnCall::call_with_types`] has produced that node, [`FnCall::args_mut`] and
/// [`FnCall::emit_cranelift_ir`] operate on its typed representation.
pub(crate) trait FnCall: std::fmt::Debug + Into<FnCallEnum> {
const ARG_COUNT: ArgumentCount;
/// Constrains the call and its arguments to the types accepted by its parent expression.
///
/// Implementations validate their signature, recursively infer every argument, update
/// `inferred_types` with the accepted types for variables, and return the possible result
/// types that remain after intersecting with `target_type`.
///
/// As we recursively visit the argument, we encounter variables.
/// This method is meant to restrict the set of accept types associated to each
/// variable name by mutating their inferred type set, found in inferred_types map.
fn infer_types<'a>(
args: &'a [UntypedExpr],
target_type: InferredTypeSet,
inferred_types: &mut HashMap<&'a str, InferredTypeSet>,
) -> Result<InferredTypeSet, TypeError>
where
Self: Sized;
fn validate_args(args: &[UntypedExpr]) -> Result<(), InvalidFnCall> {
Self::ARG_COUNT.validate(args)?;
Ok(())
}
/// Builds the typed call after concrete variable types have been supplied.
///
/// `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 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(
args: &[UntypedExpr],
target_type_set: InferredTypeSet,
context: &mut CompileFnBuilder<'_, '_>,
) -> Result<TypedExpr, CompileError>
where
Self: Sized;
/// Returns the typed child expressions that participate in recursive AST passes.
///
/// This is only used, to assign and deduplicate variable input slots. Compile-time
/// configuration stored directly on a call does not need to be returned.
///
/// Today this is only used as a cheap visitor to allocate variable ids.
fn args_mut(&mut self) -> &mut [TypedExpr];
/// Serializes the function name and its normalized typed arguments.
fn serialize(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result;
/// Emits Cranelift IR for an already typed call and returns its result SSA value.
///
/// `return_type` is the concrete type selected during typed-AST construction. Implementations
/// lower child expressions through `context` and append their own instructions to `builder`.
fn emit_cranelift_ir(
&self,
return_type: VarType,
context: &mut LoweringContext<'_>,
builder: &mut FunctionBuilder<'_>,
) -> Result<LoweredValue, CompileError>;
}
#[cfg(test)]
mod tests {
use super::*;
use crate::compile::compile;
fn call_error(function: Function, args: Vec<UntypedExpr>) -> InvalidFnCall {
match UntypedExpr::new_fn_call(function, args) {
Ok(_) => panic!("expected the function call to be rejected"),
Err(error) => error,
}
}
#[test]
fn test_argument_count_display() {
assert_eq!(ArgumentCount::Any.to_string(), "any number of arguments");
assert_eq!(ArgumentCount::Exactly(1).to_string(), "exactly 1 argument");
assert_eq!(ArgumentCount::Exactly(2).to_string(), "exactly 2 arguments");
assert_eq!(ArgumentCount::AtLeast(1).to_string(), "at least 1 argument");
assert_eq!(
ArgumentCount::AtLeast(2).to_string(),
"at least 2 arguments"
);
assert_eq!(
ArgumentCount::Between { min: 2, max: 3 }.to_string(),
"between 2 and 3 arguments"
);
}
#[test]
fn test_argument_count_validation() {
assert_eq!(
call_error(Function::IsNull, Vec::new()),
InvalidFnCall::InvalidNumberOfArguments {
expected: ArgumentCount::Exactly(1),
provided: 0,
}
);
assert_eq!(
ArgumentCount::AtLeast(1).validate(&[]).unwrap_err(),
InvalidFnCall::InvalidNumberOfArguments {
expected: ArgumentCount::AtLeast(1),
provided: 0,
}
);
assert_eq!(
call_error(Function::RegexpExtract, Vec::new()),
InvalidFnCall::InvalidNumberOfArguments {
expected: ArgumentCount::Between { min: 2, max: 3 },
provided: 0,
}
);
assert!(UntypedExpr::new_fn_call(Function::Add, Vec::new()).is_ok());
}
#[test]
fn test_literal_argument_validation() {
assert_eq!(
call_error(
Function::RegexpExtract,
vec![
UntypedExpr::variable("input"),
UntypedExpr::variable("pattern")
],
),
InvalidFnCall::ExpectedLiteral {
argument: 2,
expected: VarType::Str,
}
);
assert_eq!(
call_error(
Function::RegexpExtract,
vec![
UntypedExpr::variable("input"),
UntypedExpr::literal("pattern"),
UntypedExpr::literal(1i64),
],
),
InvalidFnCall::ExpectedLiteral {
argument: 3,
expected: VarType::U64,
}
);
}
#[test]
fn test_typed_construction_validates_unchecked_ast() {
// Bypass construction-time validation to exercise the compiler's checks.
let expression = UntypedExpr::FnCall {
function: Function::IsNull,
args: Vec::new(),
};
let error = compile(&expression, &HashMap::new()).err().unwrap();
assert!(matches!(
error,
CompileError::InvalidArguments(InvalidFnCall::InvalidNumberOfArguments {
expected: ArgumentCount::Exactly(1),
provided: 0,
})
));
}
}
+32
View File
@@ -0,0 +1,32 @@
use cranelift::codegen::ir::{FuncRef, Function as CraneliftFunction, Type};
use cranelift_jit::{JITBuilder, JITModule};
use super::regexp_extract;
use crate::compile::CompileError;
/// References to native functions imported into the current Cranelift function.
pub(crate) struct NativeFunctions {
regexp_extract: FuncRef,
}
impl NativeFunctions {
pub(crate) fn regexp_extract(&self) -> FuncRef {
self.regexp_extract
}
}
/// Registers the process symbols that native calls may reference from generated code.
pub(crate) fn register_jit_symbols(jit_builder: &mut JITBuilder) {
regexp_extract::register_jit_symbol(jit_builder);
}
/// Declares every native function imported by the expression being compiled.
pub(crate) fn declare_native_functions(
module: &mut JITModule,
function: &mut CraneliftFunction,
pointer_type: Type,
) -> Result<NativeFunctions, CompileError> {
Ok(NativeFunctions {
regexp_extract: regexp_extract::declare_native_function(module, function, pointer_type)?,
})
}
+434
View File
@@ -0,0 +1,434 @@
// RegexpExtract extracts a regular-expression match from a string.
//
// It takes two or three arguments:
// - string: the input string
// - const string: a regular-expression pattern literal. This one CANNOT be the result of another
// expression
// - optional const u64: capture index literal. It defaults to 0 during conversion to the typed
// expression. Capture index 0 returns the full match, while indexes 1 and above return the
// corresponding explicit capture group.
//
// It returns None when the input is None, the pattern does not match, or the requested capture
// 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;
use cranelift::prelude::{AbiParam, InstBuilder};
use cranelift_jit::{JITBuilder, JITModule};
use cranelift_module::{Linkage, Module};
use regex::Regex;
use crate::ast::{Function, InferredTypeSet, Literal, TypeError, UntypedExpr};
use crate::compile::{
CompileError, CompileFnBuilder, LoweredValue, LoweringContext, TypedExpr, TypedExprAst,
};
use crate::functions::{FnCall, FnCallEnum, InvalidFnCall};
use crate::types::VarType;
const SYMBOL: &str = "jitexpr_regexp_extract";
#[derive(Clone, Debug)]
pub(crate) struct RegexpExtractFnCall {
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 {
const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Between { min: 2, max: 3 };
fn validate_args(args: &[UntypedExpr]) -> Result<(), InvalidFnCall> {
Self::ARG_COUNT.validate(args)?;
super::validate_literal(args, 1, VarType::Str, |literal| {
matches!(literal, Literal::String(_))
})?;
if args.len() == 3 {
super::validate_literal(args, 2, VarType::U64, |literal| {
matches!(literal, Literal::U64(_))
})?;
}
Ok(())
}
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::RegexpExtract,
expected: target_type,
got: InferredTypeSet::STRING,
});
}
if !(2..=3).contains(&args.len()) {
return Err(TypeError::InvalidNumberOfArguments {
function: Function::RegexpExtract,
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(capture_index) = args.get(2) {
crate::ast::infer_types_aux(capture_index, InferredTypeSet::NUMERICAL, inferred_types)?;
}
Ok(InferredTypeSet::STRING)
}
fn call_with_types(
args: &[UntypedExpr],
target_type_set: InferredTypeSet,
context: &mut CompileFnBuilder<'_, '_>,
) -> Result<TypedExpr, CompileError> {
Self::ARG_COUNT.validate(args)?;
let haystack = context.apply_types(&args[0], target_type_set)?;
if haystack.return_type != VarType::Str {
return Ok(TypedExpr::none());
}
let UntypedExpr::Literal(Literal::String(pattern)) = &args[1] else {
return Err(InvalidFnCall::ExpectedLiteral {
argument: 2,
expected: VarType::Str,
}
.into());
};
let regex = Arc::new(
Regex::new(pattern).map_err(|source| CompileError::InvalidRegex {
pattern: pattern.to_string(),
source,
})?,
);
let capture_index = match args.get(2) {
None => 0,
Some(UntypedExpr::Literal(Literal::U64(capture_index))) => *capture_index,
Some(_) => {
return Err(InvalidFnCall::ExpectedLiteral {
argument: 3,
expected: VarType::U64,
}
.into());
}
};
Ok(TypedExpr {
return_type: VarType::Str,
ast: TypedExprAst::from_fn_call(RegexpExtractFnCall {
regex,
haystack: Box::new(haystack),
capture_index,
}),
})
}
fn args_mut(&mut self) -> &mut [TypedExpr] {
std::slice::from_mut(&mut self.haystack)
}
fn serialize(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
write!(formatter, "REGEXP_EXTRACT {} ", self.haystack)?;
crate::compile::format_string_literal(self.regex.as_str(), formatter)?;
write!(formatter, " {}u64", self.capture_index)
}
/// Produce CraneLift IR for the given function call.
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 haystack = context.compile_expr(&self.haystack, builder)?;
let null = builder.ins().iconst(context.pointer_type(), 0);
let haystack_ptr = builder
.ins()
.select(haystack.is_present, haystack.value, null);
let regex_ptr = builder
.ins()
.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(),
&[regex_ptr, haystack_ptr, haystack.string_len, capture_index],
);
let value = builder.inst_results(call)[0];
let string_len = builder.inst_results(call)[1];
let is_present = builder
.ins()
.icmp_imm_u(cranelift::prelude::IntCC::NotEqual, value, 0);
Ok(LoweredValue {
value,
is_present,
string_len,
})
}
}
pub(super) fn register_jit_symbol(jit_builder: &mut JITBuilder) {
jit_builder.symbol(SYMBOL, regexp_extract 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(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));
signature.returns.push(AbiParam::new(types::I64));
let function_id = module.declare_function(SYMBOL, Linkage::Import, &signature)?;
Ok(module.declare_func_in_func(function_id, function))
}
/// Raw two-word string result returned to generated code.
#[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(),
}
}
}
/// Runtime implementation called by generated code for `RegexpExtract`.
///
/// 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(
regex: *const Regex,
haystack_ptr: *const u8,
haystack_len: usize,
capture_index: u64,
) -> RawStr {
if haystack_ptr.is_null() {
return RawStr::none();
}
let Ok(capture_index) = usize::try_from(capture_index) else {
return RawStr::none();
};
// 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 {
std::str::from_utf8_unchecked(std::slice::from_raw_parts(haystack_ptr, haystack_len))
};
let Some(regex_match) = regex
.captures(haystack)
.and_then(|captures| captures.get(capture_index))
else {
return RawStr::none();
};
RawStr::some(regex_match.as_str())
}
impl From<RegexpExtractFnCall> for FnCallEnum {
fn from(call: RegexpExtractFnCall) -> Self {
FnCallEnum::RegexpExtract(call)
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use super::*;
use crate::ast::{self, infer_types};
use crate::compile::compile;
use crate::types::VariableValue;
#[test]
fn test_infer_types_constrains_haystack_to_string() {
let expression = ast::deserialize(r#"(REGEXP_EXTRACT message "([a-z]+)")"#).unwrap();
let inferred_types = infer_types(&expression).unwrap();
assert_eq!(
inferred_types.get("message"),
Some(&InferredTypeSet::STRING)
);
}
#[test]
fn test_infer_types_accepts_optional_capture_index() {
for expression in [
r#"(REGEXP_EXTRACT message "([a-z]+)")"#,
r#"(REGEXP_EXTRACT message "([a-z]+)" 1u64)"#,
] {
let expression = ast::deserialize(expression).unwrap();
assert!(infer_types(&expression).is_ok());
}
for args in [
vec![UntypedExpr::variable("message")],
vec![
UntypedExpr::variable("message"),
UntypedExpr::literal("([a-z]+)"),
UntypedExpr::literal(0u64),
UntypedExpr::literal(1u64),
],
] {
// Bypass call validation to exercise type inference on an invalid AST.
let invalid_fn_call: InvalidFnCall =
UntypedExpr::new_fn_call(Function::RegexpExtract, args).unwrap_err();
assert!(matches!(
invalid_fn_call,
InvalidFnCall::InvalidNumberOfArguments { .. }
));
}
}
#[test]
fn test_compile_returns_borrowed_capture() {
let expression =
ast::deserialize(r#"(REGEXP_EXTRACT message "([a-z]+)-(\\d+)" 1u64)"#).unwrap();
let variable_types = HashMap::from([("message", VarType::Str)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let haystack = "prefix user-123 suffix";
let input = [VariableValue::some(haystack)];
let output = unsafe { compiled.call(&input) };
let extracted = unsafe { output.as_str() }.unwrap();
assert_eq!(extracted, "user");
assert_eq!(extracted.as_ptr(), haystack[7..].as_ptr());
}
#[test]
fn test_compile_selects_capture_by_index() {
let expression =
ast::deserialize(r#"(REGEXP_EXTRACT message "([a-z]+)-(\\d+)" 2u64)"#).unwrap();
let variable_types = HashMap::from([("message", VarType::Str)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let input = [VariableValue::some("user-123")];
let output = unsafe { compiled.call(&input) };
assert_eq!(unsafe { output.as_str() }, Some("123"));
}
#[test]
fn test_compile_returns_none_without_capture() {
let expression =
ast::deserialize(r#"(REGEXP_EXTRACT message "([a-z]+)-(\\d+)" 0u64)"#).unwrap();
let variable_types = HashMap::from([("message", VarType::Str)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let input = [VariableValue::some("no digits here")];
let output = unsafe { compiled.call(&input) };
assert_eq!(unsafe { output.as_str() }, None);
}
#[test]
fn test_compile_propagates_none_haystack() {
let expression = ast::deserialize(r#"(REGEXP_EXTRACT message "([a-z]+)" 0u64)"#).unwrap();
let variable_types = HashMap::from([("message", VarType::Str)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let input = [VariableValue::none()];
let output = unsafe { compiled.call(&input) };
assert_eq!(unsafe { output.as_str() }, None);
}
#[test]
fn test_compile_propagates_compile_time_none_haystack() {
let expression = ast::deserialize(r#"(REGEXP_EXTRACT missing "([a-z]+)" 0u64)"#).unwrap();
let mut compiled = compile(&expression, &HashMap::new()).unwrap().context();
assert_eq!(compiled.result_type(), VarType::None);
let output = unsafe { compiled.call(&[]) };
assert_eq!(unsafe { output.as_str() }, None);
}
#[test]
fn test_compile_distinguishes_empty_capture_from_none() {
let expression = ast::deserialize(r#"(REGEXP_EXTRACT "b" "(a*)b" 1u64)"#).unwrap();
let mut compiled = compile(&expression, &HashMap::new()).unwrap().context();
let output = unsafe { compiled.call(&[]) };
assert_eq!(unsafe { output.as_str() }, Some(""));
}
#[test]
fn test_compile_omitted_group_defaults_to_full_match() {
let expression = ast::deserialize(r#"(REGEXP_EXTRACT message "[a-z]+-\\d+")"#).unwrap();
let variable_types = HashMap::from([("message", VarType::Str)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let input = [VariableValue::some("prefix user-123 suffix")];
let output = unsafe { compiled.call(&input) };
assert_eq!(unsafe { output.as_str() }, Some("user-123"));
}
#[test]
fn test_compile_group_zero_returns_full_match_without_capture_groups() {
let expression =
ast::deserialize(r#"(REGEXP_EXTRACT message "[a-z]+-\\d+" 0u64)"#).unwrap();
let variable_types = HashMap::from([("message", VarType::Str)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let input = [VariableValue::some("prefix user-123 suffix")];
let output = unsafe { compiled.call(&input) };
assert_eq!(unsafe { output.as_str() }, Some("user-123"));
}
#[test]
fn test_compile_nested_calls_use_their_own_regexes() {
let expression = ast::deserialize(
r#"(REGEXP_EXTRACT
(REGEXP_EXTRACT message "([a-z]+-\\d+)" 1u64)
"([a-z]+)"
1u64)"#,
)
.unwrap();
let variable_types = HashMap::from([("message", VarType::Str)]);
let mut compiled = compile(&expression, &variable_types).unwrap().context();
let input = [VariableValue::some("id=user-123!")];
let output = unsafe { compiled.call(&input) };
assert_eq!(unsafe { output.as_str() }, Some("user"));
}
#[test]
fn test_compile_rejects_invalid_pattern() {
let expression = ast::deserialize(r#"(REGEXP_EXTRACT "anything" "(" 0u64)"#).unwrap();
let error = compile(&expression, &HashMap::new()).err().unwrap();
assert!(matches!(
error,
CompileError::InvalidRegex { pattern, .. } if pattern == "("
));
}
}
+34
View File
@@ -0,0 +1,34 @@
pub mod ast;
pub mod compile;
pub mod types;
mod functions;
#[cfg(test)]
pub(crate) fn typed_expr_from_str(
untyped_expr: &str,
variable_types: &std::collections::HashMap<&str, types::VarType>,
) -> compile::TypedExpr {
let untyped_expr = ast::deserialize(untyped_expr).unwrap();
let mut context = compile::CompileFnBuilder::new(variable_types);
context
.apply_types(&untyped_expr, ast::InferredTypeSet::ALL)
.unwrap()
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use super::*;
use crate::types::VarType;
#[test]
fn test_typed_expr_from_str() {
let variable_types = HashMap::from([("value", VarType::U64)]);
let typed_expr = typed_expr_from_str("(ADD value 1i64)", &variable_types);
assert_eq!(typed_expr.return_type, VarType::U64);
}
}
+345
View File
@@ -0,0 +1,345 @@
//! Source types and nullable runtime value representations.
/// A value type supported by compiled expressions.
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq, Ord, PartialOrd)]
pub enum VarType {
Bool,
F64,
U64,
I64,
Str,
None,
}
/// The payload of a primitive runtime value.
///
/// This union is deliberately untagged. The corresponding
/// [`crate::compile::TypedVariable`] identifies the active payload field.
#[repr(C)]
#[derive(Clone, Copy)]
pub union VariablePrimitive {
pub boolean: bool,
pub float: f64,
pub int_u64: u64,
pub int_i64: i64,
}
impl From<bool> for VariablePrimitive {
fn from(value: bool) -> Self {
VariablePrimitive { boolean: value }
}
}
impl From<f64> for VariablePrimitive {
fn from(value: f64) -> Self {
VariablePrimitive { float: value }
}
}
impl From<u64> for VariablePrimitive {
fn from(value: u64) -> Self {
VariablePrimitive { int_u64: value }
}
}
impl From<i64> for VariablePrimitive {
fn from(value: i64) -> Self {
VariablePrimitive { int_i64: value }
}
}
impl Default for VariablePrimitive {
fn default() -> Self {
VariablePrimitive { int_u64: 0 }
}
}
/// A nullable primitive value.
///
/// `value` is meaningful only when `is_present` is true.
#[repr(C)]
#[derive(Clone, Copy, Default)]
pub struct VariablePrimitiveOpt {
pub value: VariablePrimitive,
pub is_present: bool,
}
impl VariablePrimitiveOpt {
/// Wraps a present primitive value.
pub fn some(value: impl Into<VariablePrimitive>) -> Self {
Self {
value: value.into(),
is_present: true,
}
}
/// Creates an absent primitive value.
pub fn none() -> Self {
Self::default()
}
}
impl<T: Into<VariablePrimitive>> From<T> for VariablePrimitiveOpt {
fn from(value: T) -> Self {
Self::some(value)
}
}
/// A nullable runtime argument or result slot.
///
/// Primitive values use the [`VariablePrimitiveOpt`] arm. Strings use the
/// nullable `string` arm: a null data pointer represents `None`, while a
/// non-null data pointer and its byte length represent a borrowed `str`.
/// Both arms occupy two machine words on the supported 64-bit targets.
///
/// string relies on the null ptr optimization to make that happen.
#[repr(C)]
#[derive(Clone, Copy)]
pub union VariableValue<'a> {
pub primitive: VariablePrimitiveOpt,
pub string: Option<&'a str>,
}
const _: () = {
assert!(std::mem::size_of::<VariablePrimitive>() == 8);
assert!(std::mem::offset_of!(VariablePrimitiveOpt, value) == 0);
assert!(std::mem::offset_of!(VariablePrimitiveOpt, is_present) == 8);
assert!(std::mem::size_of::<VariablePrimitiveOpt>() == 16);
assert!(std::mem::size_of::<Option<&str>>() == 16);
assert!(std::mem::size_of::<VariableValue>() == 16);
assert!(std::mem::align_of::<VariableValue>() == 8);
};
/// Pins down the *internal* layout of the `string` arm, which `#[repr(C)]` on
/// the union does not specify.
const _: () = {
// The `None` niche is a null data pointer in word 0.
let none_words: [usize; 2] = unsafe { std::mem::transmute::<Option<&str>, [usize; 2]>(None) };
assert!(none_words[0] == 0);
const SAMPLE: &str = "abcde";
let some_parts: (*const u8, usize) =
unsafe { std::mem::transmute::<Option<&str>, (*const u8, usize)>(Some(SAMPLE)) };
assert!(!some_parts.0.is_null());
assert!(some_parts.1 == SAMPLE.len());
};
impl<'a> VariableValue<'a> {
/// Wraps a present runtime value.
#[inline(always)]
pub fn some(value: impl Into<Self>) -> Self {
value.into()
}
/// Creates an absent runtime value for either arm.
#[inline(always)]
pub fn none() -> Self {
// SAFETY: All-zeroes is both an absent VariablePrimitiveOpt and the
// null niche used by Option<&str>.
unsafe { std::mem::zeroed() }
}
/// Returns the boolean payload, or `None` when this value is absent.
///
/// # Safety
///
/// This value must contain a primitive boolean or be absent.
#[inline(always)]
pub unsafe fn as_bool(self) -> Option<bool> {
// SAFETY: Guaranteed by the caller.
let primitive = unsafe { self.primitive };
if primitive.is_present {
// SAFETY: The caller guarantees that the active payload is `boolean`.
Some(unsafe { primitive.value.boolean })
} else {
None
}
}
/// Returns the `f64` payload, or `None` when this value is absent.
///
/// # Safety
///
/// This value must contain a primitive `f64` or be absent.
#[inline(always)]
pub unsafe fn as_f64(self) -> Option<f64> {
// SAFETY: Guaranteed by the caller.
let primitive = unsafe { self.primitive };
if primitive.is_present {
// SAFETY: The caller guarantees that the active payload is `float`.
Some(unsafe { primitive.value.float })
} else {
None
}
}
/// Returns the `u64` payload, or `None` when this value is absent.
///
/// # Safety
///
/// This value must contain a primitive `u64` or be absent.
#[inline(always)]
pub unsafe fn as_u64(self) -> Option<u64> {
// SAFETY: Guaranteed by the caller.
let primitive = unsafe { self.primitive };
if primitive.is_present {
// SAFETY: The caller guarantees that the active payload is `int_u64`.
Some(unsafe { primitive.value.int_u64 })
} else {
None
}
}
/// Returns the `i64` payload, or `None` when this value is absent.
///
/// # Safety
///
/// This value must contain a primitive `i64` or be absent.
#[inline(always)]
pub unsafe fn as_i64(self) -> Option<i64> {
// SAFETY: Guaranteed by the caller.
let primitive = unsafe { self.primitive };
if primitive.is_present {
// SAFETY: The caller guarantees that the active payload is `int_i64`.
Some(unsafe { primitive.value.int_i64 })
} else {
None
}
}
/// Returns the borrowed string payload, or `None` when it is absent.
///
/// # Safety
///
/// This value must contain the `string` arm or be the all-zero absent
/// representation returned by [`VariableValue::none`].
#[inline(always)]
pub unsafe fn as_str(self) -> Option<&'a str> {
// SAFETY: Guaranteed by the caller.
unsafe { self.string }
}
}
impl Default for VariableValue<'_> {
fn default() -> Self {
Self::none()
}
}
impl From<bool> for VariableValue<'_> {
fn from(value: bool) -> Self {
Self {
primitive: VariablePrimitiveOpt::some(value),
}
}
}
impl From<f64> for VariableValue<'_> {
fn from(value: f64) -> Self {
Self {
primitive: VariablePrimitiveOpt::some(value),
}
}
}
impl From<u64> for VariableValue<'_> {
fn from(value: u64) -> Self {
Self {
primitive: VariablePrimitiveOpt::some(value),
}
}
}
impl From<i64> for VariableValue<'_> {
fn from(value: i64) -> Self {
Self {
primitive: VariablePrimitiveOpt::some(value),
}
}
}
impl<'a> From<&'a str> for VariableValue<'a> {
fn from(value: &'a str) -> Self {
Self {
string: Some(value),
}
}
}
impl<'a> From<Option<&'a str>> for VariableValue<'a> {
fn from(value: Option<&'a str>) -> Self {
match value {
Some(value) => Self::from(value),
None => Self::none(),
}
}
}
impl<'a> From<VariablePrimitive> for VariableValue<'a> {
fn from(value: VariablePrimitive) -> Self {
Self {
primitive: VariablePrimitiveOpt::some(value),
}
}
}
impl<'a> From<VariablePrimitiveOpt> for VariableValue<'a> {
fn from(value: VariablePrimitiveOpt) -> Self {
Self { primitive: value }
}
}
#[cfg(test)]
mod tests {
use crate::types::{VariablePrimitive, VariablePrimitiveOpt, VariableValue};
#[test]
fn test_runtime_value_layouts() {
assert_eq!(std::mem::size_of::<VariablePrimitive>(), 8);
assert_eq!(std::mem::offset_of!(VariablePrimitiveOpt, value), 0);
assert_eq!(std::mem::offset_of!(VariablePrimitiveOpt, is_present), 8);
assert_eq!(std::mem::size_of::<VariablePrimitiveOpt>(), 16);
assert_eq!(std::mem::size_of::<Option<&str>>(), 16);
assert_eq!(std::mem::size_of::<VariableValue>(), 16);
assert_eq!(std::mem::align_of::<VariableValue>(), 8);
let text = "hello";
let words: [usize; 2] = unsafe { std::mem::transmute(VariableValue::some(text)) };
assert_eq!(words, [text.as_ptr() as usize, text.len()]);
let none_words: [usize; 2] = unsafe { std::mem::transmute(VariableValue::none()) };
assert_eq!(none_words, [0, 0]);
}
#[test]
fn test_variable_value_accessors() {
assert_eq!(unsafe { VariableValue::some(true).as_bool() }, Some(true));
assert_eq!(unsafe { VariableValue::some(1.5f64).as_f64() }, Some(1.5));
assert_eq!(unsafe { VariableValue::some(7u64).as_u64() }, Some(7));
assert_eq!(unsafe { VariableValue::some(-3i64).as_i64() }, Some(-3));
assert_eq!(
unsafe { VariableValue::some(VariablePrimitive { int_u64: 11 }).as_u64() },
Some(11)
);
assert_eq!(
unsafe { VariableValue::some("hello").as_str() },
Some("hello")
);
let none = VariableValue::none();
assert_eq!(unsafe { none.as_bool() }, None);
assert_eq!(unsafe { none.as_f64() }, None);
assert_eq!(unsafe { none.as_u64() }, None);
assert_eq!(unsafe { none.as_i64() }, None);
assert_eq!(unsafe { none.as_str() }, None);
assert_eq!(unsafe { VariableValue::from(None::<&str>).as_str() }, None);
}
#[test]
fn test_empty_string_is_distinct_from_none() {
let empty = VariableValue::some("");
let none = VariableValue::none();
assert_eq!(unsafe { empty.as_str() }, Some(""));
assert_eq!(unsafe { none.as_str() }, None);
}
}