"""Vector distance helpers and SQL generator registration."""
# ruff: noqa: N802
# pyright: ignore[reportConstantRedefinition]
from collections.abc import Callable, MutableMapping
from typing import TYPE_CHECKING, Any, Final, cast
from sqlglot import exp
from sqlspec.builder._generation import invalidate_generator_dispatch
if TYPE_CHECKING:
from sqlglot.generator import Generator
__all__ = (
"VectorDistance",
"has_vector_distance_ancestor",
"is_vector_distance_expression",
"render_vector_distance_bigquery",
"render_vector_distance_duckdb",
"render_vector_distance_generic",
"render_vector_distance_mysql",
"render_vector_distance_oracle",
"render_vector_distance_postgres",
"vector_distance_metric",
)
_VECTOR_DISTANCE_META_KEY: Final[str] = "sqlspec_vector_distance_metric"
_OperatorTransform = Callable[[Any, exp.Operator], str]
_SQLGLOT_VECTOR_DISTANCE_REGISTERED = False
_BASE_OPERATOR_TRANSFORM: _OperatorTransform | None = None
_POSTGRES_OPERATOR_TRANSFORM: _OperatorTransform | None = None
_MYSQL_OPERATOR_TRANSFORM: _OperatorTransform | None = None
_ORACLE_OPERATOR_TRANSFORM: _OperatorTransform | None = None
_BIGQUERY_OPERATOR_TRANSFORM: _OperatorTransform | None = None
_DUCKDB_OPERATOR_TRANSFORM: _OperatorTransform | None = None
def is_vector_distance_expression(expression: object) -> bool:
"""Return True when an Operator node is a SQLSpec vector-distance expression."""
return isinstance(expression, exp.Operator) and _VECTOR_DISTANCE_META_KEY in expression.meta
def has_vector_distance_ancestor(expression: exp.Expr) -> bool:
"""Return True when any ancestor is a SQLSpec vector-distance expression."""
parent = expression.parent
while parent is not None:
if is_vector_distance_expression(parent):
return True
parent = parent.parent
return False
def vector_distance_metric(expression: object) -> str:
"""Get the normalized vector-distance metric from an Operator node."""
if not isinstance(expression, exp.Operator):
msg = f"Expected sqlglot Operator, got {type(expression)}"
raise TypeError(msg)
metric = expression.meta.get(_VECTOR_DISTANCE_META_KEY)
if isinstance(metric, str):
return metric
operator = expression.args.get("operator")
return str(operator).lower() if operator is not None else "euclidean"
def _normalize_metric(metric: Any) -> str:
"""Normalize vector metrics to a lowercase string."""
if isinstance(metric, exp.Literal):
return str(metric.this).lower()
if isinstance(metric, exp.Identifier):
identifier = metric.this
return identifier.lower() if isinstance(identifier, str) else "euclidean"
if isinstance(metric, str):
return metric.lower()
return "euclidean"
def _build_vector_distance(this: exp.Expr, expression: exp.Expr, metric: Any = "euclidean") -> exp.Operator:
normalized_metric = _normalize_metric(metric)
node = exp.Operator(this=this, expression=expression, operator=normalized_metric)
node.meta[_VECTOR_DISTANCE_META_KEY] = normalized_metric
return node
def VectorDistance(*, this: exp.Expr, expression: exp.Expr, metric: Any = "euclidean") -> exp.Operator:
"""Build a SQLSpec vector-distance expression."""
_register_with_sqlglot()
return _build_vector_distance(this=this, expression=expression, metric=metric)