use super::{
multi_field_match_shape, Engine, MultiFieldMatchShape, SQLError, ScalarExpr, SourcePlan, Value,
};
use uqa_sql::registry::FunctionKind;
const SINGLE_FIELD_TEXT_MATCH_FUNCTIONS: [&str; 4] = [
"text_match",
"bayesian_match",
"fts_match",
"bayesian_match_with_prior",
];
pub(in crate::sql) fn is_semantic_field_argument(
function: &str,
args: &[ScalarExpr],
argument_index: usize,
) -> Result<bool, SQLError> {
let dispatch_name = crate::sql::builtin_function_dispatch_name(function);
let Some(kind) = uqa_sql::registry::lookup(&dispatch_name) else {
return Ok(false);
};
let is_field = match kind {
FunctionKind::TextMatch | FunctionKind::BayesianMatch | FunctionKind::KNNMatch => {
argument_index == 0
}
FunctionKind::FTSMatch => argument_index == 0 && !fts_query_is_jsonpath(args.get(1)),
FunctionKind::BayesianMatchWithPrior => matches!(argument_index, 0 | 2),
FunctionKind::CalibratedVectorMatch => argument_index == 0,
FunctionKind::MultiFieldMatch => match multi_field_match_shape(args)? {
MultiFieldMatchShape::FieldsThenQuery { query_idx, .. } => argument_index < query_idx,
MultiFieldMatchShape::Pairs { .. } => argument_index.is_multiple_of(2),
},
FunctionKind::StagedRetrieval => {
!matches!(args.first(), Some(ScalarExpr::Func { .. }))
&& argument_index.is_multiple_of(3)
}
FunctionKind::UQAFacets => true,
FunctionKind::ScoreBM25 | FunctionKind::ScoreBayesianBM25 => {
args.len() == 2 && argument_index == 0
}
FunctionKind::FuseLogOdds
| FunctionKind::PositiveEvidencePool
| FunctionKind::BayesianEvidenceFusion
| FunctionKind::GraphPagerank
| FunctionKind::GraphHits
| FunctionKind::GraphBetweenness
| FunctionKind::GraphTraverse
| FunctionKind::GraphNeighbors
| FunctionKind::DeepPredict
| FunctionKind::UQAHighlight
| FunctionKind::TraverseMatch
| FunctionKind::TemporalTraverse
| FunctionKind::RPQ
| FunctionKind::GraphCreate
| FunctionKind::GraphDrop
| FunctionKind::GraphExists
| FunctionKind::GraphLabelCreate
| FunctionKind::GraphLabelDrop
| FunctionKind::GraphAlter
| FunctionKind::GraphEdges
| FunctionKind::AttentionFusion
| FunctionKind::LearnedFusion
| FunctionKind::SparseThreshold
| FunctionKind::DeepLearn
| FunctionKind::Convolve
| FunctionKind::Pool
| FunctionKind::Flatten
| FunctionKind::Dense
| FunctionKind::Softmax
| FunctionKind::Layer
| FunctionKind::Model => false,
};
Ok(is_field)
}
fn walk_text_match_fields(
expr: &ScalarExpr,
validate: &mut dyn FnMut(&ScalarExpr, &str) -> Result<(), SQLError>,
) -> Result<(), SQLError> {
match expr {
ScalarExpr::Func {
name, args, filter, ..
} => {
let lower = name.to_ascii_lowercase();
if SINGLE_FIELD_TEXT_MATCH_FUNCTIONS.contains(&lower.as_str()) {
if let Some(field_arg) = args.first() {
if !(lower == "fts_match" && fts_query_is_jsonpath(args.get(1))) {
validate(field_arg, &lower)?;
}
}
} else if lower == "multi_field_match" {
match multi_field_match_shape(args)? {
MultiFieldMatchShape::FieldsThenQuery { fields, .. }
| MultiFieldMatchShape::Pairs { fields } => {
for field_arg in fields {
validate(field_arg, "multi_field_match")?;
}
}
}
}
for arg in args {
walk_text_match_fields(arg, validate)?;
}
if let Some(filter) = filter {
walk_text_match_fields(filter, validate)?;
}
Ok(())
}
ScalarExpr::And(items)
| ScalarExpr::Or(items)
| ScalarExpr::Array(items)
| ScalarExpr::Row(items) => {
for item in items {
walk_text_match_fields(item, validate)?;
}
Ok(())
}
ScalarExpr::Not(inner) | ScalarExpr::UnaryMinus(inner) => {
walk_text_match_fields(inner, validate)
}
ScalarExpr::Binary { lhs, rhs, .. } => {
walk_text_match_fields(lhs, validate)?;
walk_text_match_fields(rhs, validate)
}
ScalarExpr::IsNull { expr, .. } => walk_text_match_fields(expr, validate),
ScalarExpr::Between { expr, low, high } => {
walk_text_match_fields(expr, validate)?;
walk_text_match_fields(low, validate)?;
walk_text_match_fields(high, validate)
}
ScalarExpr::InList { expr, list, .. } => {
walk_text_match_fields(expr, validate)?;
for item in list {
walk_text_match_fields(item, validate)?;
}
Ok(())
}
_ => Ok(()),
}
}
pub(in crate::sql) fn validate_expr_text_match_fields(
engine: &Engine,
table: &str,
expr: &ScalarExpr,
) -> Result<(), SQLError> {
walk_text_match_fields(
expr,
&mut |field_arg, function_name| match text_match_field_name(field_arg) {
Some(TextMatchField::All) => {
validate_text_match_all_fields(engine, table, function_name)
}
Some(TextMatchField::Named(field)) => {
validate_text_match_field(engine, table, field, function_name)
}
None => Ok(()),
},
)
}
enum TextMatchField<'a> {
All,
Named(&'a str),
}
fn text_match_field_name(field_arg: &ScalarExpr) -> Option<TextMatchField<'_>> {
match field_arg {
ScalarExpr::Column(name) | ScalarExpr::QualifiedColumn { column: name, .. } => {
if name.is_empty() || name == "_all" {
Some(TextMatchField::All)
} else {
Some(TextMatchField::Named(name))
}
}
ScalarExpr::Literal(Value::Str(s)) if s.is_empty() || s == "_all" => {
Some(TextMatchField::All)
}
_ => None,
}
}
pub(in crate::sql) fn validate_joined_expr_text_match_fields(
engine: &Engine,
from: &SourcePlan,
expr: &ScalarExpr,
) -> Result<(), SQLError> {
let mut tables: Vec<(Option<String>, String, Vec<String>)> = Vec::new();
let mut has_opaque_source = false;
collect_from_tables(from, &mut tables, &mut has_opaque_source);
walk_text_match_fields(expr, &mut |field_arg, function_name| {
let (qualifier, column) = match field_arg {
ScalarExpr::Column(name) => (None, name.as_str()),
ScalarExpr::QualifiedColumn {
qualifier, column, ..
} => (Some(qualifier.as_str()), column.as_str()),
_ => return Ok(()),
};
if column.is_empty() || column == "_all" {
return Ok(());
}
if let Some(qualifier) = qualifier {
let resolved = tables
.iter()
.find(|(alias, name, _)| alias.as_deref() == Some(qualifier) || name == qualifier);
return match resolved {
Some((_, table, aliases)) => {
let physical = table_source_physical_column(engine, table, aliases, column)?
.unwrap_or_else(|| column.to_string());
validate_text_match_field(engine, table, &physical, function_name)
}
None => Ok(()),
};
}
let mut containing = Vec::new();
for (_, name, aliases) in &tables {
if let Some(physical) = table_source_physical_column(engine, name, aliases, column)? {
containing.push((name, physical));
}
}
for (name, physical) in &containing {
if engine
.fts_fields_for_table(name)?
.iter()
.any(|field| field == physical)
{
return Ok(());
}
}
if let Some((table, physical)) = containing.first() {
return validate_text_match_field(engine, table, physical, function_name);
}
if has_opaque_source {
return Ok(());
}
Err(SQLError::TypeMismatch(format!(
"{function_name}: column `{column}` does not exist on any joined table"
)))
})
}
fn table_source_physical_column(
engine: &Engine,
table: &str,
aliases: &[String],
visible: &str,
) -> Result<Option<String>, SQLError> {
let columns = engine
.try_query_table_columns(table)
.map_err(|error| SQLError::Internal(format!("read table schema: {error}")))?;
if columns.is_empty() {
return Ok(Some(visible.to_string()));
}
Ok(columns
.into_iter()
.enumerate()
.find_map(|(position, physical)| {
aliases
.get(position)
.map_or_else(
|| physical.eq_ignore_ascii_case(visible),
|alias| alias.eq_ignore_ascii_case(visible),
)
.then_some(physical)
}))
}
fn fts_query_is_jsonpath(query_arg: Option<&ScalarExpr>) -> bool {
matches!(
query_arg,
Some(ScalarExpr::Literal(Value::Str(path))) if path.trim_start().starts_with('$')
)
}
fn collect_from_tables(
from: &SourcePlan,
out: &mut Vec<(Option<String>, String, Vec<String>)>,
has_opaque_source: &mut bool,
) {
match from {
SourcePlan::Table {
name,
qualifier,
alias,
column_aliases,
..
} => out.push((
Some(alias.as_ref().unwrap_or(qualifier).clone()),
name.clone(),
column_aliases.clone(),
)),
SourcePlan::Join {
left, right, alias, ..
} => {
if alias.is_some() {
*has_opaque_source = true;
} else {
collect_from_tables(left, out, has_opaque_source);
collect_from_tables(right, out, has_opaque_source);
}
}
_ => *has_opaque_source = true,
}
}
pub(in crate::sql) fn validate_text_match_field(
engine: &Engine,
table: &str,
field: &str,
function_name: &str,
) -> Result<(), SQLError> {
if !engine
.try_query_has_table(table)
.map_err(|err| SQLError::Internal(format!("read table catalog: {err}")))?
{
return Err(SQLError::TypeMismatch(format!(
"{function_name}: unknown table `{table}`"
)));
}
let indexed = engine
.fts_fields_for_table(table)?
.iter()
.any(|fts| fts == field);
if !indexed {
if !engine
.try_query_table_has_column(table, field)
.map_err(|err| SQLError::Internal(format!("read table schema: {err}")))?
&& !engine
.try_query_table_columns(table)
.map_err(|err| SQLError::Internal(format!("read table schema: {err}")))?
.is_empty()
{
return Err(SQLError::TypeMismatch(format!(
"{function_name}: column `{field}` does not exist on table `{table}`"
)));
}
return Err(SQLError::TypeMismatch(format!(
"{function_name}: column `{table}.{field}` has no text index; \
create one with CREATE INDEX ... ON {table} USING gin ({field})"
)));
}
Ok(())
}
pub(in crate::sql) fn validate_text_match_all_fields(
engine: &Engine,
table: &str,
function_name: &str,
) -> Result<(), SQLError> {
if !engine
.try_query_has_table(table)
.map_err(|err| SQLError::Internal(format!("read table catalog: {err}")))?
{
return Err(SQLError::TypeMismatch(format!(
"{function_name}: unknown table `{table}`"
)));
}
if engine.fts_fields_for_table(table)?.is_empty() {
return Err(SQLError::TypeMismatch(format!(
"{function_name}: table `{table}` has no text-indexed columns; \
create one with CREATE INDEX ... ON {table} USING gin (...)"
)));
}
Ok(())
}