use akar_common::types::Value;
use akar_common::vector::DataChunk;
use hashbrown::HashMap;
use std::sync::Arc;
#[derive(Clone)]
#[allow(clippy::type_complexity)]
pub enum ScalarFunction {
Arithmetic {
op: ArithmeticOp,
},
Comparison {
op: ComparisonOp,
},
String {
op: StringOp,
},
Cast {
target_type: CastTarget,
},
Date {
op: DateOp,
},
List {
op: ListOp,
},
Map {
op: MapOp,
},
Struct {
op: StructOp,
},
Boolean {
op: BooleanOp,
},
Utility {
op: UtilityOp,
},
Schema {
op: SchemaOp,
},
Array {
op: ArrayOp,
},
Path {
op: PathOp,
},
Uuid,
CustomScalar {
name: String,
execute: Arc<dyn Fn(&[Value]) -> Result<Value, String> + Send + Sync>,
},
SequenceOp {
is_nextval: bool,
},
Hash {
op: HashOp,
},
Interval {
op: IntervalOp,
},
Blob {
op: BlobOp,
},
Union {
op: UnionOp,
},
}
impl std::fmt::Debug for ScalarFunction {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Arithmetic { op } => f.debug_struct("Arithmetic").field("op", op).finish(),
Self::Comparison { op } => f.debug_struct("Comparison").field("op", op).finish(),
Self::String { op } => f.debug_struct("String").field("op", op).finish(),
Self::Cast { target_type } => f.debug_struct("Cast").field("target_type", target_type).finish(),
Self::Date { op } => f.debug_struct("Date").field("op", op).finish(),
Self::List { op } => f.debug_struct("List").field("op", op).finish(),
Self::Map { op } => f.debug_struct("Map").field("op", op).finish(),
Self::Struct { op } => f.debug_struct("Struct").field("op", op).finish(),
Self::Boolean { op } => f.debug_struct("Boolean").field("op", op).finish(),
Self::Utility { op } => f.debug_struct("Utility").field("op", op).finish(),
Self::Schema { op } => f.debug_struct("Schema").field("op", op).finish(),
Self::Array { op } => f.debug_struct("Array").field("op", op).finish(),
Self::Path { op } => f.debug_struct("Path").field("op", op).finish(),
Self::Uuid => f.debug_struct("Uuid").finish(),
Self::CustomScalar { name, .. } => f.debug_struct("CustomScalar").field("name", name).finish(),
Self::SequenceOp { is_nextval } => f.debug_struct("SequenceOp").field("is_nextval", is_nextval).finish(),
Self::Hash { op } => f.debug_struct("Hash").field("op", op).finish(),
Self::Interval { op } => f.debug_struct("Interval").field("op", op).finish(),
Self::Blob { op } => f.debug_struct("Blob").field("op", op).finish(),
Self::Union { op } => f.debug_struct("Union").field("op", op).finish(),
}
}
}
#[derive(Debug, Clone, Copy)]
pub enum ArithmeticOp {
Add,
Sub,
Mul,
Div,
Mod,
Abs,
Ceil,
Floor,
Round,
Negate,
Power,
Sqrt,
Log,
Exp,
Sin,
Cos,
Tan,
Asin,
Acos,
Atan,
Atan2,
Sinh,
Cosh,
Tanh,
Degrees,
Radians,
Sign,
Pi,
Rand,
Cbrt,
Cot,
Log2,
Even,
Gcd,
Lcm,
Factorial,
Gamma,
Lgamma,
SetSeed,
BitwiseAnd,
BitwiseOr,
BitwiseXor,
BitShiftLeft,
BitShiftRight,
}
#[derive(Debug, Clone, Copy)]
pub enum ComparisonOp {
Eq,
NotEq,
Lt,
Lte,
Gt,
Gte,
IsNull,
IsNotNull,
}
#[derive(Debug, Clone, Copy)]
pub enum StringOp {
Concat,
Contains,
StartsWith,
EndsWith,
ToUpper,
ToLower,
Trim,
LTrim,
RTrim,
Length,
Reverse,
Repeat,
Replace,
Substring,
RegexMatches,
RegexReplace,
Split,
Head,
Tail,
Left,
Right,
Lpad,
Rpad,
InitCap,
Soundex,
ConcatWs,
SplitPart,
ArrayExtract,
RegexpFullMatch,
RegexpExtract,
RegexpExtractAll,
RegexpSplitToArray,
Levenshtein,
Like,
}
#[derive(Debug, Clone, Copy)]
pub enum CastTarget {
String,
Int64,
Int32,
Double,
Float,
Bool,
Date,
Timestamp,
Interval,
}
#[derive(Debug, Clone, Copy)]
pub enum DateOp {
DatePart,
DateTrunc,
DateDiff,
DateAdd,
CurrentDate,
CurrentTimestamp,
Year,
Month,
Day,
Hour,
Minute,
Second,
DayName,
MonthName,
LastDay,
MakeDate,
Century,
EpochMs,
ToTimestamp,
ToEpochMs,
}
#[derive(Debug, Clone, Copy)]
pub enum ListOp {
Creation,
Extract,
Concat,
Len,
Sort,
Reverse,
Contains,
Append,
Prepend,
Slice,
Range,
Distinct,
Unique,
Sum,
Product,
AnyValue,
ToString,
Position,
HasAll,
ReverseSort,
Any,
All,
None,
Single,
Count,
Min,
Max,
HasAny,
Transform,
Filter,
Reduce,
}
#[derive(Debug, Clone, Copy)]
pub enum MapOp {
Creation,
Extract,
MapFromEntries,
Keys,
Values,
Contains,
}
#[derive(Debug, Clone, Copy)]
pub enum StructOp {
Creation,
Extract,
}
#[derive(Debug, Clone, Copy)]
pub enum BooleanOp {
And,
Or,
Xor,
Not,
}
#[derive(Debug, Clone, Copy)]
pub enum UtilityOp {
Coalesce,
IfNull,
TypeOf,
NullIf,
Size,
Error,
PgIsReady,
Greatest,
Least,
ConstantOrNull,
}
#[derive(Debug, Clone, Copy)]
pub enum SchemaOp {
Offset,
Id,
StartNode,
EndNode,
Label,
Cost,
RowId,
}
#[derive(Debug, Clone, Copy)]
pub enum PathOp {
Nodes,
Rels,
Length,
Properties,
IsTrail,
IsAcyclic,
}
#[derive(Debug, Clone, Copy)]
pub enum ArrayOp {
CosineSimilarity,
Distance,
InnerProduct,
DotProduct,
CrossProduct,
SquaredDistance,
Intersect,
}
#[derive(Debug, Clone, Copy)]
pub enum HashOp {
Md5,
Sha256,
Hash,
}
#[derive(Debug, Clone, Copy)]
pub enum IntervalOp {
ToYears,
ToMonths,
ToDays,
ToHours,
ToMinutes,
ToSeconds,
ToMilliseconds,
ToMicroseconds,
}
#[derive(Debug, Clone, Copy)]
pub enum BlobOp {
Encode,
Decode,
OctetLength,
ToBase64,
FromBase64,
BlobFromBytes,
}
#[derive(Debug, Clone, Copy)]
pub enum UnionOp {
UnionValue,
UnionTag,
UnionExtract,
}
#[derive(Debug, Clone)]
pub enum AggregateFunction {
Count,
Sum,
Avg,
Min,
Max,
Collect,
CountStar,
StdDev,
Variance,
StringAgg {
delimiter: String,
},
PercentileDisc {
percentile: f64,
},
PercentileCont {
percentile: f64,
},
CountIf,
}
#[derive(Clone)]
pub enum TableFunction {
ScanCsv {
path: String,
},
ScanParquet {
path: String,
},
ScanJson {
path: String,
},
ListTables,
ShowColumns {
table_name: String,
},
CurrentSetting {
key: String,
},
Custom {
name: String,
},
#[allow(clippy::type_complexity)]
CustomTable {
name: String,
execute: Arc<dyn Fn(&[Value], &mut DataChunk) -> Result<(), String> + Send + Sync>,
},
#[allow(clippy::type_complexity)]
CustomTableWithGraph {
name: String,
execute: Arc<
dyn Fn(&[Value], Option<&dyn crate::graph::GraphDataSource>, &mut DataChunk) -> Result<(), String>
+ Send
+ Sync,
>,
},
}
impl std::fmt::Debug for TableFunction {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::ScanCsv { path } => f.debug_struct("ScanCsv").field("path", path).finish(),
Self::ScanParquet { path } => f.debug_struct("ScanParquet").field("path", path).finish(),
Self::ScanJson { path } => f.debug_struct("ScanJson").field("path", path).finish(),
Self::ListTables => write!(f, "ListTables"),
Self::ShowColumns { table_name } => f.debug_struct("ShowColumns").field("table_name", table_name).finish(),
Self::CurrentSetting { key } => f.debug_struct("CurrentSetting").field("key", key).finish(),
Self::Custom { name } => f.debug_struct("Custom").field("name", name).finish(),
Self::CustomTable { name, .. } => f.debug_struct("CustomTable").field("name", name).finish(),
Self::CustomTableWithGraph { name, .. } => {
f.debug_struct("CustomTableWithGraph").field("name", name).finish()
}
}
}
}
#[derive(Debug, Clone)]
pub enum ResolvedFunction {
Scalar(ScalarFunction),
Aggregate(AggregateFunction),
Table(TableFunction),
}
#[derive(Default)]
pub struct FunctionRegistry {
scalar_functions: HashMap<String, ScalarFunction>,
aggregate_functions: HashMap<String, AggregateFunction>,
table_functions: HashMap<String, TableFunction>,
}
impl FunctionRegistry {
pub fn new() -> Self {
let mut reg = Self::default();
reg.register_builtins();
reg
}
fn register_builtins(&mut self) {
self.register_scalar("+", ScalarFunction::Arithmetic { op: ArithmeticOp::Add });
self.register_scalar("-", ScalarFunction::Arithmetic { op: ArithmeticOp::Sub });
self.register_scalar("*", ScalarFunction::Arithmetic { op: ArithmeticOp::Mul });
self.register_scalar("/", ScalarFunction::Arithmetic { op: ArithmeticOp::Div });
self.register_scalar("%", ScalarFunction::Arithmetic { op: ArithmeticOp::Mod });
self.register_scalar("abs", ScalarFunction::Arithmetic { op: ArithmeticOp::Abs });
self.register_scalar("ceil", ScalarFunction::Arithmetic { op: ArithmeticOp::Ceil });
self.register_scalar("ceiling", ScalarFunction::Arithmetic { op: ArithmeticOp::Ceil });
self.register_scalar(
"floor",
ScalarFunction::Arithmetic {
op: ArithmeticOp::Floor,
},
);
self.register_scalar(
"round",
ScalarFunction::Arithmetic {
op: ArithmeticOp::Round,
},
);
self.register_scalar(
"^",
ScalarFunction::Arithmetic {
op: ArithmeticOp::Power,
},
);
self.register_scalar("sqrt", ScalarFunction::Arithmetic { op: ArithmeticOp::Sqrt });
self.register_scalar(
"pow",
ScalarFunction::Arithmetic {
op: ArithmeticOp::Power,
},
);
self.register_scalar("log10", ScalarFunction::Arithmetic { op: ArithmeticOp::Log });
self.register_scalar("cbrt", ScalarFunction::Arithmetic { op: ArithmeticOp::Cbrt });
self.register_scalar("cot", ScalarFunction::Arithmetic { op: ArithmeticOp::Cot });
self.register_scalar("log", ScalarFunction::Arithmetic { op: ArithmeticOp::Log });
self.register_scalar("ln", ScalarFunction::Arithmetic { op: ArithmeticOp::Log });
self.register_scalar("log2", ScalarFunction::Arithmetic { op: ArithmeticOp::Log2 });
self.register_scalar("even", ScalarFunction::Arithmetic { op: ArithmeticOp::Even });
self.register_scalar(
"factorial",
ScalarFunction::Arithmetic {
op: ArithmeticOp::Factorial,
},
);
self.register_scalar(
"gamma",
ScalarFunction::Arithmetic {
op: ArithmeticOp::Gamma,
},
);
self.register_scalar(
"lgamma",
ScalarFunction::Arithmetic {
op: ArithmeticOp::Lgamma,
},
);
self.register_scalar(
"set_seed",
ScalarFunction::Arithmetic {
op: ArithmeticOp::SetSeed,
},
);
self.register_scalar("exp", ScalarFunction::Arithmetic { op: ArithmeticOp::Exp });
self.register_scalar("sin", ScalarFunction::Arithmetic { op: ArithmeticOp::Sin });
self.register_scalar("cos", ScalarFunction::Arithmetic { op: ArithmeticOp::Cos });
self.register_scalar("tan", ScalarFunction::Arithmetic { op: ArithmeticOp::Tan });
self.register_scalar("asin", ScalarFunction::Arithmetic { op: ArithmeticOp::Asin });
self.register_scalar("acos", ScalarFunction::Arithmetic { op: ArithmeticOp::Acos });
self.register_scalar("atan", ScalarFunction::Arithmetic { op: ArithmeticOp::Atan });
self.register_scalar(
"atan2",
ScalarFunction::Arithmetic {
op: ArithmeticOp::Atan2,
},
);
self.register_scalar(
"degrees",
ScalarFunction::Arithmetic {
op: ArithmeticOp::Degrees,
},
);
self.register_scalar(
"radians",
ScalarFunction::Arithmetic {
op: ArithmeticOp::Radians,
},
);
self.register_scalar("sign", ScalarFunction::Arithmetic { op: ArithmeticOp::Sign });
self.register_scalar("pi", ScalarFunction::Arithmetic { op: ArithmeticOp::Pi });
self.register_scalar("rand", ScalarFunction::Arithmetic { op: ArithmeticOp::Rand });
self.register_scalar("sinh", ScalarFunction::Arithmetic { op: ArithmeticOp::Sinh });
self.register_scalar("cosh", ScalarFunction::Arithmetic { op: ArithmeticOp::Cosh });
self.register_scalar("tanh", ScalarFunction::Arithmetic { op: ArithmeticOp::Tanh });
self.register_scalar("gcd", ScalarFunction::Arithmetic { op: ArithmeticOp::Gcd });
self.register_scalar("lcm", ScalarFunction::Arithmetic { op: ArithmeticOp::Lcm });
self.register_scalar(
"bitwise_and",
ScalarFunction::Arithmetic {
op: ArithmeticOp::BitwiseAnd,
},
);
self.register_scalar(
"&",
ScalarFunction::Arithmetic {
op: ArithmeticOp::BitwiseAnd,
},
);
self.register_scalar(
"bitwise_or",
ScalarFunction::Arithmetic {
op: ArithmeticOp::BitwiseOr,
},
);
self.register_scalar(
"|",
ScalarFunction::Arithmetic {
op: ArithmeticOp::BitwiseOr,
},
);
self.register_scalar(
"bitwise_xor",
ScalarFunction::Arithmetic {
op: ArithmeticOp::BitwiseXor,
},
);
self.register_scalar(
"#",
ScalarFunction::Arithmetic {
op: ArithmeticOp::BitwiseXor,
},
);
self.register_scalar(
"bit_shift_left",
ScalarFunction::Arithmetic {
op: ArithmeticOp::BitShiftLeft,
},
);
self.register_scalar(
"<<",
ScalarFunction::Arithmetic {
op: ArithmeticOp::BitShiftLeft,
},
);
self.register_scalar(
"bit_shift_right",
ScalarFunction::Arithmetic {
op: ArithmeticOp::BitShiftRight,
},
);
self.register_scalar(
">>",
ScalarFunction::Arithmetic {
op: ArithmeticOp::BitShiftRight,
},
);
self.register_scalar("=", ScalarFunction::Comparison { op: ComparisonOp::Eq });
self.register_scalar(
"<>",
ScalarFunction::Comparison {
op: ComparisonOp::NotEq,
},
);
self.register_scalar("<", ScalarFunction::Comparison { op: ComparisonOp::Lt });
self.register_scalar("<=", ScalarFunction::Comparison { op: ComparisonOp::Lte });
self.register_scalar(">", ScalarFunction::Comparison { op: ComparisonOp::Gt });
self.register_scalar(">=", ScalarFunction::Comparison { op: ComparisonOp::Gte });
self.register_scalar(
"IS NULL",
ScalarFunction::Comparison {
op: ComparisonOp::IsNull,
},
);
self.register_scalar(
"IS NOT NULL",
ScalarFunction::Comparison {
op: ComparisonOp::IsNotNull,
},
);
self.register_scalar("concat", ScalarFunction::String { op: StringOp::Concat });
self.register_scalar("contains", ScalarFunction::String { op: StringOp::Contains });
self.register_scalar(
"starts_with",
ScalarFunction::String {
op: StringOp::StartsWith,
},
);
self.register_scalar("ends_with", ScalarFunction::String { op: StringOp::EndsWith });
self.register_scalar(
"prefix",
ScalarFunction::String {
op: StringOp::StartsWith,
},
);
self.register_scalar("suffix", ScalarFunction::String { op: StringOp::EndsWith });
self.register_scalar("like", ScalarFunction::String { op: StringOp::Like });
self.register_scalar("to_upper", ScalarFunction::String { op: StringOp::ToUpper });
self.register_scalar("to_lower", ScalarFunction::String { op: StringOp::ToLower });
self.register_scalar("upper", ScalarFunction::String { op: StringOp::ToUpper });
self.register_scalar("lower", ScalarFunction::String { op: StringOp::ToLower });
self.register_scalar("ucase", ScalarFunction::String { op: StringOp::ToUpper });
self.register_scalar("lcase", ScalarFunction::String { op: StringOp::ToLower });
self.register_scalar("trim", ScalarFunction::String { op: StringOp::Trim });
self.register_scalar("ltrim", ScalarFunction::String { op: StringOp::LTrim });
self.register_scalar("rtrim", ScalarFunction::String { op: StringOp::RTrim });
self.register_scalar("length", ScalarFunction::String { op: StringOp::Length });
self.register_scalar("reverse", ScalarFunction::String { op: StringOp::Reverse });
self.register_scalar("repeat", ScalarFunction::String { op: StringOp::Repeat });
self.register_scalar("replace", ScalarFunction::String { op: StringOp::Replace });
self.register_scalar(
"substring",
ScalarFunction::String {
op: StringOp::Substring,
},
);
self.register_scalar(
"regex_matches",
ScalarFunction::String {
op: StringOp::RegexMatches,
},
);
self.register_scalar(
"regex_replace",
ScalarFunction::String {
op: StringOp::RegexReplace,
},
);
self.register_scalar("split", ScalarFunction::String { op: StringOp::Split });
self.register_scalar("head", ScalarFunction::String { op: StringOp::Head });
self.register_scalar("tail", ScalarFunction::String { op: StringOp::Tail });
self.register_scalar("left", ScalarFunction::String { op: StringOp::Left });
self.register_scalar("right", ScalarFunction::String { op: StringOp::Right });
self.register_scalar("lpad", ScalarFunction::String { op: StringOp::Lpad });
self.register_scalar("rpad", ScalarFunction::String { op: StringOp::Rpad });
self.register_scalar("initcap", ScalarFunction::String { op: StringOp::InitCap });
self.register_scalar("concat_ws", ScalarFunction::String { op: StringOp::ConcatWs });
self.register_scalar(
"split_part",
ScalarFunction::String {
op: StringOp::SplitPart,
},
);
self.register_scalar(
"array_extract",
ScalarFunction::String {
op: StringOp::ArrayExtract,
},
);
self.register_scalar(
"regexp_full_match",
ScalarFunction::String {
op: StringOp::RegexpFullMatch,
},
);
self.register_scalar(
"regexp_extract",
ScalarFunction::String {
op: StringOp::RegexpExtract,
},
);
self.register_scalar(
"regexp_extract_all",
ScalarFunction::String {
op: StringOp::RegexpExtractAll,
},
);
self.register_scalar(
"regexp_split_to_array",
ScalarFunction::String {
op: StringOp::RegexpSplitToArray,
},
);
self.register_scalar(
"levenshtein",
ScalarFunction::String {
op: StringOp::Levenshtein,
},
);
self.register_scalar("soundex", ScalarFunction::String { op: StringOp::Soundex });
self.register_scalar("md5", ScalarFunction::Hash { op: HashOp::Md5 });
self.register_scalar("sha256", ScalarFunction::Hash { op: HashOp::Sha256 });
self.register_scalar("hash", ScalarFunction::Hash { op: HashOp::Hash });
self.register_scalar(
"to_years",
ScalarFunction::Interval {
op: IntervalOp::ToYears,
},
);
self.register_scalar(
"to_months",
ScalarFunction::Interval {
op: IntervalOp::ToMonths,
},
);
self.register_scalar("to_days", ScalarFunction::Interval { op: IntervalOp::ToDays });
self.register_scalar(
"to_hours",
ScalarFunction::Interval {
op: IntervalOp::ToHours,
},
);
self.register_scalar(
"to_minutes",
ScalarFunction::Interval {
op: IntervalOp::ToMinutes,
},
);
self.register_scalar(
"to_seconds",
ScalarFunction::Interval {
op: IntervalOp::ToSeconds,
},
);
self.register_scalar(
"to_milliseconds",
ScalarFunction::Interval {
op: IntervalOp::ToMilliseconds,
},
);
self.register_scalar(
"to_microseconds",
ScalarFunction::Interval {
op: IntervalOp::ToMicroseconds,
},
);
self.register_scalar("date_part", ScalarFunction::Date { op: DateOp::DatePart });
self.register_scalar("date_trunc", ScalarFunction::Date { op: DateOp::DateTrunc });
self.register_scalar("date_diff", ScalarFunction::Date { op: DateOp::DateDiff });
self.register_scalar("date_add", ScalarFunction::Date { op: DateOp::DateAdd });
self.register_scalar(
"current_date",
ScalarFunction::Date {
op: DateOp::CurrentDate,
},
);
self.register_scalar(
"current_timestamp",
ScalarFunction::Date {
op: DateOp::CurrentTimestamp,
},
);
self.register_scalar("year", ScalarFunction::Date { op: DateOp::Year });
self.register_scalar("month", ScalarFunction::Date { op: DateOp::Month });
self.register_scalar("nextval", ScalarFunction::SequenceOp { is_nextval: true });
self.register_scalar("currval", ScalarFunction::SequenceOp { is_nextval: false });
self.register_scalar("day", ScalarFunction::Date { op: DateOp::Day });
self.register_scalar("hour", ScalarFunction::Date { op: DateOp::Hour });
self.register_scalar("minute", ScalarFunction::Date { op: DateOp::Minute });
self.register_scalar("second", ScalarFunction::Date { op: DateOp::Second });
self.register_scalar("dayname", ScalarFunction::Date { op: DateOp::DayName });
self.register_scalar("monthname", ScalarFunction::Date { op: DateOp::MonthName });
self.register_scalar("last_day", ScalarFunction::Date { op: DateOp::LastDay });
self.register_scalar("make_date", ScalarFunction::Date { op: DateOp::MakeDate });
self.register_scalar("century", ScalarFunction::Date { op: DateOp::Century });
self.register_scalar("epoch_ms", ScalarFunction::Date { op: DateOp::EpochMs });
self.register_scalar(
"to_timestamp",
ScalarFunction::Date {
op: DateOp::ToTimestamp,
},
);
self.register_scalar("to_epoch_ms", ScalarFunction::Date { op: DateOp::ToEpochMs });
self.register_scalar(
"CAST",
ScalarFunction::Cast {
target_type: CastTarget::String,
},
);
self.register_scalar(
"cast_string",
ScalarFunction::Cast {
target_type: CastTarget::String,
},
);
self.register_scalar(
"cast_int64",
ScalarFunction::Cast {
target_type: CastTarget::Int64,
},
);
self.register_scalar(
"cast_double",
ScalarFunction::Cast {
target_type: CastTarget::Double,
},
);
self.register_scalar(
"cast_bool",
ScalarFunction::Cast {
target_type: CastTarget::Bool,
},
);
self.register_scalar(
"date",
ScalarFunction::Cast {
target_type: CastTarget::Date,
},
);
self.register_scalar(
"timestamp",
ScalarFunction::Cast {
target_type: CastTarget::Timestamp,
},
);
self.register_scalar(
"float",
ScalarFunction::Cast {
target_type: CastTarget::Double,
},
);
self.register_scalar(
"double",
ScalarFunction::Cast {
target_type: CastTarget::Double,
},
);
self.register_scalar(
"int64",
ScalarFunction::Cast {
target_type: CastTarget::Int64,
},
);
self.register_scalar(
"int",
ScalarFunction::Cast {
target_type: CastTarget::Int64,
},
);
self.register_scalar(
"bool",
ScalarFunction::Cast {
target_type: CastTarget::Bool,
},
);
self.register_scalar(
"boolean",
ScalarFunction::Cast {
target_type: CastTarget::Bool,
},
);
self.register_scalar(
"string",
ScalarFunction::Cast {
target_type: CastTarget::String,
},
);
self.register_scalar(
"blob",
ScalarFunction::Cast {
target_type: CastTarget::String,
},
);
self.register_scalar("encode", ScalarFunction::Blob { op: BlobOp::Encode });
self.register_scalar("decode", ScalarFunction::Blob { op: BlobOp::Decode });
self.register_scalar(
"octet_length",
ScalarFunction::Blob {
op: BlobOp::OctetLength,
},
);
self.register_scalar("list_creation", ScalarFunction::List { op: ListOp::Creation });
self.register_scalar("list_extract", ScalarFunction::List { op: ListOp::Extract });
self.register_scalar("list_concat", ScalarFunction::List { op: ListOp::Concat });
self.register_scalar("list_cat", ScalarFunction::List { op: ListOp::Concat });
self.register_scalar("list_len", ScalarFunction::List { op: ListOp::Len });
self.register_scalar("list_sort", ScalarFunction::List { op: ListOp::Sort });
self.register_scalar("list_reverse", ScalarFunction::List { op: ListOp::Reverse });
self.register_scalar("list_contains", ScalarFunction::List { op: ListOp::Contains });
self.register_scalar("list_append", ScalarFunction::List { op: ListOp::Append });
self.register_scalar("list_prepend", ScalarFunction::List { op: ListOp::Prepend });
self.register_scalar("list_slice", ScalarFunction::List { op: ListOp::Slice });
self.register_scalar("range", ScalarFunction::List { op: ListOp::Range });
self.register_scalar("list_distinct", ScalarFunction::List { op: ListOp::Distinct });
self.register_scalar("list_unique", ScalarFunction::List { op: ListOp::Unique });
self.register_scalar("list_sum", ScalarFunction::List { op: ListOp::Sum });
self.register_scalar("list_product", ScalarFunction::List { op: ListOp::Product });
self.register_scalar("list_any_value", ScalarFunction::List { op: ListOp::AnyValue });
self.register_scalar("list_to_string", ScalarFunction::List { op: ListOp::ToString });
self.register_scalar("list_position", ScalarFunction::List { op: ListOp::Position });
self.register_scalar("list_indexof", ScalarFunction::List { op: ListOp::Position });
self.register_scalar("list_has_all", ScalarFunction::List { op: ListOp::HasAll });
self.register_scalar("list_has_any", ScalarFunction::List { op: ListOp::HasAny });
self.register_scalar("list_count", ScalarFunction::List { op: ListOp::Count });
self.register_scalar("list_min", ScalarFunction::List { op: ListOp::Min });
self.register_scalar("list_max", ScalarFunction::List { op: ListOp::Max });
self.register_scalar(
"list_reverse_sort",
ScalarFunction::List {
op: ListOp::ReverseSort,
},
);
self.register_scalar("list_transform", ScalarFunction::List { op: ListOp::Transform });
self.register_scalar("list_filter", ScalarFunction::List { op: ListOp::Filter });
self.register_scalar("list_reduce", ScalarFunction::List { op: ListOp::Reduce });
self.register_scalar("any", ScalarFunction::List { op: ListOp::Any });
self.register_scalar("all", ScalarFunction::List { op: ListOp::All });
self.register_scalar("none", ScalarFunction::List { op: ListOp::None });
self.register_scalar("single", ScalarFunction::List { op: ListOp::Single });
self.register_scalar("map_creation", ScalarFunction::Map { op: MapOp::Creation });
self.register_scalar("map_extract", ScalarFunction::Map { op: MapOp::Extract });
self.register_scalar("element_at", ScalarFunction::Map { op: MapOp::Extract });
self.register_scalar("map_keys", ScalarFunction::Map { op: MapOp::Keys });
self.register_scalar("map_values", ScalarFunction::Map { op: MapOp::Values });
self.register_scalar("struct_creation", ScalarFunction::Struct { op: StructOp::Creation });
self.register_scalar("struct_extract", ScalarFunction::Struct { op: StructOp::Extract });
self.register_scalar(
"union_value",
ScalarFunction::Union {
op: UnionOp::UnionValue,
},
);
self.register_scalar(
"union_extract",
ScalarFunction::Union {
op: UnionOp::UnionExtract,
},
);
self.register_scalar("union_tag", ScalarFunction::Union { op: UnionOp::UnionTag });
self.register_scalar("AND", ScalarFunction::Boolean { op: BooleanOp::And });
self.register_scalar("OR", ScalarFunction::Boolean { op: BooleanOp::Or });
self.register_scalar("XOR", ScalarFunction::Boolean { op: BooleanOp::Xor });
self.register_scalar("NOT", ScalarFunction::Boolean { op: BooleanOp::Not });
self.register_scalar(
"coalesce",
ScalarFunction::Utility {
op: UtilityOp::Coalesce,
},
);
self.register_scalar("ifnull", ScalarFunction::Utility { op: UtilityOp::IfNull });
self.register_scalar("nullif", ScalarFunction::Utility { op: UtilityOp::NullIf });
self.register_scalar("size", ScalarFunction::Utility { op: UtilityOp::Size });
self.register_scalar("cardinality", ScalarFunction::Utility { op: UtilityOp::Size });
self.register_scalar("typeof", ScalarFunction::Utility { op: UtilityOp::TypeOf });
self.register_scalar("error", ScalarFunction::Utility { op: UtilityOp::Error });
self.register_scalar(
"pg_isready",
ScalarFunction::Utility {
op: UtilityOp::PgIsReady,
},
);
self.register_scalar(
"greatest",
ScalarFunction::Utility {
op: UtilityOp::Greatest,
},
);
self.register_scalar("least", ScalarFunction::Utility { op: UtilityOp::Least });
self.register_scalar(
"constant_or_null",
ScalarFunction::Utility {
op: UtilityOp::ConstantOrNull,
},
);
self.register_scalar("OFFSET", ScalarFunction::Schema { op: SchemaOp::Offset });
self.register_scalar("ID", ScalarFunction::Schema { op: SchemaOp::Id });
self.register_scalar(
"START_NODE",
ScalarFunction::Schema {
op: SchemaOp::StartNode,
},
);
self.register_scalar("END_NODE", ScalarFunction::Schema { op: SchemaOp::EndNode });
self.register_scalar("LABEL", ScalarFunction::Schema { op: SchemaOp::Label });
self.register_scalar("COST", ScalarFunction::Schema { op: SchemaOp::Cost });
self.register_scalar("ROWID", ScalarFunction::Schema { op: SchemaOp::RowId });
self.register_scalar(
"array_cosine_similarity",
ScalarFunction::Array {
op: ArrayOp::CosineSimilarity,
},
);
self.register_scalar("array_distance", ScalarFunction::Array { op: ArrayOp::Distance });
self.register_scalar(
"array_inner_product",
ScalarFunction::Array {
op: ArrayOp::InnerProduct,
},
);
self.register_scalar(
"array_dot_product",
ScalarFunction::Array {
op: ArrayOp::DotProduct,
},
);
self.register_scalar(
"array_cross_product",
ScalarFunction::Array {
op: ArrayOp::CrossProduct,
},
);
self.register_scalar(
"array_squared_distance",
ScalarFunction::Array {
op: ArrayOp::SquaredDistance,
},
);
self.register_scalar("array_intersect", ScalarFunction::Array { op: ArrayOp::Intersect });
self.register_scalar("nodes", ScalarFunction::Path { op: PathOp::Nodes });
self.register_scalar("rels", ScalarFunction::Path { op: PathOp::Rels });
self.register_scalar("relationships", ScalarFunction::Path { op: PathOp::Rels });
self.register_scalar("properties", ScalarFunction::Path { op: PathOp::Properties });
self.register_scalar("is_trail", ScalarFunction::Path { op: PathOp::IsTrail });
self.register_scalar("is_acyclic", ScalarFunction::Path { op: PathOp::IsAcyclic });
self.register_scalar("gen_random_uuid", ScalarFunction::Uuid);
self.register_scalar(
"map_from_entries",
ScalarFunction::Map {
op: MapOp::MapFromEntries,
},
);
self.register_scalar(
"blob_from_bytes",
ScalarFunction::Blob {
op: BlobOp::BlobFromBytes,
},
);
self.register_scalar("to_base64", ScalarFunction::Blob { op: BlobOp::ToBase64 });
self.register_scalar("from_base64", ScalarFunction::Blob { op: BlobOp::FromBase64 });
self.register_scalar("array_concat", ScalarFunction::List { op: ListOp::Concat });
self.register_scalar("array_cat", ScalarFunction::List { op: ListOp::Concat });
self.register_scalar("array_append", ScalarFunction::List { op: ListOp::Append });
self.register_scalar("array_push_back", ScalarFunction::List { op: ListOp::Append });
self.register_scalar("array_prepend", ScalarFunction::List { op: ListOp::Prepend });
self.register_scalar("array_push_front", ScalarFunction::List { op: ListOp::Prepend });
self.register_scalar("array_contains", ScalarFunction::List { op: ListOp::Contains });
self.register_scalar("array_has", ScalarFunction::List { op: ListOp::Contains });
self.register_scalar("array_slice", ScalarFunction::List { op: ListOp::Slice });
self.register_scalar("array_value", ScalarFunction::List { op: ListOp::Creation });
self.register_aggregate("COUNT", AggregateFunction::Count);
self.register_aggregate("COUNT(*)", AggregateFunction::CountStar);
self.register_aggregate("COUNT_IF", AggregateFunction::CountIf);
self.register_aggregate("SUM", AggregateFunction::Sum);
self.register_aggregate("AVG", AggregateFunction::Avg);
self.register_aggregate("MIN", AggregateFunction::Min);
self.register_aggregate("MAX", AggregateFunction::Max);
self.register_aggregate("COLLECT", AggregateFunction::Collect);
self.register_aggregate("STDDEV", AggregateFunction::StdDev);
self.register_aggregate("VARIANCE", AggregateFunction::Variance);
self.register_aggregate(
"STRING_AGG",
AggregateFunction::StringAgg {
delimiter: ",".to_string(),
},
);
self.register_aggregate(
"GROUP_CONCAT",
AggregateFunction::StringAgg {
delimiter: ",".to_string(),
},
);
self.register_aggregate("PERCENTILE_DISC", AggregateFunction::PercentileDisc { percentile: 0.5 });
self.register_aggregate("PERCENTILE_CONT", AggregateFunction::PercentileCont { percentile: 0.5 });
self.register_table("list_tables", TableFunction::ListTables);
}
pub fn register_scalar(&mut self, name: &str, func: ScalarFunction) {
self.scalar_functions.insert(name.to_lowercase(), func);
}
pub fn register_aggregate(&mut self, name: &str, func: AggregateFunction) {
self.aggregate_functions.insert(name.to_lowercase(), func);
}
pub fn register_table(&mut self, name: &str, func: TableFunction) {
self.table_functions.insert(name.to_lowercase(), func);
}
pub fn resolve(&self, name: &str) -> Option<ResolvedFunction> {
let lower = name.to_lowercase();
if let Some(f) = self.scalar_functions.get(&lower) {
return Some(ResolvedFunction::Scalar(f.clone()));
}
if let Some(f) = self.aggregate_functions.get(&lower) {
return Some(ResolvedFunction::Aggregate(f.clone()));
}
if let Some(f) = self.table_functions.get(&lower) {
return Some(ResolvedFunction::Table(f.clone()));
}
None
}
pub fn get_scalar(&self, name: &str) -> Option<&ScalarFunction> {
self.scalar_functions.get(&name.to_lowercase())
}
pub fn get_aggregate(&self, name: &str) -> Option<&AggregateFunction> {
self.aggregate_functions.get(&name.to_lowercase())
}
pub fn get_table(&self, name: &str) -> Option<&TableFunction> {
self.table_functions.get(&name.to_lowercase())
}
pub fn list_all(&self) -> Vec<(String, String)> {
let mut result = Vec::new();
for name in self.scalar_functions.keys() {
result.push((name.clone(), "SCALAR".to_string()));
}
for name in self.aggregate_functions.keys() {
result.push((name.clone(), "AGGREGATE".to_string()));
}
for name in self.table_functions.keys() {
result.push((name.clone(), "TABLE".to_string()));
}
result.sort_by(|a, b| a.0.cmp(&b.0));
result
}
pub fn contains(&self, name: &str) -> bool {
let lower = name.to_lowercase();
self.scalar_functions.contains_key(&lower)
|| self.aggregate_functions.contains_key(&lower)
|| self.table_functions.contains_key(&lower)
}
pub fn scalar_count(&self) -> usize {
self.scalar_functions.len()
}
pub fn aggregate_count(&self) -> usize {
self.aggregate_functions.len()
}
pub fn table_count(&self) -> usize {
self.table_functions.len()
}
pub fn total_count(&self) -> usize {
self.scalar_count() + self.aggregate_count() + self.table_count()
}
pub fn execute_table_function(
&self,
name: &str,
args: &[Value],
graph: Option<&dyn crate::graph::GraphDataSource>,
) -> Result<Vec<Vec<Value>>, String> {
use akar_common::vector::DataChunk;
let func = self
.get_table(name)
.ok_or_else(|| format!("Table function '{}' not found", name))?;
match func {
TableFunction::ListTables => Err("ListTables requires catalog access — handled at connection level".into()),
TableFunction::ShowColumns { .. } => {
Err("ShowColumns requires catalog access — handled at connection level".into())
}
TableFunction::Custom { name: custom_name } => Err(format!(
"Table function '{}' requires an extension or external context to be loaded. \
Use LOAD EXTENSION or CALL with the appropriate handler.",
custom_name
)),
TableFunction::CustomTable { name: _, execute } => {
let mut chunk = DataChunk {
fields: Vec::new(),
field_types: Vec::new(),
size: 0,
field_names: vec![],
sel_vector: None,
};
execute(args, &mut chunk).map(|_| {
let mut rows = Vec::new();
for row in 0..chunk.size {
let mut row_vals = Vec::new();
for field_idx in 0..chunk.fields.len() {
row_vals.push(chunk.get_value(field_idx, row).unwrap_or(Value::Null));
}
rows.push(row_vals);
}
rows
})
}
TableFunction::CustomTableWithGraph { name: _, execute } => {
let mut chunk = DataChunk {
fields: Vec::new(),
field_types: Vec::new(),
size: 0,
field_names: vec![],
sel_vector: None,
};
execute(args, graph, &mut chunk).map(|_| {
let mut rows = Vec::new();
for row in 0..chunk.size {
let mut row_vals = Vec::new();
for field_idx in 0..chunk.fields.len() {
row_vals.push(chunk.get_value(field_idx, row).unwrap_or(Value::Null));
}
rows.push(row_vals);
}
rows
})
}
TableFunction::ScanCsv { .. }
| TableFunction::ScanParquet { .. }
| TableFunction::ScanJson { .. }
| TableFunction::CurrentSetting { .. } => Err(format!(
"Table function '{}' cannot be executed via CALL — it requires file path or catalog context. \
Use COPY FROM 'file' FORMAT CSV/PARQUET/JSON or CALL current_setting('key') via the connection layer.",
name
)),
}
}
}