use rudb_common::{Error, Field, LogicalType, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TableFunction {
Range,
GenerateSeries,
ReadParquet,
ReadCsv,
}
impl TableFunction {
#[must_use]
pub const fn name(self) -> &'static str {
match self {
Self::Range => "range",
Self::GenerateSeries => "generate_series",
Self::ReadParquet => "read_parquet",
Self::ReadCsv => "read_csv",
}
}
#[must_use]
pub const fn inclusive(self) -> bool {
matches!(self, Self::GenerateSeries)
}
#[must_use]
pub fn parameters(self) -> &'static [(&'static str, LogicalType)] {
static READ_PARQUET: &[(&str, LogicalType)] = &[("binary_as_string", LogicalType::Boolean)];
static READ_CSV: &[(&str, LogicalType)] = &[
("all_varchar", LogicalType::Boolean),
("delim", LogicalType::Varchar),
("escape", LogicalType::Varchar),
("header", LogicalType::Boolean),
("quote", LogicalType::Varchar),
("sep", LogicalType::Varchar),
];
match self {
Self::ReadParquet => READ_PARQUET,
Self::ReadCsv => READ_CSV,
_ => &[],
}
}
#[must_use]
pub fn lookup(name: &str) -> Option<Self> {
if name.eq_ignore_ascii_case("range") {
return Some(Self::Range);
}
if name.eq_ignore_ascii_case("generate_series") {
return Some(Self::GenerateSeries);
}
if name.eq_ignore_ascii_case("read_parquet") || name.eq_ignore_ascii_case("parquet_scan") {
return Some(Self::ReadParquet);
}
if name.eq_ignore_ascii_case("read_csv") || name.eq_ignore_ascii_case("read_csv_auto") {
return Some(Self::ReadCsv);
}
None
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Columns {
Fixed(Vec<Field>),
Parquet,
Csv,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ResolvedTable {
pub function: TableFunction,
pub arguments: Vec<LogicalType>,
pub columns: Columns,
}
pub fn resolve_table(name: &str, arguments: &[LogicalType]) -> Result<ResolvedTable> {
let Some(function) = TableFunction::lookup(name) else {
return Err(Error::catalog(format!("Table Function with name {name} does not exist!")));
};
if let Some(columns) = file_columns(function) {
let list = LogicalType::list(LogicalType::Varchar);
let single = arguments.len() == 1 && arguments[0] == LogicalType::Varchar;
let many = arguments.len() == 1 && arguments[0] == list;
let nothing = arguments.len() == 1 && arguments[0] == LogicalType::Null;
if !single && !many && !nothing {
return Err(no_overload(function, arguments));
}
let wanted = if many {
list
} else if nothing {
LogicalType::Null
} else {
LogicalType::Varchar
};
return Ok(ResolvedTable { function, arguments: vec![wanted], columns });
}
let arity = arguments.len();
if !(1..=3).contains(&arity) {
return Err(Error::binder(format!(
"Table function {}() takes between 1 and 3 arguments, {arity} were given",
function.name()
)));
}
Ok(ResolvedTable {
function,
arguments: vec![LogicalType::BigInt; arity],
columns: Columns::Fixed(vec![Field::new(function.name(), LogicalType::BigInt)]),
})
}
fn file_columns(function: TableFunction) -> Option<Columns> {
match function {
TableFunction::ReadParquet => Some(Columns::Parquet),
TableFunction::ReadCsv => Some(Columns::Csv),
TableFunction::Range | TableFunction::GenerateSeries => None,
}
}
fn no_overload(function: TableFunction, arguments: &[LogicalType]) -> Error {
let written: Vec<String> = arguments.iter().map(ToString::to_string).collect();
let name = function.name();
Error::binder(format!(
"No function matches the given name and argument types '{name}({})'. You might need to \
add explicit type casts.\n\tCandidate functions:\n\t{name}(VARCHAR)\n\t{name}(VARCHAR[])\n",
written.join(", ")
))
}
pub fn series(function: TableFunction, start: i64, stop: i64, step: i64) -> Result<Vec<i64>> {
let count = series_length(function, start, stop, step)?;
let mut out = Vec::with_capacity(count);
let mut at = start;
for _ in 0..count {
out.push(at);
at = at.saturating_add(step);
}
Ok(out)
}
pub fn series_length(function: TableFunction, start: i64, stop: i64, step: i64) -> Result<usize> {
if step == 0 {
return Err(Error::binder("interval cannot be 0!"));
}
Ok(length(function, start, stop, step))
}
fn length(function: TableFunction, start: i64, stop: i64, step: i64) -> usize {
let start = i128::from(start);
let stop = i128::from(stop);
let step = i128::from(step);
let span = if function.inclusive() {
if step > 0 { stop - start + 1 } else { stop - start - 1 }
} else {
stop - start
};
if (span > 0) != (step > 0) {
return 0;
}
let count = (span + step - step.signum()) / step;
usize::try_from(count).unwrap_or(usize::MAX)
}
#[cfg(test)]
mod tests {
use super::*;
fn fixed(resolved: &ResolvedTable) -> &[Field] {
match &resolved.columns {
Columns::Fixed(fields) => fields,
Columns::Parquet | Columns::Csv => {
panic!("{} resolves to a file", resolved.function.name())
}
}
}
fn integers(count: usize) -> Vec<LogicalType> {
vec![LogicalType::BigInt; count]
}
#[test]
fn a_name_that_is_not_a_table_function_says_so_rather_than_binding() {
let error = resolve_table("read_csv", &integers(1)).unwrap_err();
assert!(error.to_string().contains("read_csv"), "{error}");
}
#[test]
fn both_names_resolve_and_each_one_names_its_own_column() {
let range = resolve_table("range", &integers(1)).unwrap();
assert_eq!(fixed(&range)[0].name, "range");
let series = resolve_table("GENERATE_SERIES", &integers(3)).unwrap();
assert_eq!(fixed(&series)[0].name, "generate_series");
assert_eq!(series.arguments.len(), 3);
}
#[test]
fn no_arguments_and_four_arguments_are_both_the_arity_error() {
assert!(resolve_table("range", &integers(0)).is_err());
assert!(resolve_table("range", &integers(4)).is_err());
}
#[test]
fn a_series_call_ignores_the_types_it_was_given_and_casts_them_all_to_bigint() {
let resolved =
resolve_table("range", &[LogicalType::Varchar, LogicalType::Double]).unwrap();
assert_eq!(resolved.arguments, integers(2));
}
#[test]
fn read_parquet_takes_one_string_and_says_its_columns_are_in_the_file() {
let resolved = resolve_table("read_parquet", &[LogicalType::Varchar]).unwrap();
assert_eq!(resolved.function, TableFunction::ReadParquet);
assert_eq!(resolved.arguments, vec![LogicalType::Varchar]);
assert_eq!(resolved.columns, Columns::Parquet);
}
#[test]
fn parquet_scan_is_the_same_function_under_duckdbs_other_name_for_it() {
assert_eq!(TableFunction::lookup("parquet_scan"), Some(TableFunction::ReadParquet));
let resolved = resolve_table("parquet_scan", &[LogicalType::Varchar]).unwrap();
assert_eq!(resolved.function.name(), "read_parquet");
}
#[test]
fn a_path_that_is_not_a_string_is_the_message_duckdb_gives_for_it() {
let error = resolve_table("read_parquet", &[LogicalType::Integer]).unwrap_err();
assert!(
error.message().starts_with(
"No function matches the given name and argument types 'read_parquet(INTEGER)'."
),
"{error}"
);
assert!(error.message().contains("read_parquet(VARCHAR)"), "{error}");
}
#[test]
fn read_parquet_of_no_arguments_or_two_is_the_same_no_overload_message() {
let two = resolve_table("read_parquet", &[LogicalType::Varchar, LogicalType::Varchar]);
assert!(two.unwrap_err().message().contains("read_parquet(VARCHAR, VARCHAR)"));
let none = resolve_table("read_parquet", &[]);
assert!(none.unwrap_err().message().contains("read_parquet()"));
}
#[test]
fn range_stops_before_the_end_and_generate_series_stops_on_it() {
assert_eq!(series(TableFunction::Range, 0, 3, 1).unwrap(), vec![0, 1, 2]);
assert_eq!(series(TableFunction::GenerateSeries, 0, 3, 1).unwrap(), vec![0, 1, 2, 3]);
}
#[test]
fn a_step_that_does_not_divide_the_span_stops_before_the_end_of_it() {
assert_eq!(series(TableFunction::Range, 2, 7, 2).unwrap(), vec![2, 4, 6]);
assert_eq!(series(TableFunction::GenerateSeries, 2, 7, 2).unwrap(), vec![2, 4, 6]);
}
#[test]
fn a_negative_step_counts_down_and_stops_on_the_same_rule() {
assert_eq!(series(TableFunction::Range, 5, 1, -2).unwrap(), vec![5, 3]);
assert_eq!(series(TableFunction::GenerateSeries, 5, 1, -2).unwrap(), vec![5, 3, 1]);
}
#[test]
fn a_step_going_the_wrong_way_produces_nothing_rather_than_running_forever() {
assert!(series(TableFunction::Range, 0, 10, -1).unwrap().is_empty());
assert!(series(TableFunction::Range, 10, 0, 1).unwrap().is_empty());
}
#[test]
fn an_empty_range_and_a_single_value_series_are_the_boundary_between_the_two() {
assert!(series(TableFunction::Range, 4, 4, 1).unwrap().is_empty());
assert_eq!(series(TableFunction::GenerateSeries, 4, 4, 1).unwrap(), vec![4]);
}
#[test]
fn a_step_of_zero_is_the_one_case_that_is_an_error_rather_than_nothing() {
let error = series(TableFunction::Range, 1, 5, 0).unwrap_err();
assert!(error.to_string().contains("interval cannot be 0"), "{error}");
}
#[test]
fn a_span_that_does_not_fit_in_an_i64_does_not_overflow_the_length() {
assert_eq!(length(TableFunction::Range, i64::MIN, i64::MAX, 1), usize::MAX);
}
}